Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions .changeset/warm-pages-cache.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,5 @@
---
"apollo": patch
---

global_chat: cache the planner's history across tool-loop rounds
6 changes: 6 additions & 0 deletions services/global_chat/planner.py
Original file line number Diff line number Diff line change
Expand Up @@ -53,6 +53,10 @@
},
}

# API puts a breakpoint on the last block of each request,
# so each tool-loop round reads the rounds before it from cache.
_HISTORY_CACHE = {"cache_control": {"type": "ephemeral"}}

_FINAL_ROUND_NOTICE = (
"Stop and reply to the user now. Say what you changed. Mention unfinished work "
"only if there is any, and then offer to continue next turn — otherwise don't "
Expand Down Expand Up @@ -614,6 +618,7 @@ def _call_api(self, system_prompt, messages, stream, stream_manager, tool_choice
output_config={"effort": "medium"},
**choice,
**_CONTEXT_MANAGEMENT,
**_HISTORY_CACHE,
) as stream_obj:
for event in stream_obj:
if event.type == "content_block_delta" and event.delta.type == "text_delta":
Expand Down Expand Up @@ -644,6 +649,7 @@ def _call_api(self, system_prompt, messages, stream, stream_manager, tool_choice
# which the SDK otherwise rejects.
timeout=httpx.Timeout(600.0, connect=5.0),
**_CONTEXT_MANAGEMENT,
**_HISTORY_CACHE,
)
return response

Expand Down
98 changes: 98 additions & 0 deletions services/global_chat/tests/unit/test_planner.py
Original file line number Diff line number Diff line change
Expand Up @@ -1592,3 +1592,101 @@ def always_fail(*_args: object, **_kwargs: object) -> FakeResponse:
assert "something else" not in excinfo.value.details["upstream_message"]
# The retry failed, so nothing was claimed about web search.
assert planner._segments == []


class RecordingMessages:
"""Records the kwargs of every stream/create call and replays scripted outcomes.

Each outcome is a FakeResponse to return, or an exception to raise.
"""

def __init__(self, outcomes: list) -> None:
self._outcomes = list(outcomes)
self.calls: list[dict] = []

def _next(self, kwargs: dict) -> FakeResponse:
self.calls.append(kwargs)
outcome = self._outcomes.pop(0)
if isinstance(outcome, Exception):
raise outcome
return outcome

def create(self, **kwargs: object) -> FakeResponse:
return self._next(kwargs)

def stream(self, **kwargs: object) -> FakeStream:
return FakeStream([], self._next(kwargs))


class RecordingClient:
def __init__(self, outcomes: list) -> None:
self.messages = RecordingMessages(outcomes)
self.beta = self


HISTORY_CACHE = {"type": "ephemeral"}


def test_a_streamed_call_sends_history_cache() -> None:
planner = make_run_planner()
planner.client = RecordingClient([FakeResponse("end_turn", [FakeTextBlock("Hi")])])

planner._call_api([], [], True, StubStreamManager())

assert planner.client.messages.calls[0]["cache_control"] == HISTORY_CACHE


def test_a_non_streamed_call_sends_history_cache() -> None:
planner = make_run_planner()
planner.client = RecordingClient([FakeResponse("end_turn", [FakeTextBlock("Hi")])])

planner._call_api([], [], False, StubStreamManager())

assert planner.client.messages.calls[0]["cache_control"] == HISTORY_CACHE


def test_every_round_of_a_turn_carries_history_cache() -> None:
"""Pause continuation, tool round and the no-tools wrap-up all re-send history."""
planner = make_run_planner(max_tool_calls=1)
planner.tools = TOOL_DEFINITIONS
planner.config_loader = StubPromptLoader(planner_system_prompt="BASE PROMPT")
inspect = FakeToolUse("inspect_job_code", {"job_keys": ["fetch-patients"]})
planner.client = RecordingClient([
FakeResponse("pause_turn", [FakeTextBlock("Looking. ")]),
FakeResponse("tool_use", [inspect]),
FakeResponse("end_turn", [FakeTextBlock("Done.")]),
])

result = planner.run("q", WORKFLOW_YAML, None, [], stream=False)

calls = planner.client.messages.calls
assert result.response == "Done."
assert len(calls) == 3
# The last call is the budget-spent wrap-up round.
assert calls[-1]["tool_choice"] == {"type": "none"}
for call in calls:
assert call["cache_control"] == HISTORY_CACHE
# The existing explicit breakpoints are unchanged: one on the tools,
# one on the system prompt, plus the automatic one is 3 of the API's 4.
explicit = [t for t in call["tools"] if "cache_control" in t] + \
[b for b in call["system"] if "cache_control" in b]
assert len(explicit) == 2


def test_web_downgrade_retry_keeps_history_cache() -> None:
planner = make_run_planner()
planner.web_tools = build_web_tools(WEB_CONFIG)
planner.tools = TOOL_DEFINITIONS + planner.web_tools
planner.web_search_enabled = True
planner.config_loader = StubPromptLoader(planner_system_prompt="BASE PROMPT")
planner.client = RecordingClient([
make_bad_request(),
FakeResponse("end_turn", [FakeTextBlock("Answered without the web.")]),
])

result = planner.run("q", None, None, [], stream=False)

calls = planner.client.messages.calls
assert result.meta["web_search_downgraded"] is True
assert len(calls) == 2
assert all(call["cache_control"] == HISTORY_CACHE for call in calls)
Loading