fix(provider): stabilize Codex prompt cache routing (#5540)

This commit is contained in:
chengyongru
2026-08-26 01:36:00 +08:00
committed by GitHub
parent 3ee3791626
commit 42f37dc4c0
10 changed files with 91 additions and 26 deletions
+1
View File
@@ -495,6 +495,7 @@ class AgentRunner:
model=spec.runtime.model, model=spec.runtime.model,
messages=messages, messages=messages,
state=spec.provider_state, state=spec.provider_state,
session_id=spec.session_key,
) )
governance_config = ContextGovernanceConfig( governance_config = ContextGovernanceConfig(
provider=spec.runtime.provider, provider=spec.runtime.provider,
+4
View File
@@ -252,10 +252,13 @@ class ProviderCallContext:
The regular ``chat`` contract stays provider-agnostic. Responses-capable The regular ``chat`` contract stays provider-agnostic. Responses-capable
providers consume this context through the opt-in ``chat_with_context`` providers consume this context through the opt-in ``chat_with_context``
hooks, while every other provider inherits the context-free delegation. hooks, while every other provider inherits the context-free delegation.
``session_id`` gives providers a stable conversation-scoped routing key
without exposing that identity in the public message transcript.
""" """
conversation_state: ProviderConversationState | None = field(default=None, repr=False) conversation_state: ProviderConversationState | None = field(default=None, repr=False)
context_window_tokens: int | None = None context_window_tokens: int | None = None
session_id: str | None = field(default=None, repr=False)
@dataclass(frozen=True, slots=True) @dataclass(frozen=True, slots=True)
@@ -1640,6 +1643,7 @@ class LLMProvider(ABC):
context_window_tokens=( context_window_tokens=(
provider_context.context_window_tokens provider_context.context_window_tokens
), ),
session_id=provider_context.session_id,
) )
if stripped is not None or stripped_context is not None: if stripped is not None or stripped_context is not None:
logger.warning( logger.warning(
+8 -2
View File
@@ -42,9 +42,11 @@ class ProviderConversationStateController:
model: str | None, model: str | None,
messages: list[dict[str, Any]], messages: list[dict[str, Any]],
state: ProviderConversationState | None = None, state: ProviderConversationState | None = None,
session_id: str | None = None,
) -> None: ) -> None:
self._provider = provider self._provider = provider
self._model = model self._model = model
self._session_id = session_id
self._state = ( self._state = (
state state
if state is not None if state is not None
@@ -60,9 +62,12 @@ class ProviderConversationStateController:
context_window_tokens: int | None, context_window_tokens: int | None,
) -> ProviderCallContext | None: ) -> ProviderCallContext | None:
"""Return typed provider context for a request that does not resume state.""" """Return typed provider context for a request that does not resume state."""
if context_window_tokens is None: if context_window_tokens is None and self._session_id is None:
return None return None
return ProviderCallContext(context_window_tokens=context_window_tokens) return ProviderCallContext(
context_window_tokens=context_window_tokens,
session_id=self._session_id,
)
def prepare_request( def prepare_request(
self, self,
@@ -112,6 +117,7 @@ class ProviderConversationStateController:
if independent_context is not None if independent_context is not None
else None else None
), ),
session_id=self._session_id,
) )
def observe_response( def observe_response(
+2
View File
@@ -186,6 +186,7 @@ class FallbackProvider(LLMProvider):
return ProviderCallContext( return ProviderCallContext(
conversation_state=provider_context.conversation_state, conversation_state=provider_context.conversation_state,
context_window_tokens=context_window_tokens, context_window_tokens=context_window_tokens,
session_id=provider_context.session_id,
) )
def _primary_available(self) -> bool: def _primary_available(self) -> bool:
@@ -541,6 +542,7 @@ class FallbackProvider(LLMProvider):
fallback_kwargs["provider_context"] = ProviderCallContext( fallback_kwargs["provider_context"] = ProviderCallContext(
conversation_state=state, conversation_state=state,
context_window_tokens=context_window_tokens, context_window_tokens=context_window_tokens,
session_id=provider_context.session_id,
) )
if fallback.reasoning_effort is None: if fallback.reasoning_effort is None:
fallback_kwargs.pop("reasoning_effort", None) fallback_kwargs.pop("reasoning_effort", None)
+5 -4
View File
@@ -103,6 +103,7 @@ class OpenAICodexProvider(LLMProvider):
provider=self._responses_state_provider(), provider=self._responses_state_provider(),
model=_strip_model_prefix(model), model=_strip_model_prefix(model),
) )
session_id = provider_context.session_id if provider_context is not None else None
body: dict[str, Any] = { body: dict[str, Any] = {
"model": _strip_model_prefix(model), "model": _strip_model_prefix(model),
@@ -111,10 +112,11 @@ class OpenAICodexProvider(LLMProvider):
"instructions": system_prompt, "instructions": system_prompt,
"input": input_items, "input": input_items,
"text": {"verbosity": "medium"}, "text": {"verbosity": "medium"},
"prompt_cache_key": _prompt_cache_key(messages[:2]),
"tool_choice": tool_choice or "auto", "tool_choice": tool_choice or "auto",
"parallel_tool_calls": True, "parallel_tool_calls": True,
} }
if session_id:
body["prompt_cache_key"] = _prompt_cache_key(session_id)
body["include"] = ["reasoning.encrypted_content"] body["include"] = ["reasoning.encrypted_content"]
reasoning_options = _build_reasoning_options(reasoning_effort) reasoning_options = _build_reasoning_options(reasoning_effort)
if replayed and "gpt-5.6" in _strip_model_prefix(model).lower(): if replayed and "gpt-5.6" in _strip_model_prefix(model).lower():
@@ -496,9 +498,8 @@ async def _request_codex(
return result return result
def _prompt_cache_key(messages: list[dict[str, Any]]) -> str: def _prompt_cache_key(session_id: str) -> str:
raw = json.dumps(messages, ensure_ascii=True, sort_keys=True) return hashlib.sha256(session_id.encode("utf-8")).hexdigest()
return hashlib.sha256(raw.encode("utf-8")).hexdigest()
def _friendly_error(status_code: int, raw: str) -> str: def _friendly_error(status_code: int, raw: str) -> str:
+5 -1
View File
@@ -384,12 +384,16 @@ class TestFallbackOnPrimaryError:
messages=[{"role": "user", "content": "hi"}], messages=[{"role": "user", "content": "hi"}],
model="gpt-5.6", model="gpt-5.6",
max_tokens=10_000, max_tokens=10_000,
provider_context=ProviderCallContext(context_window_tokens=50_000), provider_context=ProviderCallContext(
context_window_tokens=50_000,
session_id="webui:cache-test",
),
) )
primary_context = primary.context_calls[0] primary_context = primary.context_calls[0]
assert primary_context is not None assert primary_context is not None
assert primary_context.context_window_tokens == 200_000 assert primary_context.context_window_tokens == 200_000
assert primary_context.session_id == "webui:cache-test"
assert resolve_compact_threshold( assert resolve_compact_threshold(
primary_context.context_window_tokens, primary_context.context_window_tokens,
10_000, 10_000,
@@ -8,6 +8,7 @@ from nanobot.providers.base import (
GenerationSettings, GenerationSettings,
LLMProvider, LLMProvider,
LLMResponse, LLMResponse,
ProviderCallContext,
ToolCallRequest, ToolCallRequest,
) )
from nanobot.utils.llm_runtime import LLMRuntime from nanobot.utils.llm_runtime import LLMRuntime
@@ -22,6 +23,7 @@ async def test_active_run_keeps_provider_captured_at_admission() -> None:
first_calls = 0 first_calls = 0
second_calls = 0 second_calls = 0
request_temperatures: list[float] = [] request_temperatures: list[float] = []
request_session_ids: list[str | None] = []
selected_runtime = LLMRuntime.capture( selected_runtime = LLMRuntime.capture(
first_provider, first_provider,
"captured-model", "captured-model",
@@ -33,6 +35,9 @@ async def test_active_run_keeps_provider_captured_at_admission() -> None:
nonlocal first_calls, selected_runtime nonlocal first_calls, selected_runtime
first_calls += 1 first_calls += 1
request_temperatures.append(kwargs["temperature"]) request_temperatures.append(kwargs["temperature"])
provider_context = kwargs["provider_context"]
assert isinstance(provider_context, ProviderCallContext)
request_session_ids.append(provider_context.session_id)
selected_runtime = LLMRuntime.capture( selected_runtime = LLMRuntime.capture(
second_provider, second_provider,
"future-model", "future-model",
@@ -63,9 +68,11 @@ async def test_active_run_keeps_provider_captured_at_admission() -> None:
runtime=selected_runtime, runtime=selected_runtime,
max_iterations=2, max_iterations=2,
max_tool_result_chars=AgentDefaults().max_tool_result_chars, max_tool_result_chars=AgentDefaults().max_tool_result_chars,
session_key="webui:cache-test",
)) ))
assert first_calls == 2 assert first_calls == 2
assert second_calls == 0 assert second_calls == 0
assert request_temperatures == [0.2, 0.2] assert request_temperatures == [0.2, 0.2]
assert request_session_ids == ["webui:cache-test", "webui:cache-test"]
assert selected_runtime.provider is second_provider assert selected_runtime.provider is second_provider
@@ -9,6 +9,7 @@ import pytest
from nanobot.providers.base import ( from nanobot.providers.base import (
LLMProvider, LLMProvider,
LLMResponse, LLMResponse,
ProviderCallContext,
ProviderConversationState, ProviderConversationState,
ToolCallRequest, ToolCallRequest,
) )
@@ -289,3 +290,23 @@ def test_independent_request_exposes_context_without_capability_check() -> None:
assert provider_context.conversation_state is None assert provider_context.conversation_state is None
assert provider_context.context_window_tokens == 200_000 assert provider_context.context_window_tokens == 200_000
provider.supports_native_compaction.assert_not_called() provider.supports_native_compaction.assert_not_called()
def test_independent_request_exposes_session_id_without_token_budget() -> None:
provider = _provider(compact=False)
messages = [{"role": "user", "content": "hello"}]
controller = ProviderConversationStateController(
provider=provider,
model="gpt-5.6",
messages=messages,
session_id="webui:cache-test",
)
provider_context = controller.prepare_request(
messages,
context_window_tokens=None,
)
assert provider_context == ProviderCallContext(
session_id="webui:cache-test",
)
+28 -13
View File
@@ -260,8 +260,8 @@ async def test_codex_request_uses_configured_proxy(monkeypatch) -> None:
@pytest.mark.asyncio @pytest.mark.asyncio
async def test_codex_prompt_cache_key_uses_stable_conversation_prefix(monkeypatch) -> None: async def test_codex_omits_prompt_cache_key_without_session_id(monkeypatch) -> None:
bodies: list[dict] = [] bodies: list[dict[str, Any]] = []
_mock_codex_token(monkeypatch) _mock_codex_token(monkeypatch)
@@ -289,25 +289,40 @@ async def test_codex_prompt_cache_key_uses_stable_conversation_prefix(monkeypatc
{"role": "assistant", "content": "first answer"}, {"role": "assistant", "content": "first answer"},
], ],
) )
assert "prompt_cache_key" not in bodies[0]
assert "service_tier" not in bodies[0]
@pytest.mark.asyncio
async def test_codex_prompt_cache_key_prefers_stable_session_id(monkeypatch) -> None:
bodies: list[dict[str, Any]] = []
_mock_codex_token(monkeypatch)
async def fake_request(_url, _headers, body, **_kwargs):
bodies.append(body)
return provider_base.LLMResponse(content="ok")
monkeypatch.setattr("nanobot.providers.openai_codex_provider._request_codex", fake_request)
provider = OpenAICodexProvider()
for session_id, first_request in (
("session-a", "first request"),
("session-a", "different visible prefix"),
("session-b", "first request"),
):
await provider.chat( await provider.chat(
[ [
{"role": "system", "content": "You are nanobot."}, {"role": "system", "content": "You are nanobot."},
{"role": "user", "content": "first request"}, {"role": "user", "content": first_request},
{"role": "assistant", "content": "first answer"},
{"role": "user", "content": "follow up"},
],
)
await provider.chat(
[
{"role": "system", "content": "You are nanobot."},
{"role": "user", "content": "different request"},
{"role": "assistant", "content": "first answer"},
], ],
provider_context=provider_base.ProviderCallContext(
session_id=session_id,
),
) )
assert bodies[0]["prompt_cache_key"] == bodies[1]["prompt_cache_key"] assert bodies[0]["prompt_cache_key"] == bodies[1]["prompt_cache_key"]
assert bodies[0]["prompt_cache_key"] != bodies[2]["prompt_cache_key"] assert bodies[0]["prompt_cache_key"] != bodies[2]["prompt_cache_key"]
assert all("service_tier" not in body for body in bodies)
@pytest.mark.asyncio @pytest.mark.asyncio
+5 -1
View File
@@ -431,13 +431,17 @@ async def test_image_retry_discards_provider_state_with_images(
response = await provider.chat_with_retry( response = await provider.chat_with_retry(
messages=messages, messages=messages,
provider_context=ProviderCallContext(conversation_state=state), provider_context=ProviderCallContext(
conversation_state=state,
session_id="webui:cache-test",
),
) )
assert response.content == "ok, no image" assert response.content == "ok, no image"
retry_context = provider.contexts[-1] retry_context = provider.contexts[-1]
assert isinstance(retry_context, ProviderCallContext) assert isinstance(retry_context, ProviderCallContext)
assert retry_context.conversation_state is None assert retry_context.conversation_state is None
assert retry_context.session_id == "webui:cache-test"
public_content = messages[0]["content"] public_content = messages[0]["content"]
if isinstance(public_content, list): if isinstance(public_content, list):
assert all(block.get("type") != "image_url" for block in public_content) assert all(block.get("type") != "image_url" for block in public_content)