refactor(agent): pass-2 structural compaction of prompt_caching and native_compaction

- prompt_caching: fold _count_cache_markers' nested _marked into one sum over
  [*messages, *parts, *tools]; collapse the open-turn guard + endpoint append in
  _completed_transaction_endpoint_indexes into one predicate; enumerate() in
  apply_anthropic_cache_control's non_sys scan; effective_cache_ttl tail -> one
  conditional; pack wrapped calls/comprehensions to <=120 cols; tighten
  docstrings (every WHY kept).
- native_compaction: resolve_compact_threshold upper -> one conditional expr;
  warn-once guard -> single if; _extract_item_text part loop drops the str/dict/
  else ladder (candidates default (part,), non-str filtered by the same isinstance
  check); _retain_summary early-return -> guarded body; pack wrapped calls and
  docstrings.
Zero behavior change: E corpus (224MB) and C corpus byte-identical vs base
113f04616bd; extra corpus for warn-once/_extract_item_text edges identical;
68 test files / 1359 tests green; no symbol added/removed vs pass 1.
This commit is contained in:
Teknium
2026-09-02 21:44:48 -07:00
parent adb241dfc4
commit acfa6e3691
2 changed files with 82 additions and 150 deletions
+40 -77
View File
@@ -38,10 +38,8 @@ def resolve_native_compaction_capabilities(
"""Resolve the native-compaction capability for a runtime destination (a resolved ``False``
is distinct from "unresolved" and must survive model switches unchanged)."""
direct_default = (provider or "").strip().lower() == "openai" and not base_url
eligible = is_native_compaction_model(model) and (
direct_default or is_direct_openai_route(base_url, is_codex_backend=is_codex_backend)
)
return {"native_compaction": eligible}
return {"native_compaction": is_native_compaction_model(model) and (
direct_default or is_direct_openai_route(base_url, is_codex_backend=is_codex_backend))}
def is_direct_openai_route(base_url: Optional[str], *, is_codex_backend: bool = False) -> bool:
@@ -74,10 +72,8 @@ def resolve_compact_threshold(configured_threshold: Any, local_trigger_tokens: A
fires first. Booleans are never thresholds.
"""
local = _positive_int(local_trigger_tokens)
upper = None
if local is not None:
upper = max(1_024, local - LOCAL_TRIGGER_SAFETY_MARGIN if local > LOCAL_TRIGGER_SAFETY_MARGIN else int(local * 0.8))
upper = None if local is None else max(
1_024, local - LOCAL_TRIGGER_SAFETY_MARGIN if local > LOCAL_TRIGGER_SAFETY_MARGIN else int(local * 0.8))
configured = _positive_int(configured_threshold, reject=(bool, float))
if configured is None:
return upper if upper is not None else DEFAULT_COMPACT_THRESHOLD
@@ -92,19 +88,17 @@ _checkpoint_suppression_logged = False
def _warn_native_compaction_suppressed_by_checkpoint_gate() -> None:
"""Log once per process; the suppression itself is re-evaluated per request."""
global _checkpoint_suppression_logged
if _checkpoint_suppression_logged:
return
_checkpoint_suppression_logged = True
logger.warning(
"compression.checkpoint_required is enabled: server-side native "
"compaction (context_management) is disabled for this agent so the "
"checkpoint-aware Hermes compressor stays authoritative."
)
if not _checkpoint_suppression_logged:
_checkpoint_suppression_logged = True
logger.warning(
"compression.checkpoint_required is enabled: server-side native "
"compaction (context_management) is disabled for this agent so the "
"checkpoint-aware Hermes compressor stays authoritative."
)
def native_compaction_context_management(
agent: Any, *, is_codex_backend: bool, is_xai_responses: bool = False, is_github_responses: bool = False,
) -> Optional[List[Dict[str, Any]]]:
def native_compaction_context_management(agent: Any, *, is_codex_backend: bool, is_xai_responses: bool = False,
is_github_responses: bool = False) -> Optional[List[Dict[str, Any]]]:
"""Return the ``context_management`` payload for this request, or None ("do not send").
Every gate is re-checked per request so a mid-session model switch or the in-session
@@ -129,10 +123,8 @@ def native_compaction_context_management(
return None
compressor = getattr(agent, "context_compressor", None)
threshold = resolve_compact_threshold(
getattr(agent, "codex_responses_compact_threshold", None),
getattr(compressor, "threshold_tokens", None) if compressor is not None else None,
)
local_trigger = getattr(compressor, "threshold_tokens", None) if compressor is not None else None
threshold = resolve_compact_threshold(getattr(agent, "codex_responses_compact_threshold", None), local_trigger)
return [{"type": "compaction", "compact_threshold": threshold}]
@@ -151,28 +143,20 @@ def _extract_item_text(item: Any) -> Optional[str]:
"""Measurable text from a Responses item (string/multipart/metadata), or None."""
if not isinstance(item, dict):
return None
content = item.get("content")
if content is None and "output_text" in item:
content = item.get("output_text")
if isinstance(content, str):
return content if content.strip() else None
if not isinstance(content, list):
return None
parts = []
for part in content:
if isinstance(part, str):
candidates = (part,)
elif isinstance(part, dict):
candidates: tuple = (part,) # non-str, non-dict parts filter out below
if isinstance(part, dict):
part_meta = part.get("metadata")
candidates = (
part.get("text") or part.get("input_text") or part.get("output_text"),
part_meta.get("text") if isinstance(part_meta, dict) else None,
)
else:
continue
candidates = (part.get("text") or part.get("input_text") or part.get("output_text"),
part_meta.get("text") if isinstance(part_meta, dict) else None)
parts.extend(c.strip() for c in candidates if isinstance(c, str) and c.strip())
text = " ".join(parts)
return text if text.strip() else None
@@ -183,11 +167,8 @@ def _has_retainable_image_content(item: Any) -> bool:
adapter-owned shape counts, so empty multipart placeholders never become durable history)."""
content = item.get("content") if isinstance(item, dict) else None
return isinstance(content, list) and any(
isinstance(part, dict)
and str(part.get("type") or "").strip().lower() == "input_image"
and isinstance(part.get("image_url"), str)
and part["image_url"].strip()
for part in content
isinstance(part, dict) and str(part.get("type") or "").strip().lower() == "input_image"
and isinstance(part.get("image_url"), str) and part["image_url"].strip() for part in content
)
@@ -204,8 +185,7 @@ def prune_pre_checkpoint_items(
items: List[Dict[str, Any]],
retained_user_token_budget: int = RETAINED_USER_MESSAGE_TOKEN_BUDGET,
retained_summary_token_budget: int = RETAINED_SUMMARY_TOKEN_BUDGET,
enable_summary_retention: bool = True,
item_sources: Optional[List[Any]] = None,
enable_summary_retention: bool = True, item_sources: Optional[List[Any]] = None,
) -> List[Dict[str, Any]]:
"""Restructure Responses input around the newest compaction checkpoint.
@@ -214,12 +194,12 @@ def prune_pre_checkpoint_items(
[checkpoint run] + [retained user & summary messages (newest-first budget)] + [post]
- The NEWEST contiguous run of checkpoints wins.
- The NEWEST contiguous run of checkpoints wins; relative order is preserved.
- User messages are kept verbatim within ``retained_user_token_budget``; the boundary
message is head-truncated when it only partially fits (string content only). A
recognized image-only user message is retained whole at one-token cost.
- Summaries are retained whole within ``retained_summary_token_budget``, never sliced
(framing would corrupt) and never duplicated. Relative order is preserved.
(framing would corrupt) and never duplicated.
- ``item_sources`` (parallel to ``items``) is the raw chat message each item came from.
Conversion can be lossy for summaries (merge-into-tail carrier → typed
``function_call_output``; assistant carrier shadowed by a stale replay), so a source
@@ -246,48 +226,39 @@ def prune_pre_checkpoint_items(
seen_summary_texts: set = set()
def _retain_summary(text: Optional[str], retained_item: Dict[str, Any]) -> None:
"""Retain a summary whole when it fits the budget and is not a duplicate."""
"""Retain a summary whole when it fits the budget and is not a duplicate (never sliced)."""
nonlocal summary_remaining
if not text or summary_remaining <= 0 or text in seen_summary_texts:
return
cost = _approx_tokens(text)
if cost > summary_remaining:
return # never slice a summary's structural framing
seen_summary_texts.add(text)
retained_reversed.append(retained_item)
summary_remaining -= cost
if cost <= summary_remaining:
seen_summary_texts.add(text)
retained_reversed.append(retained_item)
summary_remaining -= cost
for item, source in zip(reversed(pre), reversed(pre_sources)):
if not isinstance(item, dict):
continue
# Source-based detection sees past a lossy conversion; it only fires
# when the source itself is a provenance-tagged summary carrier.
if enable_summary_retention and isinstance(source, dict) and _is_summary_item(source):
text = flatten_message_text(source.get("content"))
_src_role = source.get("role")
_retain_summary(text if text.strip() else None, {
"role": _src_role if _src_role in ("user", "assistant") else "assistant",
"content": text,
})
_retain_summary(text if text.strip() else None,
{"role": _src_role if _src_role in ("user", "assistant") else "assistant", "content": text})
continue
# Typed non-message items never carry role=user or a summary flag.
if "type" in item and item.get("type") != "message":
continue
is_summary = enable_summary_retention and _is_summary_item(item)
is_user = item.get("role") == "user"
if not is_user and not is_summary:
continue
text = _extract_item_text(item)
if text is None:
if not (is_user and _has_retainable_image_content(item)):
continue
text = ""
if is_summary:
_retain_summary(text, item)
elif user_remaining > 0:
@@ -302,11 +273,8 @@ def prune_pre_checkpoint_items(
user_remaining = 0
result = items[first_cp : last_cp + 1] + list(reversed(retained_reversed)) + items[last_cp + 1 :]
logger.debug(
"Pruned pre-checkpoint items: %d input -> %d retained (user_rem=%d, summary_rem=%d)",
len(items), len(result), user_remaining, summary_remaining,
)
logger.debug("Pruned pre-checkpoint items: %d input -> %d retained (user_rem=%d, summary_rem=%d)",
len(items), len(result), user_remaining, summary_remaining)
return result
@@ -320,10 +288,9 @@ _REJECTION_MARKERS = (
def is_native_compaction_rejection(error: Any, status_code: Any = None) -> bool:
"""True when a provider error is a STRUCTURED rejection of ``context_management``.
Drives one-shot recovery (strip, disable for the session, retry), so matching is
narrow: a transient 5xx that merely ECHOES the request must not downgrade native
compaction. Requires ``status_code`` 400 (or unknown) AND the field name with rejection
language.
Drives one-shot recovery (strip, disable for the session, retry), so matching is narrow:
a transient 5xx that merely ECHOES the request must not downgrade native compaction.
Requires ``status_code`` 400 (or unknown) AND the field name with rejection language.
"""
text = str(error or "").lower()
if "context_management" not in text and "compact_threshold" not in text:
@@ -337,11 +304,8 @@ def is_native_compaction_rejection(error: Any, status_code: Any = None) -> bool:
def has_compaction_checkpoint(items: Any) -> bool:
"""Does this ``codex_reasoning_items`` sidecar carry a compaction checkpoint?
A compaction item is cumulative context that exists in exactly one place: anything
that rewrites or discards the sidecar must ask this first or lose the history.
"""
"""Does this ``codex_reasoning_items`` sidecar carry a compaction checkpoint? A checkpoint is
cumulative context living in exactly one place: rewrite/discard the sidecar only after asking."""
return isinstance(items, list) and any(_is_compaction_item(item) for item in items)
@@ -352,9 +316,8 @@ def merge_interim_reasoning_items(prior_items: Any, new_items: Any) -> List[Dict
overwrite drops the only copy: newer items win, prior checkpoints are prepended unless
the newer payload has its own.
"""
kept_checkpoints = [
item for item in (prior_items if isinstance(prior_items, list) else []) if _is_compaction_item(item)
]
prior = prior_items if isinstance(prior_items, list) else []
kept_checkpoints = [item for item in prior if _is_compaction_item(item)]
new_list = list(new_items) if isinstance(new_items, list) else []
if has_compaction_checkpoint(new_list) or not kept_checkpoints:
return new_list
+42 -73
View File
@@ -23,10 +23,9 @@ class PromptCachePlan:
def envelope_tool_part_cache_markers_supported(provider: str | None, base_url: str | None) -> bool:
"""Whether the envelope-layout route honors part-level markers on role:tool.
OpenRouter (and Nous Portal) relocate a part-level ``cache_control`` onto the
``tool_result`` block; LiteLLM-style proxies copy parts verbatim, so the marker lands at
``tool_result.content[0]`` — a non-retryable 400. There, tool messages carry no part
markers and the breakpoint budget reallocates to the nearest eligible message.
OpenRouter/Nous Portal relocate a part-level ``cache_control`` onto the ``tool_result``
block; LiteLLM-style proxies copy parts verbatim, so it lands at ``tool_result.content[0]``
(non-retryable 400). There, role:tool carries no part markers and the budget reallocates.
"""
from agent.agent_runtime_helpers import _is_litellm_route
@@ -40,9 +39,8 @@ def _text_part(text: str, cache_marker: dict | None = None) -> dict:
return part
def _apply_cache_marker(
msg: dict, cache_marker: dict, native_anthropic: bool = False, tool_part_markers: bool = True,
) -> None:
def _apply_cache_marker(msg: dict, cache_marker: dict, native_anthropic: bool = False,
tool_part_markers: bool = True) -> None:
"""Add cache_control to a single message, handling all format variations."""
role = msg.get("role", "")
content = msg.get("content")
@@ -71,10 +69,9 @@ def _apply_cache_marker(
def _can_carry_marker(msg: dict, native_anthropic: bool, tool_part_markers: bool = True) -> bool:
"""True if a marker on this message is actually honored by the provider.
Native Anthropic honors every message. The envelope layout only honors markers inside
content parts, so empty-content messages would waste a breakpoint; with
``tool_part_markers=False`` every role:tool message is excluded too (400). Must agree
with :func:`_apply_cache_marker`, which marks only the LAST part.
Native Anthropic honors every message; the envelope layout only honors markers inside
content parts (empty content wastes a breakpoint) and ``tool_part_markers=False`` excludes
role:tool too (400). Must agree with :func:`_apply_cache_marker` (marks the LAST part).
"""
if native_anthropic:
return True
@@ -117,18 +114,15 @@ def is_qwen_model(model: str) -> bool:
def effective_cache_ttl(ttl: str | None, *, model: str = "", provider: str = "") -> str:
"""Clamp a requested cache TTL to what the destination route supports (``None`` → ``5m``).
Qwen/Alibaba routes drop the ``1h`` tier, so ``1h`` regresses to ``5m`` there — except
on ``MEASURED_1H_PROVIDERS`` (minus ``NO_1H_TIER_MODELS``). The measured-route check
runs BEFORE the generic Qwen clamp, which would otherwise swallow every Qwen model on it.
Qwen/Alibaba routes drop ``1h`` (→ ``5m``) except on ``MEASURED_1H_PROVIDERS`` minus
``NO_1H_TIER_MODELS``; that check runs BEFORE the generic Qwen clamp, which would swallow it.
"""
if ttl != "1h":
return ttl or "5m"
provider_lower = (provider or "").lower()
if provider_lower in MEASURED_1H_PROVIDERS:
return "5m" if _flat_model(model) in NO_1H_TIER_MODELS else "1h"
if is_qwen_model(model) or provider_lower in ALIBABA_FAMILY_PROVIDERS:
return "5m"
return "1h"
return "5m" if is_qwen_model(model) or provider_lower in ALIBABA_FAMILY_PROVIDERS else "1h"
def _apply_system_cache_markers(
@@ -137,19 +131,17 @@ def _apply_system_cache_markers(
) -> int:
"""Mark the static system prefix (and optionally the full prompt); returns markers applied.
The system prompt stays one stored string, split only in the outgoing request.
``mark_suffix=False`` is the tool-cache-plan layout (suffix budget spent on the tools
array). ``fallback_to_whole=False`` marks nothing when the split is impossible. When the
prompt IS the prefix the whole message is one block — never an empty text block (400).
The stored system prompt stays one string, split only in the request. ``mark_suffix=False``
is the tool-cache-plan layout (suffix budget spent on the tools array); ``fallback_to_whole=
False`` marks nothing when the split is impossible. When the prompt IS the prefix the whole
message is one block — never an empty text block (400).
"""
content = message.get("content")
if isinstance(static_system_prefix, str) and static_system_prefix and isinstance(content, str) and content.startswith(static_system_prefix):
suffix = content[len(static_system_prefix):]
if suffix.strip():
message["content"] = [
_text_part(static_system_prefix, cache_marker),
_text_part(suffix, cache_marker if mark_suffix else None),
]
message["content"] = [_text_part(static_system_prefix, cache_marker),
_text_part(suffix, cache_marker if mark_suffix else None)]
return 2 if mark_suffix else 1
elif not fallback_to_whole:
return 0
@@ -158,20 +150,17 @@ def _apply_system_cache_markers(
def _has_part_marker(content: Any) -> bool:
return isinstance(content, list) and any(
isinstance(part, dict) and "cache_control" in part for part in content
)
return isinstance(content, list) and any(isinstance(part, dict) and "cache_control" in part for part in content)
def strip_anthropic_cache_control(api_messages: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
"""Remove ``cache_control`` markers and undo decoration-produced list shapes (in place).
Used before re-decorating after a mid-turn provider failover. Flattening back to a
string is restricted to the exact shapes :func:`apply_anthropic_cache_control` produces
from string content — a single text part, the two-part ``[static, volatile]`` system
split, or the two-part skill split — so the ``""``-join is provably byte-exact. Marker
removal is copy-on-write on part dicts: parts can alias caller-held lists and stripping
must never rewrite the stored transcript.
Used before re-decorating after a mid-turn failover. Flattening to a string is restricted
to the exact shapes :func:`apply_anthropic_cache_control` produces from string content
(single text part, two-part system split, two-part skill split) so the ``""``-join is
byte-exact. Marker removal is copy-on-write on part dicts: parts can alias caller-held
lists and stripping must never rewrite the stored transcript.
"""
for msg in api_messages:
if not isinstance(msg, dict):
@@ -183,21 +172,16 @@ def strip_anthropic_cache_control(api_messages: List[Dict[str, Any]]) -> List[Di
role = msg.get("role")
# The skill split is the only decoration marking the FIRST part of a user message,
# so the shape alone identifies it even after the prefix registry evicted the entry.
skill_split_shape = (
role == "user" and len(content) == 2
and all(isinstance(p, dict) for p in content)
and "cache_control" in content[0] and "cache_control" not in content[1]
)
skill_split_shape = (role == "user" and len(content) == 2 and all(isinstance(p, dict) for p in content)
and "cache_control" in content[0] and "cache_control" not in content[1])
if _has_part_marker(content):
content = msg["content"] = [
{k: v for k, v in part.items() if k != "cache_control"}
if isinstance(part, dict) and "cache_control" in part else part
for part in content
if isinstance(part, dict) and "cache_control" in part else part for part in content
]
plain_text_parts = content and all(
isinstance(part, dict) and part.get("type", "text") == "text"
and isinstance(part.get("text"), str) and set(part.keys()) <= {"type", "text"}
for part in content
and isinstance(part.get("text"), str) and set(part.keys()) <= {"type", "text"} for part in content
)
if plain_text_parts and (len(content) == 1 or (role == "system" and len(content) == 2) or skill_split_shape):
msg["content"] = "".join(part["text"] for part in content)
@@ -215,11 +199,8 @@ def strip_anthropic_tool_cache_control(tools: List[Dict[str, Any]] | None) -> Li
def _count_cache_markers(messages: List[Dict[str, Any]], tools: List[Dict[str, Any]]) -> int:
"""Count the wire-visible cache markers in a request-local plan."""
def _marked(items: Any) -> int:
return sum(1 for item in items if isinstance(item, dict) and "cache_control" in item)
parts = [p for m in messages if isinstance(m, dict) and isinstance(m.get("content"), list) for p in m["content"]]
return _marked(messages) + _marked(parts) + _marked(tools)
return sum(1 for item in [*messages, *parts, *tools] if isinstance(item, dict) and "cache_control" in item)
def _completed_transaction_endpoint_indexes(messages: List[Dict[str, Any]], *, native_anthropic: bool) -> List[int]:
@@ -251,13 +232,9 @@ def _completed_transaction_endpoint_indexes(messages: List[Dict[str, Any]], *, n
index = _tool_run_end(index)
continue
if (role == "user" and index + 1 < len(messages)) or (
role == "assistant" and message.get("content") in (None, "")
):
index += 1
continue
if _can_carry_marker(message, native_anthropic):
open_turn = (role == "user" and index + 1 < len(messages)) or (
role == "assistant" and message.get("content") in (None, ""))
if not open_turn and _can_carry_marker(message, native_anthropic):
endpoints.append(index)
index += 1
return endpoints
@@ -277,18 +254,15 @@ def build_prompt_cache_plan(
if not direct_native_tool_cache or not planned_tools:
planned_messages = apply_anthropic_cache_control(
messages, cache_ttl=cache_ttl, native_anthropic=native_anthropic,
static_system_prefix=static_system_prefix, tool_part_markers=tool_part_markers,
)
static_system_prefix=static_system_prefix, tool_part_markers=tool_part_markers)
return PromptCachePlan(messages=planned_messages, tools=planned_tools)
marker = _build_marker(cache_ttl)
if messages and isinstance(messages[0], dict) and messages[0].get("role") == "system":
# Tool-cache layout: only the static prefix carries a system-side marker; the
# volatile suffix's budget is spent on the tools array.
_apply_system_cache_markers(
messages[0], marker, static_system_prefix,
native_anthropic=True, mark_suffix=False, fallback_to_whole=False,
)
_apply_system_cache_markers(messages[0], marker, static_system_prefix,
native_anthropic=True, mark_suffix=False, fallback_to_whole=False)
planned_tools[-1]["cache_control"] = dict(marker)
for endpoint in _completed_transaction_endpoint_indexes(messages, native_anthropic=True)[-2:]:
_apply_cache_marker(messages[endpoint], marker, native_anthropic=True)
@@ -302,11 +276,10 @@ def apply_anthropic_cache_control(
) -> List[Dict[str, Any]]:
"""Apply Anthropic cache-control markers to API messages.
With a matching ``static_system_prefix`` the prefix and the full system prompt each get
a marker and the remaining two go to the latest cacheable non-system messages; without
it, the legacy system-and-3 layout applies. Idempotent: pre-existing markers are
stripped from a per-message copy first (shallow copy suffices — stripping is
copy-on-write on parts). Returns a shallow list copy with deep copies of modified messages.
With a matching ``static_system_prefix`` the prefix and full system prompt each get a
marker and the remaining two go to the latest cacheable non-system messages; otherwise
the legacy system-and-3 layout applies. Idempotent: pre-existing markers are stripped from
a per-message copy first. Returns a shallow list copy with deep copies of modified messages.
"""
if not api_messages:
return api_messages
@@ -321,15 +294,11 @@ def apply_anthropic_cache_control(
breakpoints_used = 0
if messages[0].get("role") == "system":
messages[0] = copy.deepcopy(messages[0])
breakpoints_used = _apply_system_cache_markers(
messages[0], marker, static_system_prefix, native_anthropic=native_anthropic,
)
breakpoints_used = _apply_system_cache_markers(messages[0], marker, static_system_prefix,
native_anthropic=native_anthropic)
non_sys = [
i for i in range(len(messages))
if messages[i].get("role") != "system"
and _can_carry_marker(messages[i], native_anthropic=native_anthropic, tool_part_markers=tool_part_markers)
]
non_sys = [i for i, m in enumerate(messages) if m.get("role") != "system"
and _can_carry_marker(m, native_anthropic=native_anthropic, tool_part_markers=tool_part_markers)]
for idx in non_sys[-(4 - breakpoints_used):]:
messages[idx] = copy.deepcopy(messages[idx])
_apply_cache_marker(messages[idx], marker, native_anthropic=native_anthropic, tool_part_markers=tool_part_markers)