Skip to content
Merged
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
6 changes: 6 additions & 0 deletions agent_core/messages.py
Original file line number Diff line number Diff line change
Expand Up @@ -84,6 +84,12 @@ class Message(TypedDict, total=False):
# endpoint's pydantic union, whereas a message-level key outside
# ``WIRE_MESSAGE_KEYS`` is dropped by ``for_wire``.
image_meta: list[dict[str, Any]]
# Marks a per-request projection that is NOT in persistent history (the
# loop's ``system_addendum_per_call``). A provider that places a rolling
# prompt-cache breakpoint must anchor it BEFORE such messages: the next
# request drops them and appends the real turn in their place, so a prefix
# ending on one is never reused. Filtered out by ``for_wire``.
transient: bool


# ── Wire boundary ────────────────────────────────────────────────────────
Expand Down
37 changes: 25 additions & 12 deletions agent_core/providers/anthropic.py
Original file line number Diff line number Diff line change
Expand Up @@ -95,15 +95,21 @@ def _build_kwargs(
) -> dict[str, Any]:
"""Shared request-shape builder for :meth:`chat` and :meth:`stream`."""
system, msgs = _split_system(messages)
# ``_to_anthropic_msg`` returns None for a message with nothing
# sendable (a contentless assistant turn); those are dropped.
pairs = [
(converted, bool(m.get("transient")))
for m in msgs
if (converted := _to_anthropic_msg(m)) is not None
]
transient_tail = 0
for _, is_transient in reversed(pairs):
if not is_transient:
break
transient_tail += 1
kwargs: dict[str, Any] = {
"model": self.model,
# ``_to_anthropic_msg`` returns None for a message with nothing
# sendable (a contentless assistant turn); those are dropped.
"messages": [
converted
for converted in (_to_anthropic_msg(m) for m in msgs)
if converted is not None
],
"messages": [converted for converted, _ in pairs],
"max_tokens": max_tokens or self.default_max_tokens or 4096,
}
if system:
Expand All @@ -128,7 +134,7 @@ def _build_kwargs(
kwargs["timeout"] = timeout
elif self.default_timeout is not None:
kwargs["timeout"] = self.default_timeout
_add_prompt_cache(kwargs)
_add_prompt_cache(kwargs, transient_tail=transient_tail)
return kwargs

async def chat(
Expand Down Expand Up @@ -380,7 +386,7 @@ def _split_system(messages: list[Message]) -> tuple[str, list[Message]]:
return "", list(messages)


def _add_prompt_cache(kwargs: dict[str, Any]) -> None:
def _add_prompt_cache(kwargs: dict[str, Any], *, transient_tail: int = 0) -> None:
"""Set Anthropic prompt-cache breakpoints on ``kwargs`` in place.

Anthropic caching is opt-in per content block (unlike OpenAI's automatic
Expand All @@ -391,6 +397,13 @@ def _add_prompt_cache(kwargs: dict[str, Any]) -> None:
to 4 and serves the longest matching cached prefix, so these two cover the
static head and the moving tail. This lives inside ``AnthropicClient`` so it
only ever touches Anthropic requests. Disable with ``ANTHROPIC_PROMPT_CACHE=0``.

``transient_tail`` counts trailing messages that exist only in this request
(``Message.transient``, e.g. the per-call runtime addendum). The rolling
breakpoint skips them: the next request replaces them with the real turn,
so a cached prefix ending on one never matches again and every turn would
re-write the whole conversation at the cache-write rate while reading only
the static head.
"""
if os.getenv("ANTHROPIC_PROMPT_CACHE", "1") == "0":
return
Expand All @@ -402,11 +415,11 @@ def _add_prompt_cache(kwargs: dict[str, Any]) -> None:
"text": system,
"cache_control": {"type": "ephemeral"},
}]
# Rolling tail: mark the last message's final content block.
# Rolling tail: mark the last persistent message's final content block.
msgs = kwargs.get("messages")
if not msgs:
if not msgs or transient_tail >= len(msgs):
return
last = msgs[-1]
last = msgs[-1 - transient_tail]
content = last.get("content")
if isinstance(content, str):
if content:
Expand Down
6 changes: 5 additions & 1 deletion agent_core/runtime/loop/agent_loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -715,7 +715,11 @@ async def _prepare_llm_request(
else system_msg
)
if addendum_text:
messages_for_call = [*messages, addendum_factory(addendum_text)]
addendum = addendum_factory(addendum_text)
# Not in ``messages``: tell cache-aware providers not to anchor the
# rolling breakpoint on it (see ``Message.transient``).
addendum["transient"] = True
messages_for_call = [*messages, addendum]

