feat: add cache_salt dimension to LLM response cache
This commit is contained in:
+8
-2
@@ -274,6 +274,7 @@ class GovernedLLMClient:
|
|||||||
*,
|
*,
|
||||||
session_id: str | None = None,
|
session_id: str | None = None,
|
||||||
parent_call_id: str | None = None,
|
parent_call_id: str | None = None,
|
||||||
|
cache_salt: str | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
"""发起 LLM 调用,经四层治理栈:熔断 → 缓存 → 重试+流式 → 遥测。
|
"""发起 LLM 调用,经四层治理栈:熔断 → 缓存 → 重试+流式 → 遥测。
|
||||||
|
|
||||||
@@ -281,6 +282,7 @@ class GovernedLLMClient:
|
|||||||
messages: OpenAI 格式消息列表。
|
messages: OpenAI 格式消息列表。
|
||||||
session_id: 会话 ID(传递到遥测)。
|
session_id: 会话 ID(传递到遥测)。
|
||||||
parent_call_id: 父调用 ID(传递到遥测)。
|
parent_call_id: 父调用 ID(传递到遥测)。
|
||||||
|
cache_salt: 可选缓存盐,透传到 Redis 缓存键(如跨 epoch 重采样)。
|
||||||
|
|
||||||
返回:
|
返回:
|
||||||
LLMResponse 统一响应。
|
LLMResponse 统一响应。
|
||||||
@@ -296,7 +298,11 @@ class GovernedLLMClient:
|
|||||||
raise CircuitOpenError(f"熔断器已开启,拒绝调用 provider={self._provider}")
|
raise CircuitOpenError(f"熔断器已开启,拒绝调用 provider={self._provider}")
|
||||||
|
|
||||||
# ② 缓存查询(cache 为 None 时跳过)— call_id 在缓存路径独立生成
|
# ② 缓存查询(cache 为 None 时跳过)— call_id 在缓存路径独立生成
|
||||||
cached = await self._cache.get(self._model, messages) if self._cache is not None else None
|
cached = (
|
||||||
|
await self._cache.get(self._model, messages, cache_salt)
|
||||||
|
if self._cache is not None
|
||||||
|
else None
|
||||||
|
)
|
||||||
if cached is not None:
|
if cached is not None:
|
||||||
cache_call_id = str(uuid4())
|
cache_call_id = str(uuid4())
|
||||||
response = LLMResponse(
|
response = LLMResponse(
|
||||||
@@ -370,7 +376,7 @@ class GovernedLLMClient:
|
|||||||
|
|
||||||
# ④ 写缓存(cache 为 None 时跳过)
|
# ④ 写缓存(cache 为 None 时跳过)
|
||||||
if self._cache is not None:
|
if self._cache is not None:
|
||||||
await self._cache.set(self._model, messages, response)
|
await self._cache.set(self._model, messages, response, cache_salt)
|
||||||
|
|
||||||
# ⑤ 遥测
|
# ⑤ 遥测
|
||||||
await self._telemetry.record_llm_call(
|
await self._telemetry.record_llm_call(
|
||||||
|
|||||||
+23
-9
@@ -29,36 +29,48 @@ class RedisResponseCache:
|
|||||||
self._redis = redis
|
self._redis = redis
|
||||||
self._ttl_s = ttl_s
|
self._ttl_s = ttl_s
|
||||||
|
|
||||||
def _build_key(self, model: str, messages: list[dict[str, str]]) -> str:
|
def _build_key(
|
||||||
|
self,
|
||||||
|
model: str,
|
||||||
|
messages: list[dict[str, str]],
|
||||||
|
cache_salt: str | None = None,
|
||||||
|
) -> str:
|
||||||
"""构造 content-addressed 缓存键。
|
"""构造 content-addressed 缓存键。
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
model: 模型名称。
|
model: 模型名称。
|
||||||
messages: 消息列表。
|
messages: 消息列表。
|
||||||
|
cache_salt: 可选缓存盐(如跨 epoch 强制重采样)。仅当非 None 时才加入
|
||||||
|
键 payload,保证默认 None 时键结构与旧缓存一字节不差、旧键不失效。
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
sha256 哈希字符串作为 Redis 键。
|
sha256 哈希字符串作为 Redis 键。
|
||||||
"""
|
"""
|
||||||
payload = json.dumps(
|
key_obj: dict[str, Any] = {"model": model, "messages": messages}
|
||||||
{"model": model, "messages": messages},
|
if cache_salt is not None:
|
||||||
sort_keys=True,
|
key_obj["salt"] = cache_salt
|
||||||
ensure_ascii=False,
|
payload = json.dumps(key_obj, sort_keys=True, ensure_ascii=False)
|
||||||
)
|
|
||||||
digest = hashlib.sha256(payload.encode("utf-8")).hexdigest()
|
digest = hashlib.sha256(payload.encode("utf-8")).hexdigest()
|
||||||
return f"llm_cache:{digest}"
|
return f"llm_cache:{digest}"
|
||||||
|
|
||||||
async def get(self, model: str, messages: list[dict[str, str]]) -> LLMResponse | None:
|
async def get(
|
||||||
|
self,
|
||||||
|
model: str,
|
||||||
|
messages: list[dict[str, str]],
|
||||||
|
cache_salt: str | None = None,
|
||||||
|
) -> LLMResponse | None:
|
||||||
"""从缓存读取 LLM 响应。
|
"""从缓存读取 LLM 响应。
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
model: 模型名称。
|
model: 模型名称。
|
||||||
messages: 消息列表。
|
messages: 消息列表。
|
||||||
|
cache_salt: 可选缓存盐,透传到键构造。
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
缓存命中时返回 LLMResponse,未命中或 Redis 异常时返回 None。
|
缓存命中时返回 LLMResponse,未命中或 Redis 异常时返回 None。
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
key = self._build_key(model, messages)
|
key = self._build_key(model, messages, cache_salt)
|
||||||
raw = await self._redis.get(key)
|
raw = await self._redis.get(key)
|
||||||
except Exception:
|
except Exception:
|
||||||
logger.warning("Redis 缓存读取失败,降级为未命中")
|
logger.warning("Redis 缓存读取失败,降级为未命中")
|
||||||
@@ -75,6 +87,7 @@ class RedisResponseCache:
|
|||||||
model: str,
|
model: str,
|
||||||
messages: list[dict[str, str]],
|
messages: list[dict[str, str]],
|
||||||
response: LLMResponse,
|
response: LLMResponse,
|
||||||
|
cache_salt: str | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""将 LLM 响应写入缓存。
|
"""将 LLM 响应写入缓存。
|
||||||
|
|
||||||
@@ -82,9 +95,10 @@ class RedisResponseCache:
|
|||||||
model: 模型名称。
|
model: 模型名称。
|
||||||
messages: 消息列表。
|
messages: 消息列表。
|
||||||
response: 待缓存的 LLMResponse。
|
response: 待缓存的 LLMResponse。
|
||||||
|
cache_salt: 可选缓存盐,透传到键构造。
|
||||||
"""
|
"""
|
||||||
try:
|
try:
|
||||||
key = self._build_key(model, messages)
|
key = self._build_key(model, messages, cache_salt)
|
||||||
value = json.dumps(dataclasses.asdict(response), ensure_ascii=False)
|
value = json.dumps(dataclasses.asdict(response), ensure_ascii=False)
|
||||||
if self._ttl_s:
|
if self._ttl_s:
|
||||||
await self._redis.set(key, value, ex=self._ttl_s)
|
await self._redis.set(key, value, ex=self._ttl_s)
|
||||||
|
|||||||
@@ -36,6 +36,7 @@ class GovernedVLMClient:
|
|||||||
*,
|
*,
|
||||||
session_id: str | None = None,
|
session_id: str | None = None,
|
||||||
parent_call_id: str | None = None,
|
parent_call_id: str | None = None,
|
||||||
|
cache_salt: str | None = None,
|
||||||
) -> LLMResponse:
|
) -> LLMResponse:
|
||||||
"""图文调用:将图片编码为 base64 嵌入 messages,委托给 LLM 客户端。
|
"""图文调用:将图片编码为 base64 嵌入 messages,委托给 LLM 客户端。
|
||||||
|
|
||||||
@@ -45,6 +46,7 @@ class GovernedVLMClient:
|
|||||||
images: 图片文件路径列表。
|
images: 图片文件路径列表。
|
||||||
session_id: 会话 ID(遥测用)。
|
session_id: 会话 ID(遥测用)。
|
||||||
parent_call_id: 父调用 ID(遥测用)。
|
parent_call_id: 父调用 ID(遥测用)。
|
||||||
|
cache_salt: 可选缓存盐,透传到底层 LLM 缓存键。
|
||||||
|
|
||||||
返回:
|
返回:
|
||||||
LLMResponse。
|
LLMResponse。
|
||||||
@@ -54,6 +56,7 @@ class GovernedVLMClient:
|
|||||||
vision_messages,
|
vision_messages,
|
||||||
session_id=session_id,
|
session_id=session_id,
|
||||||
parent_call_id=parent_call_id,
|
parent_call_id=parent_call_id,
|
||||||
|
cache_salt=cache_salt,
|
||||||
)
|
)
|
||||||
|
|
||||||
@staticmethod
|
@staticmethod
|
||||||
|
|||||||
@@ -25,6 +25,7 @@ class LLMProvider(Protocol):
|
|||||||
*,
|
*,
|
||||||
session_id: str | None = None,
|
session_id: str | None = None,
|
||||||
parent_call_id: str | None = None,
|
parent_call_id: str | None = None,
|
||||||
|
cache_salt: str | None = None,
|
||||||
) -> LLMResponse: ...
|
) -> LLMResponse: ...
|
||||||
|
|
||||||
|
|
||||||
@@ -39,6 +40,7 @@ class VLMProvider(Protocol):
|
|||||||
*,
|
*,
|
||||||
session_id: str | None = None,
|
session_id: str | None = None,
|
||||||
parent_call_id: str | None = None,
|
parent_call_id: str | None = None,
|
||||||
|
cache_salt: str | None = None,
|
||||||
) -> LLMResponse: ...
|
) -> LLMResponse: ...
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -23,8 +23,13 @@ class _FakeRedisCache:
|
|||||||
def __init__(self) -> None:
|
def __init__(self) -> None:
|
||||||
self._store: dict[str, LLMResponse] = {}
|
self._store: dict[str, LLMResponse] = {}
|
||||||
|
|
||||||
async def get(self, model: str, messages: list[dict[str, str]]) -> LLMResponse | None:
|
async def get(
|
||||||
key = f"{model}:{json.dumps(messages, sort_keys=True)}"
|
self,
|
||||||
|
model: str,
|
||||||
|
messages: list[dict[str, str]],
|
||||||
|
cache_salt: str | None = None,
|
||||||
|
) -> LLMResponse | None:
|
||||||
|
key = f"{model}:{cache_salt}:{json.dumps(messages, sort_keys=True)}"
|
||||||
return self._store.get(key)
|
return self._store.get(key)
|
||||||
|
|
||||||
async def set(
|
async def set(
|
||||||
@@ -32,8 +37,9 @@ class _FakeRedisCache:
|
|||||||
model: str,
|
model: str,
|
||||||
messages: list[dict[str, str]],
|
messages: list[dict[str, str]],
|
||||||
response: LLMResponse,
|
response: LLMResponse,
|
||||||
|
cache_salt: str | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
key = f"{model}:{json.dumps(messages, sort_keys=True)}"
|
key = f"{model}:{cache_salt}:{json.dumps(messages, sort_keys=True)}"
|
||||||
self._store[key] = response
|
self._store[key] = response
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
@@ -109,6 +109,18 @@ async def test_different_models_different_keys(
|
|||||||
assert await cache.get("claude-3-opus", MESSAGES) is None
|
assert await cache.get("claude-3-opus", MESSAGES) is None
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_salt_changes_key(fake_redis: object) -> None:
|
||||||
|
"""cache_salt 进入键:不同 salt 产生不同键,None 保持旧键结构。"""
|
||||||
|
cache = RedisResponseCache(redis=fake_redis, ttl_s=None)
|
||||||
|
k_none = cache._build_key("m", [{"role": "user", "content": "x"}], None)
|
||||||
|
k_e1 = cache._build_key("m", [{"role": "user", "content": "x"}], "run:e1")
|
||||||
|
k_e2 = cache._build_key("m", [{"role": "user", "content": "x"}], "run:e2")
|
||||||
|
assert k_none != k_e1
|
||||||
|
assert k_none != k_e2
|
||||||
|
assert k_e1 != k_e2
|
||||||
|
|
||||||
|
|
||||||
@pytest.mark.asyncio
|
@pytest.mark.asyncio
|
||||||
async def test_graceful_degradation_on_error() -> None:
|
async def test_graceful_degradation_on_error() -> None:
|
||||||
"""Redis 不可用时静默降级:get→None,set→不报错。"""
|
"""Redis 不可用时静默降级:get→None,set→不报错。"""
|
||||||
|
|||||||
Reference in New Issue
Block a user