diff --git a/adapters/llm.py b/adapters/llm.py index a8f646d..8443301 100644 --- a/adapters/llm.py +++ b/adapters/llm.py @@ -274,6 +274,7 @@ class GovernedLLMClient: *, session_id: str | None = None, parent_call_id: str | None = None, + cache_salt: str | None = None, ) -> LLMResponse: """发起 LLM 调用,经四层治理栈:熔断 → 缓存 → 重试+流式 → 遥测。 @@ -281,6 +282,7 @@ class GovernedLLMClient: messages: OpenAI 格式消息列表。 session_id: 会话 ID(传递到遥测)。 parent_call_id: 父调用 ID(传递到遥测)。 + cache_salt: 可选缓存盐,透传到 Redis 缓存键(如跨 epoch 重采样)。 返回: LLMResponse 统一响应。 @@ -296,7 +298,11 @@ class GovernedLLMClient: raise CircuitOpenError(f"熔断器已开启,拒绝调用 provider={self._provider}") # ② 缓存查询(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: cache_call_id = str(uuid4()) response = LLMResponse( @@ -370,7 +376,7 @@ class GovernedLLMClient: # ④ 写缓存(cache 为 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( diff --git a/adapters/redis_cache.py b/adapters/redis_cache.py index 4af573c..39420e6 100644 --- a/adapters/redis_cache.py +++ b/adapters/redis_cache.py @@ -29,36 +29,48 @@ class RedisResponseCache: self._redis = redis 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 缓存键。 Args: model: 模型名称。 messages: 消息列表。 + cache_salt: 可选缓存盐(如跨 epoch 强制重采样)。仅当非 None 时才加入 + 键 payload,保证默认 None 时键结构与旧缓存一字节不差、旧键不失效。 Returns: sha256 哈希字符串作为 Redis 键。 """ - payload = json.dumps( - {"model": model, "messages": messages}, - sort_keys=True, - ensure_ascii=False, - ) + key_obj: dict[str, Any] = {"model": model, "messages": messages} + if cache_salt is not None: + key_obj["salt"] = cache_salt + payload = json.dumps(key_obj, sort_keys=True, ensure_ascii=False) digest = hashlib.sha256(payload.encode("utf-8")).hexdigest() 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 响应。 Args: model: 模型名称。 messages: 消息列表。 + cache_salt: 可选缓存盐,透传到键构造。 Returns: 缓存命中时返回 LLMResponse,未命中或 Redis 异常时返回 None。 """ try: - key = self._build_key(model, messages) + key = self._build_key(model, messages, cache_salt) raw = await self._redis.get(key) except Exception: logger.warning("Redis 缓存读取失败,降级为未命中") @@ -75,6 +87,7 @@ class RedisResponseCache: model: str, messages: list[dict[str, str]], response: LLMResponse, + cache_salt: str | None = None, ) -> None: """将 LLM 响应写入缓存。 @@ -82,9 +95,10 @@ class RedisResponseCache: model: 模型名称。 messages: 消息列表。 response: 待缓存的 LLMResponse。 + cache_salt: 可选缓存盐,透传到键构造。 """ try: - key = self._build_key(model, messages) + key = self._build_key(model, messages, cache_salt) value = json.dumps(dataclasses.asdict(response), ensure_ascii=False) if self._ttl_s: await self._redis.set(key, value, ex=self._ttl_s) diff --git a/adapters/vlm.py b/adapters/vlm.py index cf3342a..b959cb5 100644 --- a/adapters/vlm.py +++ b/adapters/vlm.py @@ -36,6 +36,7 @@ class GovernedVLMClient: *, session_id: str | None = None, parent_call_id: str | None = None, + cache_salt: str | None = None, ) -> LLMResponse: """图文调用:将图片编码为 base64 嵌入 messages,委托给 LLM 客户端。 @@ -45,6 +46,7 @@ class GovernedVLMClient: images: 图片文件路径列表。 session_id: 会话 ID(遥测用)。 parent_call_id: 父调用 ID(遥测用)。 + cache_salt: 可选缓存盐,透传到底层 LLM 缓存键。 返回: LLMResponse。 @@ -54,6 +56,7 @@ class GovernedVLMClient: vision_messages, session_id=session_id, parent_call_id=parent_call_id, + cache_salt=cache_salt, ) @staticmethod diff --git a/core/protocols.py b/core/protocols.py index 3a08d27..2d4e89f 100644 --- a/core/protocols.py +++ b/core/protocols.py @@ -25,6 +25,7 @@ class LLMProvider(Protocol): *, session_id: str | None = None, parent_call_id: str | None = None, + cache_salt: str | None = None, ) -> LLMResponse: ... @@ -39,6 +40,7 @@ class VLMProvider(Protocol): *, session_id: str | None = None, parent_call_id: str | None = None, + cache_salt: str | None = None, ) -> LLMResponse: ... diff --git a/tests/unit/test_governed_llm.py b/tests/unit/test_governed_llm.py index 923b6ec..9f487e4 100644 --- a/tests/unit/test_governed_llm.py +++ b/tests/unit/test_governed_llm.py @@ -23,8 +23,13 @@ class _FakeRedisCache: def __init__(self) -> None: self._store: dict[str, LLMResponse] = {} - async def get(self, model: str, messages: list[dict[str, str]]) -> LLMResponse | None: - key = f"{model}:{json.dumps(messages, sort_keys=True)}" + async def get( + 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) async def set( @@ -32,8 +37,9 @@ class _FakeRedisCache: model: str, messages: list[dict[str, str]], response: LLMResponse, + cache_salt: str | 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 diff --git a/tests/unit/test_redis_cache.py b/tests/unit/test_redis_cache.py index 8523384..8b5945e 100644 --- a/tests/unit/test_redis_cache.py +++ b/tests/unit/test_redis_cache.py @@ -109,6 +109,18 @@ async def test_different_models_different_keys( 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 async def test_graceful_degradation_on_error() -> None: """Redis 不可用时静默降级:get→None,set→不报错。"""