# Publish the estimate of THIS request, after observer injections and the
# addendum. An observer comparing its own estimate against the provider's
Expand Down
1 change: 1 addition & 0 deletions changes/48.fix.md
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
Fix Anthropic prompt caching for per-call runtime addenda by anchoring the rolling cache breakpoint on the last persistent conversation message.
20 changes: 20 additions & 0 deletions tests/test_agent_loop_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -1009,3 +1009,23 @@ def to_openai_schema(self) -> dict[str, Any]:
for message in llm.calls[1]
)
assert _orphan_tool_call_ids(llm.calls[1]) == set()


@pytest.mark.asyncio
async def test_per_call_addendum_is_flagged_transient_and_kept_out_of_history() -> None:
llm = SequenceLLM([LLMResponse(content="finished")])
config = LoopConfig(
max_turns=1,
loop_policy=LoopPolicy(no_tool_behavior="stop"),
max_llm_retries=1,
system_addendum_per_call="[env]",
system_addendum_per_call_role="user",
)
result = await run_agent_loop(
system_prompt="system", user_message="start", llm=llm, tools=[],
config=config, model_profile=ModelProfile(model_id="test", provider="test"),
)
sent = llm.calls[0]
assert sent[-1]["content"] == "[env]" and sent[-1]["transient"] is True
assert not any(m.get("transient") for m in sent[:-1])
assert not any(m.get("content") == "[env]" for m in result.messages)
70 changes: 70 additions & 0 deletions tests/test_provider_native_clients.py
Original file line number Diff line number Diff line change
Expand Up @@ -1120,3 +1120,73 @@ def test_session_resolver_hooks_are_reachable_from_the_package_facade():
):
assert hasattr(providers, name), name
assert name in providers.__all__, name


def _tail_breakpoints(kwargs):
"""Indices of messages carrying a ``cache_control`` block."""
return [
i for i, m in enumerate(kwargs["messages"])
if isinstance(m["content"], list)
and any(isinstance(b, dict) and "cache_control" in b for b in m["content"])
]


def _transient(text):
m = user_msg(text)
m["transient"] = True
return m


def _strip_cache(msgs):
return [
{**m, "content": [
{k: v for k, v in b.items() if k != "cache_control"} for b in m["content"]
] if isinstance(m["content"], list) else m["content"]}
for m in msgs
]


def test_anthropic_rolling_breakpoint_skips_transient_addendum(monkeypatch):
"""The per-call addendum is dropped next turn, so the rolling breakpoint
must sit on the last persistent message; otherwise no conversation prefix
is ever reused and only the static system head hits the cache."""
monkeypatch.delenv("ANTHROPIC_PROMPT_CACHE", raising=False)
c = ac.AnthropicClient("claude-x", api_key="x")
call = {"id": "t1", "type": "function",
"function": {"name": "bash", "arguments": "{}"}}
history = [system_msg("s"), user_msg("q"),
assistant_msg("", tool_calls=[call]), tool_msg("t1", "out")]
build = lambda msgs: c._build_kwargs( # noqa: E731
msgs, tools=None, temperature=None, max_tokens=None,
extra_headers=None, timeout=None,
)

first = build([*history, _transient("[Runtime environment metadata]")])
# Anchor is the tool result (index 2 after the system split), not the addendum.
assert _tail_breakpoints(first) == [2]
assert "cache_control" not in first["messages"][-1]["content"][-1]

call2 = {**call, "id": "t2"}
history += [assistant_msg("", tool_calls=[call2]), tool_msg("t2", "out2")]
second = build([*history, _transient("[Runtime environment metadata]")])
# Everything up to last turn's breakpoint is byte-identical this turn.
assert _strip_cache(first["messages"][:3]) == _strip_cache(second["messages"][:3])
assert _tail_breakpoints(second) == [4]


def test_anthropic_rolling_breakpoint_without_addendum_marks_last(monkeypatch):
"""Turns before ``system_addendum_min_turn`` carry no addendum: no offset."""
monkeypatch.delenv("ANTHROPIC_PROMPT_CACHE", raising=False)
c = ac.AnthropicClient("claude-x", api_key="x")
kwargs = c._build_kwargs(
[system_msg("s"), user_msg("q")],
tools=None, temperature=None, max_tokens=None,
extra_headers=None, timeout=None,
)
assert _tail_breakpoints(kwargs) == [0]


def test_transient_flag_is_not_sent_on_openai_wire():
from agent_core.messages import for_wire

assert for_wire([_transient("x")]) == [user_msg("x")]
Loading