feat(experiments): DeepSeek stage-1 compaction arms + prefix-preserving summarization
Browse filesThe compaction-experiment middleware stack behind evals.md F35-F38:
- StableToolOutputCapMiddleware: one persistent 40k-byte (nominal 10k-token)
cap when a tool output enters the checkpoint, so every later model call and
the summarizer read identical stable text (prefix-cache friendly), with
capped/original/retained-bytes telemetry and a sha256 audit trail.
- InstrumentedSummarizationMiddleware: trigger-evidence + summary-cost
telemetry (pre/post tokens, summary input/output, retry reasons) on the
stock XML summarization path.
- PrefixPreservingCompactionMiddleware (summarization_strategy=
"structured_prefix"): generates the summary by sending the unchanged
request prefix + one checkpoint instruction with the same model/tools/
settings bound, so the summary call rides the provider cache (94% vs 0%
cache hit on DeepSeek; -14%/trajectory vs the XML arm at identical
compaction dose). Verified analog of Codex's local compaction.
- DeepSeekCacheIsolationMiddleware: per arm/session/trial user_id injection
+ a request-size guard for experiment runs.
- Presets exp_fh_raw / exp_fh_cap10k / exp_c200_raw / exp_c200_cap10k /
exp_c200_cap10k_structured (+ stage-2 exp_c400/c800, built, unrun) in
memory_presets; per-call model_calls usage/cost telemetry with
summarization cost breakdown in telemetry.
- Tests: cap persistence end-to-end, structured retry/tool-call handling,
XML-vs-structured boundary equivalence, multi-compaction, the
no-double-count usage invariant, the langchain-openai cached_tokens
mapping contract, and KB tool-call argument robustness.
Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
- app/chat_service.py +819 -7
- app/chat_types.py +3 -0
- app/memory_presets.py +102 -0
- app/telemetry.py +149 -25
- tests/test_chat_service.py +69 -0
- tests/test_memory_presets.py +58 -0
- tests/test_memory_variants.py +1194 -2
- tests/test_telemetry.py +136 -0
|
@@ -14,11 +14,13 @@ from threading import Lock
|
|
| 14 |
from typing import Any, AsyncIterator
|
| 15 |
from uuid import uuid4
|
| 16 |
|
|
|
|
| 17 |
from langchain.agents import create_agent
|
| 18 |
from langchain.agents.middleware import (
|
| 19 |
AgentMiddleware,
|
| 20 |
ClearToolUsesEdit,
|
| 21 |
ContextEditingMiddleware,
|
|
|
|
| 22 |
SummarizationMiddleware,
|
| 23 |
)
|
| 24 |
from langchain.tools import ToolRuntime, tool
|
|
@@ -27,11 +29,16 @@ from langchain_core.messages import (
|
|
| 27 |
AIMessageChunk,
|
| 28 |
BaseMessage,
|
| 29 |
HumanMessage,
|
|
|
|
| 30 |
SystemMessage,
|
|
|
|
| 31 |
)
|
|
|
|
| 32 |
from langchain_openai import ChatOpenAI
|
| 33 |
from langgraph.checkpoint.memory import InMemorySaver
|
|
|
|
| 34 |
from langgraph.store.memory import InMemoryStore
|
|
|
|
| 35 |
|
| 36 |
from .chat_types import ChatEvent, ChatRequest, ChatTurn, SourceMatch
|
| 37 |
from .memory_presets import (
|
|
@@ -42,9 +49,13 @@ from .memory_presets import (
|
|
| 42 |
)
|
| 43 |
from .telemetry import (
|
| 44 |
TurnUsageHandler,
|
|
|
|
| 45 |
context_window_stats,
|
| 46 |
estimate_cost_usd,
|
|
|
|
| 47 |
pop_turn_signals,
|
|
|
|
|
|
|
| 48 |
record_turn_signal_max,
|
| 49 |
reset_turn_signals,
|
| 50 |
usage_totals,
|
|
@@ -138,6 +149,9 @@ class AppContext:
|
|
| 138 |
kb_session_id: str = ""
|
| 139 |
kb_command_limit: int = DEFAULT_KB_COMMAND_LIMIT
|
| 140 |
student_id: str = ""
|
|
|
|
|
|
|
|
|
|
| 141 |
# Per-request retrieval token budget (Part C / Axis B sweep); None keeps the
|
| 142 |
# retriever's DEFAULT_CONTEXT_TOKEN_BUDGET.
|
| 143 |
retrieval_budget: int | None = None
|
|
@@ -347,12 +361,24 @@ RETRIEVE_TUTOR_CONTEXT_SCHEMA = {
|
|
| 347 |
},
|
| 348 |
},
|
| 349 |
"required": ["query"],
|
|
|
|
| 350 |
}
|
| 351 |
|
| 352 |
|
| 353 |
@tool(args_schema=RETRIEVE_TUTOR_CONTEXT_SCHEMA)
|
| 354 |
-
def retrieve_tutor_context(
|
|
|
|
|
|
|
| 355 |
"""Retrieve relevant course and documentation context for an AI tutor question."""
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 356 |
try:
|
| 357 |
results = select_retriever(
|
| 358 |
getattr(runtime.context, "retriever_kind", "")
|
|
@@ -391,14 +417,28 @@ RUN_KB_COMMAND_SCHEMA = {
|
|
| 391 |
"type": "integer",
|
| 392 |
"description": "Command timeout in seconds, capped by the runtime.",
|
| 393 |
"default": 8,
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 394 |
},
|
| 395 |
"max_output_chars": {
|
| 396 |
"type": "integer",
|
| 397 |
"description": "Maximum stdout/stderr characters to return, capped by the runtime.",
|
| 398 |
"default": 40000,
|
|
|
|
|
|
|
| 399 |
},
|
| 400 |
},
|
| 401 |
"required": ["command"],
|
|
|
|
| 402 |
}
|
| 403 |
|
| 404 |
|
|
@@ -406,10 +446,23 @@ RUN_KB_COMMAND_SCHEMA = {
|
|
| 406 |
def run_kb_command(
|
| 407 |
command: str,
|
| 408 |
runtime: ToolRuntime[AppContext],
|
| 409 |
-
timeout_seconds: int =
|
| 410 |
max_output_chars: int = 40000,
|
|
|
|
|
|
|
| 411 |
) -> str:
|
| 412 |
"""Run a safe, read-only terminal-style command inside the local KB."""
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 413 |
allowed, used = _claim_kb_command_budget(
|
| 414 |
runtime.context.kb_session_id,
|
| 415 |
runtime.context.kb_command_limit,
|
|
@@ -428,11 +481,11 @@ def run_kb_command(
|
|
| 428 |
try:
|
| 429 |
result = execute_kb_command(
|
| 430 |
command,
|
| 431 |
-
timeout_seconds=
|
| 432 |
max_output_chars=max_output_chars,
|
| 433 |
)
|
| 434 |
return format_command_payload(result)
|
| 435 |
-
except (KbCommandError, OSError) as exc:
|
| 436 |
return f"$ {command}\nerror: {exc}"
|
| 437 |
|
| 438 |
|
|
@@ -1200,6 +1253,674 @@ def _turn_id_for(request: Any) -> str:
|
|
| 1200 |
return getattr(ctx, "kb_session_id", "") if ctx else ""
|
| 1201 |
|
| 1202 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1203 |
class SlidingWindowMiddleware(AgentMiddleware):
|
| 1204 |
"""Keep only the last N messages in the model's view; drop older ones.
|
| 1205 |
|
|
@@ -1572,7 +2293,29 @@ def build_agent_middleware(
|
|
| 1572 |
model: Any, memory_config: MemoryConfig
|
| 1573 |
) -> list[AgentMiddleware]:
|
| 1574 |
"""Assemble the compaction/memory middleware stack for one preset."""
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1575 |
middleware: list[AgentMiddleware] = []
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1576 |
if memory_config.context_editing:
|
| 1577 |
middleware.append(
|
| 1578 |
ContextEditingMiddleware(
|
|
@@ -1596,16 +2339,41 @@ def build_agent_middleware(
|
|
| 1596 |
)
|
| 1597 |
)
|
| 1598 |
if memory_config.summarization:
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1599 |
summarization_kwargs: dict[str, Any] = {
|
| 1600 |
"model": model,
|
| 1601 |
"trigger": ("tokens", memory_config.summarization_trigger_tokens),
|
| 1602 |
-
"keep":
|
|
|
|
| 1603 |
}
|
| 1604 |
# A custom summary prompt (selective_retention / context_reset) overrides
|
| 1605 |
# the library default; None keeps it.
|
| 1606 |
if memory_config.summary_prompt:
|
| 1607 |
summarization_kwargs["summary_prompt"] = memory_config.summary_prompt
|
| 1608 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1609 |
# Part C per-call-view mechanisms (each preset enables at most one). They
|
| 1610 |
# reshape only the request, not the checkpoint, and report via the
|
| 1611 |
# turn-signal registry.
|
|
@@ -1640,6 +2408,11 @@ def build_agent_middleware(
|
|
| 1640 |
if memory_config.longterm_memory:
|
| 1641 |
middleware.append(StudentProfileMiddleware())
|
| 1642 |
middleware.append(SourcePreferenceMiddleware())
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 1643 |
return middleware
|
| 1644 |
|
| 1645 |
|
|
@@ -1835,6 +2608,7 @@ def agent_run_config(
|
|
| 1835 |
"include_reasoning": bool(request.include_reasoning),
|
| 1836 |
"memory_preset": preset,
|
| 1837 |
"student_id": request.student_id,
|
|
|
|
| 1838 |
},
|
| 1839 |
}
|
| 1840 |
)
|
|
@@ -1984,6 +2758,7 @@ async def stream_chat(request: ChatRequest) -> AsyncIterator[ChatEvent]:
|
|
| 1984 |
kb_session_id=message_id,
|
| 1985 |
kb_command_limit=DEFAULT_KB_COMMAND_LIMIT,
|
| 1986 |
student_id=request.student_id,
|
|
|
|
| 1987 |
retrieval_budget=request.retrieval_budget,
|
| 1988 |
retriever_kind=request.retriever,
|
| 1989 |
),
|
|
@@ -2078,6 +2853,12 @@ async def stream_chat(request: ChatRequest) -> AsyncIterator[ChatEvent]:
|
|
| 2078 |
# ToolMessage and must not re-emit it as new tool activity.
|
| 2079 |
if step == "tools" and getattr(message, "type", None) == "tool":
|
| 2080 |
payload = message_content_to_text(message.content)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 2081 |
tool_call_id = str(
|
| 2082 |
getattr(message, "tool_call_id", "") or uuid4().hex
|
| 2083 |
)
|
|
@@ -2114,6 +2895,11 @@ async def stream_chat(request: ChatRequest) -> AsyncIterator[ChatEvent]:
|
|
| 2114 |
"args": tool_call.get("args"),
|
| 2115 |
"args_text": format_tool_args(tool_call.get("args")),
|
| 2116 |
"output_text": payload,
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 2117 |
"matches": [
|
| 2118 |
source_match_payload(
|
| 2119 |
match,
|
|
@@ -2282,6 +3068,15 @@ async def stream_chat(request: ChatRequest) -> AsyncIterator[ChatEvent]:
|
|
| 2282 |
"messages", []
|
| 2283 |
)
|
| 2284 |
totals = usage_totals(usage_handler.usage_metadata)
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 2285 |
total_ms = int((time.monotonic() - turn_started) * 1000)
|
| 2286 |
yield ChatEvent(
|
| 2287 |
"context_stats",
|
|
@@ -2296,8 +3091,25 @@ async def stream_chat(request: ChatRequest) -> AsyncIterator[ChatEvent]:
|
|
| 2296 |
model_key: dict(usage)
|
| 2297 |
for model_key, usage in usage_handler.usage_metadata.items()
|
| 2298 |
},
|
|
|
|
| 2299 |
**totals,
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 2300 |
"est_cost_usd": estimate_cost_usd(usage_handler.usage_metadata),
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 2301 |
"ttft_ms": (
|
| 2302 |
int((first_text_at - turn_started) * 1000)
|
| 2303 |
if first_text_at is not None
|
|
@@ -2315,7 +3127,7 @@ async def stream_chat(request: ChatRequest) -> AsyncIterator[ChatEvent]:
|
|
| 2315 |
**context_window_stats(state_messages, CLEARED_TOOL_OUTPUT_PLACEHOLDER),
|
| 2316 |
# Signals from per-call-view middlewares (this turn only); absent
|
| 2317 |
# keys simply mean that mechanism did not fire.
|
| 2318 |
-
**
|
| 2319 |
},
|
| 2320 |
)
|
| 2321 |
yield ChatEvent(
|
|
|
|
| 14 |
from typing import Any, AsyncIterator
|
| 15 |
from uuid import uuid4
|
| 16 |
|
| 17 |
+
import httpx
|
| 18 |
from langchain.agents import create_agent
|
| 19 |
from langchain.agents.middleware import (
|
| 20 |
AgentMiddleware,
|
| 21 |
ClearToolUsesEdit,
|
| 22 |
ContextEditingMiddleware,
|
| 23 |
+
ExtendedModelResponse,
|
| 24 |
SummarizationMiddleware,
|
| 25 |
)
|
| 26 |
from langchain.tools import ToolRuntime, tool
|
|
|
|
| 29 |
AIMessageChunk,
|
| 30 |
BaseMessage,
|
| 31 |
HumanMessage,
|
| 32 |
+
RemoveMessage,
|
| 33 |
SystemMessage,
|
| 34 |
+
ToolMessage,
|
| 35 |
)
|
| 36 |
+
from langchain_core.messages.utils import count_tokens_approximately, get_buffer_string
|
| 37 |
from langchain_openai import ChatOpenAI
|
| 38 |
from langgraph.checkpoint.memory import InMemorySaver
|
| 39 |
+
from langgraph.graph.message import REMOVE_ALL_MESSAGES
|
| 40 |
from langgraph.store.memory import InMemoryStore
|
| 41 |
+
from langgraph.types import Command
|
| 42 |
|
| 43 |
from .chat_types import ChatEvent, ChatRequest, ChatTurn, SourceMatch
|
| 44 |
from .memory_presets import (
|
|
|
|
| 49 |
)
|
| 50 |
from .telemetry import (
|
| 51 |
TurnUsageHandler,
|
| 52 |
+
aggregate_cost_breakdown,
|
| 53 |
context_window_stats,
|
| 54 |
estimate_cost_usd,
|
| 55 |
+
pop_turn_events,
|
| 56 |
pop_turn_signals,
|
| 57 |
+
record_turn_event,
|
| 58 |
+
record_turn_signal,
|
| 59 |
record_turn_signal_max,
|
| 60 |
reset_turn_signals,
|
| 61 |
usage_totals,
|
|
|
|
| 149 |
kb_session_id: str = ""
|
| 150 |
kb_command_limit: int = DEFAULT_KB_COMMAND_LIMIT
|
| 151 |
student_id: str = ""
|
| 152 |
+
# DeepSeek experiment-only cache namespace. Stable within one trajectory,
|
| 153 |
+
# distinct across arm/session/trial, and intentionally contains no PII.
|
| 154 |
+
cache_user_id: str = ""
|
| 155 |
# Per-request retrieval token budget (Part C / Axis B sweep); None keeps the
|
| 156 |
# retriever's DEFAULT_CONTEXT_TOKEN_BUDGET.
|
| 157 |
retrieval_budget: int | None = None
|
|
|
|
| 361 |
},
|
| 362 |
},
|
| 363 |
"required": ["query"],
|
| 364 |
+
"additionalProperties": False,
|
| 365 |
}
|
| 366 |
|
| 367 |
|
| 368 |
@tool(args_schema=RETRIEVE_TUTOR_CONTEXT_SCHEMA)
|
| 369 |
+
def retrieve_tutor_context(
|
| 370 |
+
query: str, runtime: ToolRuntime[AppContext], **unsupported: Any
|
| 371 |
+
) -> str:
|
| 372 |
"""Retrieve relevant course and documentation context for an AI tutor question."""
|
| 373 |
+
if unsupported:
|
| 374 |
+
names = ", ".join(sorted(unsupported))
|
| 375 |
+
logger.warning(
|
| 376 |
+
"retrieve_tutor_context received unsupported arguments: %s", names
|
| 377 |
+
)
|
| 378 |
+
return (
|
| 379 |
+
"retrieve_tutor_context could not run because it received unsupported "
|
| 380 |
+
f"argument(s): {names}. Retry the tool with only the query argument."
|
| 381 |
+
)
|
| 382 |
try:
|
| 383 |
results = select_retriever(
|
| 384 |
getattr(runtime.context, "retriever_kind", "")
|
|
|
|
| 417 |
"type": "integer",
|
| 418 |
"description": "Command timeout in seconds, capped by the runtime.",
|
| 419 |
"default": 8,
|
| 420 |
+
"minimum": 1,
|
| 421 |
+
"maximum": 30,
|
| 422 |
+
},
|
| 423 |
+
"timeout": {
|
| 424 |
+
"type": "integer",
|
| 425 |
+
"description": (
|
| 426 |
+
"Alias for timeout_seconds. The runtime still caps the command "
|
| 427 |
+
"at 30 seconds."
|
| 428 |
+
),
|
| 429 |
+
"minimum": 1,
|
| 430 |
+
"maximum": 30,
|
| 431 |
},
|
| 432 |
"max_output_chars": {
|
| 433 |
"type": "integer",
|
| 434 |
"description": "Maximum stdout/stderr characters to return, capped by the runtime.",
|
| 435 |
"default": 40000,
|
| 436 |
+
"minimum": 1000,
|
| 437 |
+
"maximum": 80000,
|
| 438 |
},
|
| 439 |
},
|
| 440 |
"required": ["command"],
|
| 441 |
+
"additionalProperties": False,
|
| 442 |
}
|
| 443 |
|
| 444 |
|
|
|
|
| 446 |
def run_kb_command(
|
| 447 |
command: str,
|
| 448 |
runtime: ToolRuntime[AppContext],
|
| 449 |
+
timeout_seconds: int | None = None,
|
| 450 |
max_output_chars: int = 40000,
|
| 451 |
+
timeout: int | None = None,
|
| 452 |
+
**unsupported: Any,
|
| 453 |
) -> str:
|
| 454 |
"""Run a safe, read-only terminal-style command inside the local KB."""
|
| 455 |
+
if unsupported:
|
| 456 |
+
names = ", ".join(sorted(unsupported))
|
| 457 |
+
logger.warning("run_kb_command received unsupported arguments: %s", names)
|
| 458 |
+
return (
|
| 459 |
+
f"$ {command}\n"
|
| 460 |
+
f"error: unsupported run_kb_command argument(s): {names}. "
|
| 461 |
+
"Use command, timeout_seconds (or timeout), and max_output_chars."
|
| 462 |
+
)
|
| 463 |
+
effective_timeout = timeout_seconds if timeout_seconds is not None else timeout
|
| 464 |
+
if effective_timeout is None:
|
| 465 |
+
effective_timeout = 8
|
| 466 |
allowed, used = _claim_kb_command_budget(
|
| 467 |
runtime.context.kb_session_id,
|
| 468 |
runtime.context.kb_command_limit,
|
|
|
|
| 481 |
try:
|
| 482 |
result = execute_kb_command(
|
| 483 |
command,
|
| 484 |
+
timeout_seconds=effective_timeout,
|
| 485 |
max_output_chars=max_output_chars,
|
| 486 |
)
|
| 487 |
return format_command_payload(result)
|
| 488 |
+
except (KbCommandError, OSError, TypeError, ValueError) as exc:
|
| 489 |
return f"$ {command}\nerror: {exc}"
|
| 490 |
|
| 491 |
|
|
|
|
| 1253 |
return getattr(ctx, "kb_session_id", "") if ctx else ""
|
| 1254 |
|
| 1255 |
|
| 1256 |
+
def _cache_user_id_for_runtime(runtime: Any) -> str:
|
| 1257 |
+
ctx = getattr(runtime, "context", None) if runtime else None
|
| 1258 |
+
return str(getattr(ctx, "cache_user_id", "") or "") if ctx else ""
|
| 1259 |
+
|
| 1260 |
+
|
| 1261 |
+
class DeepSeekCacheIsolationMiddleware(AgentMiddleware):
|
| 1262 |
+
"""Attach DeepSeek ``user_id`` and guard the experimental request size."""
|
| 1263 |
+
|
| 1264 |
+
def __init__(self, max_request_tokens: int | None = None) -> None:
|
| 1265 |
+
super().__init__()
|
| 1266 |
+
self.max_request_tokens = max_request_tokens
|
| 1267 |
+
|
| 1268 |
+
def _isolate(self, request: Any) -> Any:
|
| 1269 |
+
request_messages = list(getattr(request, "messages", None) or [])
|
| 1270 |
+
system_message = getattr(request, "system_message", None)
|
| 1271 |
+
if system_message is not None:
|
| 1272 |
+
request_messages.insert(0, system_message)
|
| 1273 |
+
request_tokens = int(count_tokens_approximately(request_messages))
|
| 1274 |
+
if (
|
| 1275 |
+
self.max_request_tokens is not None
|
| 1276 |
+
and request_tokens > self.max_request_tokens
|
| 1277 |
+
):
|
| 1278 |
+
raise RuntimeError(
|
| 1279 |
+
"Agent request exceeds the experiment safety guard: "
|
| 1280 |
+
f"{request_tokens:,} > {self.max_request_tokens:,} "
|
| 1281 |
+
"approximate tokens."
|
| 1282 |
+
)
|
| 1283 |
+
user_id = _cache_user_id_for_runtime(getattr(request, "runtime", None))
|
| 1284 |
+
if not user_id:
|
| 1285 |
+
return request
|
| 1286 |
+
settings = dict(getattr(request, "model_settings", None) or {})
|
| 1287 |
+
extra_body = dict(settings.get("extra_body") or {})
|
| 1288 |
+
extra_body["user_id"] = user_id
|
| 1289 |
+
settings["extra_body"] = extra_body
|
| 1290 |
+
return request.override(model_settings=settings)
|
| 1291 |
+
|
| 1292 |
+
def wrap_model_call(self, request, handler):
|
| 1293 |
+
return handler(self._isolate(request))
|
| 1294 |
+
|
| 1295 |
+
async def awrap_model_call(self, request, handler):
|
| 1296 |
+
return await handler(self._isolate(request))
|
| 1297 |
+
|
| 1298 |
+
|
| 1299 |
+
class StableToolOutputCapMiddleware(AgentMiddleware):
|
| 1300 |
+
"""Persistently cap tool output once, when it first enters agent history."""
|
| 1301 |
+
|
| 1302 |
+
def __init__(self, max_bytes: int) -> None:
|
| 1303 |
+
super().__init__()
|
| 1304 |
+
self.max_bytes = max(1_024, int(max_bytes))
|
| 1305 |
+
|
| 1306 |
+
@staticmethod
|
| 1307 |
+
def _decode_fragment(fragment: bytes) -> str:
|
| 1308 |
+
return fragment.decode("utf-8", errors="ignore")
|
| 1309 |
+
|
| 1310 |
+
def _cap(self, request: Any, result: Any) -> Any:
|
| 1311 |
+
if not isinstance(result, ToolMessage):
|
| 1312 |
+
return result
|
| 1313 |
+
text = message_content_to_text(result.content)
|
| 1314 |
+
raw = text.encode("utf-8")
|
| 1315 |
+
if len(raw) <= self.max_bytes:
|
| 1316 |
+
return result
|
| 1317 |
+
|
| 1318 |
+
marker = (
|
| 1319 |
+
f"\n\n[... tool output truncated at stable {self.max_bytes}-byte cap; "
|
| 1320 |
+
"middle omitted ...]\n\n"
|
| 1321 |
+
).encode("utf-8")
|
| 1322 |
+
payload_budget = max(0, self.max_bytes - len(marker))
|
| 1323 |
+
head_bytes = payload_budget // 2
|
| 1324 |
+
tail_bytes = payload_budget - head_bytes
|
| 1325 |
+
capped = (
|
| 1326 |
+
self._decode_fragment(raw[:head_bytes])
|
| 1327 |
+
+ marker.decode("utf-8")
|
| 1328 |
+
+ self._decode_fragment(raw[-tail_bytes:] if tail_bytes else b"")
|
| 1329 |
+
)
|
| 1330 |
+
capped_bytes = len(capped.encode("utf-8"))
|
| 1331 |
+
metadata = {
|
| 1332 |
+
"original_bytes": len(raw),
|
| 1333 |
+
"original_chars": len(text),
|
| 1334 |
+
"retained_bytes": capped_bytes,
|
| 1335 |
+
"sha256": hashlib.sha256(raw).hexdigest(),
|
| 1336 |
+
"max_bytes": self.max_bytes,
|
| 1337 |
+
}
|
| 1338 |
+
turn = _turn_id_for(request)
|
| 1339 |
+
record_turn_signal(turn, "tool_outputs_capped", 1)
|
| 1340 |
+
record_turn_signal(turn, "tool_output_original_bytes", len(raw))
|
| 1341 |
+
record_turn_signal(turn, "tool_output_retained_bytes", capped_bytes)
|
| 1342 |
+
additional = dict(result.additional_kwargs or {})
|
| 1343 |
+
additional["stable_tool_cap"] = metadata
|
| 1344 |
+
return result.model_copy(
|
| 1345 |
+
update={"content": capped, "additional_kwargs": additional}
|
| 1346 |
+
)
|
| 1347 |
+
|
| 1348 |
+
def wrap_tool_call(self, request, handler):
|
| 1349 |
+
return self._cap(request, handler(request))
|
| 1350 |
+
|
| 1351 |
+
async def awrap_tool_call(self, request, handler):
|
| 1352 |
+
return self._cap(request, await handler(request))
|
| 1353 |
+
|
| 1354 |
+
|
| 1355 |
+
class InstrumentedSummarizationMiddleware(SummarizationMiddleware):
|
| 1356 |
+
"""SummarizationMiddleware with full event telemetry and loud failures."""
|
| 1357 |
+
|
| 1358 |
+
MAX_SUMMARY_ATTEMPTS = 3
|
| 1359 |
+
RETRY_BASE_DELAY_SECONDS = 1.0
|
| 1360 |
+
|
| 1361 |
+
def __init__(
|
| 1362 |
+
self,
|
| 1363 |
+
*args: Any,
|
| 1364 |
+
summary_input_guard_tokens: int | None = None,
|
| 1365 |
+
**kwargs: Any,
|
| 1366 |
+
) -> None:
|
| 1367 |
+
super().__init__(*args, **kwargs)
|
| 1368 |
+
self.summary_input_guard_tokens = summary_input_guard_tokens
|
| 1369 |
+
|
| 1370 |
+
def _configured_token_trigger(self) -> int | None:
|
| 1371 |
+
return next(
|
| 1372 |
+
(
|
| 1373 |
+
int(clause["tokens"])
|
| 1374 |
+
for clause in self._trigger_clauses
|
| 1375 |
+
if "tokens" in clause
|
| 1376 |
+
),
|
| 1377 |
+
None,
|
| 1378 |
+
)
|
| 1379 |
+
|
| 1380 |
+
@staticmethod
|
| 1381 |
+
def _last_ai_reported_tokens(messages: list[Any]) -> int:
|
| 1382 |
+
last_ai_message = next(
|
| 1383 |
+
(
|
| 1384 |
+
message
|
| 1385 |
+
for message in reversed(messages)
|
| 1386 |
+
if isinstance(message, AIMessage)
|
| 1387 |
+
),
|
| 1388 |
+
None,
|
| 1389 |
+
)
|
| 1390 |
+
usage = getattr(last_ai_message, "usage_metadata", None) or {}
|
| 1391 |
+
return int(usage.get("total_tokens") or 0)
|
| 1392 |
+
|
| 1393 |
+
def _token_trigger_evidence(
|
| 1394 |
+
self, messages: list[Any], approximate_tokens: int
|
| 1395 |
+
) -> tuple[str, int]:
|
| 1396 |
+
trigger_tokens = self._configured_token_trigger()
|
| 1397 |
+
if trigger_tokens is None:
|
| 1398 |
+
return "non_token", 0
|
| 1399 |
+
approximate_met = approximate_tokens >= trigger_tokens
|
| 1400 |
+
provider_reported_met = self._should_summarize_based_on_reported_tokens(
|
| 1401 |
+
messages, float(trigger_tokens)
|
| 1402 |
+
)
|
| 1403 |
+
reported_tokens = self._last_ai_reported_tokens(messages)
|
| 1404 |
+
if approximate_met and provider_reported_met:
|
| 1405 |
+
return "approximate_and_provider_reported", reported_tokens
|
| 1406 |
+
if approximate_met:
|
| 1407 |
+
return "approximate", reported_tokens
|
| 1408 |
+
if provider_reported_met:
|
| 1409 |
+
return "provider_reported", reported_tokens
|
| 1410 |
+
return "other", reported_tokens
|
| 1411 |
+
|
| 1412 |
+
def _plan_compaction(self, state: Any) -> dict[str, Any] | None:
|
| 1413 |
+
messages = state["messages"]
|
| 1414 |
+
self._ensure_message_ids(messages)
|
| 1415 |
+
total_tokens = int(self.token_counter(messages))
|
| 1416 |
+
if not self._should_summarize(messages, total_tokens):
|
| 1417 |
+
return None
|
| 1418 |
+
trigger_source, trigger_reported_tokens = self._token_trigger_evidence(
|
| 1419 |
+
messages, total_tokens
|
| 1420 |
+
)
|
| 1421 |
+
cutoff_index = self._determine_cutoff_index(messages)
|
| 1422 |
+
if cutoff_index <= 0:
|
| 1423 |
+
return None
|
| 1424 |
+
selected, preserved = self._partition_messages(messages, cutoff_index)
|
| 1425 |
+
trimmed = self._trim_messages_for_summary(selected)
|
| 1426 |
+
if not trimmed:
|
| 1427 |
+
raise RuntimeError("Summarization selected no usable input messages.")
|
| 1428 |
+
summary_input_tokens = int(self._partial_token_counter(trimmed))
|
| 1429 |
+
if (
|
| 1430 |
+
self.summary_input_guard_tokens is not None
|
| 1431 |
+
and summary_input_tokens > self.summary_input_guard_tokens
|
| 1432 |
+
):
|
| 1433 |
+
raise RuntimeError(
|
| 1434 |
+
"Summarization input exceeds the experiment safety guard: "
|
| 1435 |
+
f"{summary_input_tokens:,} > "
|
| 1436 |
+
f"{self.summary_input_guard_tokens:,} approximate tokens."
|
| 1437 |
+
)
|
| 1438 |
+
return {
|
| 1439 |
+
"messages": messages,
|
| 1440 |
+
"pre_tokens": total_tokens,
|
| 1441 |
+
"trigger_source": trigger_source,
|
| 1442 |
+
"trigger_reported_tokens": trigger_reported_tokens,
|
| 1443 |
+
"selected": selected,
|
| 1444 |
+
"preserved": preserved,
|
| 1445 |
+
"trimmed": trimmed,
|
| 1446 |
+
"summary_input_tokens": summary_input_tokens,
|
| 1447 |
+
}
|
| 1448 |
+
|
| 1449 |
+
def _summary_prompt_text(self, trimmed: list[Any]) -> str:
|
| 1450 |
+
formatted = get_buffer_string(trimmed, format="xml")
|
| 1451 |
+
return self.summary_prompt.format(messages=formatted).rstrip()
|
| 1452 |
+
|
| 1453 |
+
def _summary_model(self, runtime: Any) -> Any:
|
| 1454 |
+
user_id = _cache_user_id_for_runtime(runtime)
|
| 1455 |
+
if not user_id:
|
| 1456 |
+
return self.model
|
| 1457 |
+
return self.model.bind(extra_body={"user_id": user_id})
|
| 1458 |
+
|
| 1459 |
+
@staticmethod
|
| 1460 |
+
def _is_retryable_exception(exc: BaseException) -> bool:
|
| 1461 |
+
if isinstance(exc, (TimeoutError, ConnectionError, httpx.TransportError)):
|
| 1462 |
+
return True
|
| 1463 |
+
status = getattr(exc, "status_code", None)
|
| 1464 |
+
if status is None:
|
| 1465 |
+
response = getattr(exc, "response", None)
|
| 1466 |
+
status = getattr(response, "status_code", None)
|
| 1467 |
+
if isinstance(status, int):
|
| 1468 |
+
return status in {408, 409, 425, 429} or status >= 500
|
| 1469 |
+
return type(exc).__name__ in {
|
| 1470 |
+
"APIConnectionError",
|
| 1471 |
+
"APITimeoutError",
|
| 1472 |
+
"InternalServerError",
|
| 1473 |
+
"RateLimitError",
|
| 1474 |
+
}
|
| 1475 |
+
|
| 1476 |
+
@staticmethod
|
| 1477 |
+
def _retry_reason(exc: BaseException) -> str:
|
| 1478 |
+
detail = str(exc).strip()
|
| 1479 |
+
return f"{type(exc).__name__}: {detail}" if detail else type(exc).__name__
|
| 1480 |
+
|
| 1481 |
+
@classmethod
|
| 1482 |
+
def _retry_delay(cls, attempt: int) -> float:
|
| 1483 |
+
return cls.RETRY_BASE_DELAY_SECONDS * (2 ** (attempt - 1))
|
| 1484 |
+
|
| 1485 |
+
def _log_summary_retry(
|
| 1486 |
+
self, runtime: Any, attempt: int, reason: str, delay: float
|
| 1487 |
+
) -> None:
|
| 1488 |
+
ctx = getattr(runtime, "context", None)
|
| 1489 |
+
message_id = str(getattr(ctx, "kb_session_id", "") or "") if ctx else ""
|
| 1490 |
+
logger.warning(
|
| 1491 |
+
"Retrying summarization after attempt %d/%d in %.1fs. "
|
| 1492 |
+
"message_id=%s reason=%s",
|
| 1493 |
+
attempt,
|
| 1494 |
+
self.MAX_SUMMARY_ATTEMPTS,
|
| 1495 |
+
delay,
|
| 1496 |
+
message_id,
|
| 1497 |
+
reason,
|
| 1498 |
+
)
|
| 1499 |
+
|
| 1500 |
+
def _record_compaction(
|
| 1501 |
+
self, runtime: Any, plan: dict[str, Any], summary: str
|
| 1502 |
+
) -> dict[str, Any]:
|
| 1503 |
+
if not summary.strip():
|
| 1504 |
+
raise RuntimeError("Summarization returned an empty summary.")
|
| 1505 |
+
if summary.startswith("Error generating summary:"):
|
| 1506 |
+
raise RuntimeError(summary)
|
| 1507 |
+
new_messages = self._build_new_messages(summary)
|
| 1508 |
+
preserved = plan["preserved"]
|
| 1509 |
+
turn = _cache_user_id_for_runtime(runtime)
|
| 1510 |
+
ctx = getattr(runtime, "context", None)
|
| 1511 |
+
message_id = str(getattr(ctx, "kb_session_id", "") or "") if ctx else ""
|
| 1512 |
+
event = {
|
| 1513 |
+
"event": "summarization",
|
| 1514 |
+
"summary_strategy": str(plan.get("summary_strategy") or "xml"),
|
| 1515 |
+
"configured_trigger_tokens": self._configured_token_trigger(),
|
| 1516 |
+
"trigger_source": plan["trigger_source"],
|
| 1517 |
+
"trigger_reported_tokens": int(plan["trigger_reported_tokens"]),
|
| 1518 |
+
"pre_compaction_tokens_approx": int(plan["pre_tokens"]),
|
| 1519 |
+
"pre_compaction_messages": len(plan["messages"]),
|
| 1520 |
+
"selected_messages": len(plan["selected"]),
|
| 1521 |
+
"selected_tokens_approx": int(
|
| 1522 |
+
self._partial_token_counter(plan["selected"])
|
| 1523 |
+
),
|
| 1524 |
+
"summary_input_messages": len(plan["trimmed"]),
|
| 1525 |
+
"summary_input_tokens_approx": int(plan["summary_input_tokens"]),
|
| 1526 |
+
"summary_input_untrimmed": self.trim_tokens_to_summarize is None
|
| 1527 |
+
and len(plan["trimmed"]) == len(plan["selected"]),
|
| 1528 |
+
"summary_attempts": int(plan.get("summary_attempts") or 1),
|
| 1529 |
+
"summary_retry_reasons": list(plan.get("summary_retry_reasons") or []),
|
| 1530 |
+
"summary_output_tokens_approx": int(
|
| 1531 |
+
count_tokens_approximately([HumanMessage(content=summary)])
|
| 1532 |
+
),
|
| 1533 |
+
"retained_tail_messages": len(preserved),
|
| 1534 |
+
"retained_tail_tokens_approx": int(self._partial_token_counter(preserved)),
|
| 1535 |
+
"post_compaction_tokens_approx": int(
|
| 1536 |
+
self._partial_token_counter([*new_messages, *preserved])
|
| 1537 |
+
),
|
| 1538 |
+
"cache_user_id_present": bool(turn),
|
| 1539 |
+
}
|
| 1540 |
+
event.update(plan.get("summary_request_telemetry") or {})
|
| 1541 |
+
event.update(plan.get("summary_provider_telemetry") or {})
|
| 1542 |
+
record_turn_signal(message_id, "compactions_this_turn", 1)
|
| 1543 |
+
record_turn_signal_max(
|
| 1544 |
+
message_id,
|
| 1545 |
+
"max_pre_compaction_tokens_approx",
|
| 1546 |
+
int(plan["pre_tokens"]),
|
| 1547 |
+
)
|
| 1548 |
+
record_turn_event(message_id, event)
|
| 1549 |
+
return {
|
| 1550 |
+
"messages": [
|
| 1551 |
+
RemoveMessage(id=REMOVE_ALL_MESSAGES),
|
| 1552 |
+
*new_messages,
|
| 1553 |
+
*preserved,
|
| 1554 |
+
]
|
| 1555 |
+
}
|
| 1556 |
+
|
| 1557 |
+
def before_model(self, state, runtime):
|
| 1558 |
+
plan = self._plan_compaction(state)
|
| 1559 |
+
if plan is None:
|
| 1560 |
+
return None
|
| 1561 |
+
model = self._summary_model(runtime)
|
| 1562 |
+
prompt = self._summary_prompt_text(plan["trimmed"])
|
| 1563 |
+
retry_reasons: list[str] = []
|
| 1564 |
+
for attempt in range(1, self.MAX_SUMMARY_ATTEMPTS + 1):
|
| 1565 |
+
try:
|
| 1566 |
+
response = model.invoke(
|
| 1567 |
+
prompt,
|
| 1568 |
+
config={"metadata": {"lc_source": "summarization"}},
|
| 1569 |
+
)
|
| 1570 |
+
plan["summary_provider_telemetry"] = self._summary_provider_telemetry(
|
| 1571 |
+
response
|
| 1572 |
+
)
|
| 1573 |
+
summary = message_content_to_text(response.content).strip()
|
| 1574 |
+
except Exception as exc:
|
| 1575 |
+
if (
|
| 1576 |
+
attempt >= self.MAX_SUMMARY_ATTEMPTS
|
| 1577 |
+
or not self._is_retryable_exception(exc)
|
| 1578 |
+
):
|
| 1579 |
+
raise
|
| 1580 |
+
reason = self._retry_reason(exc)
|
| 1581 |
+
else:
|
| 1582 |
+
if summary:
|
| 1583 |
+
plan["summary_attempts"] = attempt
|
| 1584 |
+
plan["summary_retry_reasons"] = retry_reasons
|
| 1585 |
+
return self._record_compaction(runtime, plan, summary)
|
| 1586 |
+
if attempt >= self.MAX_SUMMARY_ATTEMPTS:
|
| 1587 |
+
raise RuntimeError(
|
| 1588 |
+
"Summarization returned an empty summary after "
|
| 1589 |
+
f"{attempt} attempts."
|
| 1590 |
+
)
|
| 1591 |
+
reason = "empty response"
|
| 1592 |
+
retry_reasons.append(reason)
|
| 1593 |
+
delay = self._retry_delay(attempt)
|
| 1594 |
+
self._log_summary_retry(runtime, attempt, reason, delay)
|
| 1595 |
+
time.sleep(delay)
|
| 1596 |
+
raise AssertionError("unreachable")
|
| 1597 |
+
|
| 1598 |
+
async def abefore_model(self, state, runtime):
|
| 1599 |
+
plan = self._plan_compaction(state)
|
| 1600 |
+
if plan is None:
|
| 1601 |
+
return None
|
| 1602 |
+
model = self._summary_model(runtime)
|
| 1603 |
+
prompt = self._summary_prompt_text(plan["trimmed"])
|
| 1604 |
+
retry_reasons: list[str] = []
|
| 1605 |
+
for attempt in range(1, self.MAX_SUMMARY_ATTEMPTS + 1):
|
| 1606 |
+
try:
|
| 1607 |
+
response = await model.ainvoke(
|
| 1608 |
+
prompt,
|
| 1609 |
+
config={"metadata": {"lc_source": "summarization"}},
|
| 1610 |
+
)
|
| 1611 |
+
plan["summary_provider_telemetry"] = self._summary_provider_telemetry(
|
| 1612 |
+
response
|
| 1613 |
+
)
|
| 1614 |
+
summary = message_content_to_text(response.content).strip()
|
| 1615 |
+
except Exception as exc:
|
| 1616 |
+
if (
|
| 1617 |
+
attempt >= self.MAX_SUMMARY_ATTEMPTS
|
| 1618 |
+
or not self._is_retryable_exception(exc)
|
| 1619 |
+
):
|
| 1620 |
+
raise
|
| 1621 |
+
reason = self._retry_reason(exc)
|
| 1622 |
+
else:
|
| 1623 |
+
if summary:
|
| 1624 |
+
plan["summary_attempts"] = attempt
|
| 1625 |
+
plan["summary_retry_reasons"] = retry_reasons
|
| 1626 |
+
return self._record_compaction(runtime, plan, summary)
|
| 1627 |
+
if attempt >= self.MAX_SUMMARY_ATTEMPTS:
|
| 1628 |
+
raise RuntimeError(
|
| 1629 |
+
"Summarization returned an empty summary after "
|
| 1630 |
+
f"{attempt} attempts."
|
| 1631 |
+
)
|
| 1632 |
+
reason = "empty response"
|
| 1633 |
+
retry_reasons.append(reason)
|
| 1634 |
+
delay = self._retry_delay(attempt)
|
| 1635 |
+
self._log_summary_retry(runtime, attempt, reason, delay)
|
| 1636 |
+
await asyncio.sleep(delay)
|
| 1637 |
+
raise AssertionError("unreachable")
|
| 1638 |
+
|
| 1639 |
+
@staticmethod
|
| 1640 |
+
def _summary_provider_telemetry(response: Any) -> dict[str, Any]:
|
| 1641 |
+
"""Expose provider cache accounting on the compaction event itself."""
|
| 1642 |
+
usage = dict(getattr(response, "usage_metadata", None) or {})
|
| 1643 |
+
details = dict(usage.get("input_token_details") or {})
|
| 1644 |
+
input_tokens = int(usage.get("input_tokens") or 0)
|
| 1645 |
+
cache_read = int(details.get("cache_read") or 0)
|
| 1646 |
+
cache_creation = int(details.get("cache_creation") or 0)
|
| 1647 |
+
cache_miss = max(0, input_tokens - cache_read - cache_creation)
|
| 1648 |
+
return {
|
| 1649 |
+
"summary_provider_usage_reported": bool(usage),
|
| 1650 |
+
"summary_provider_cache_details_reported": "cache_read" in details,
|
| 1651 |
+
"summary_provider_input_tokens": input_tokens,
|
| 1652 |
+
"summary_provider_cache_read_tokens": cache_read,
|
| 1653 |
+
"summary_provider_cache_creation_tokens": cache_creation,
|
| 1654 |
+
"summary_provider_cache_miss_tokens": cache_miss,
|
| 1655 |
+
"summary_provider_output_tokens": int(usage.get("output_tokens") or 0),
|
| 1656 |
+
"summary_provider_cache_hit_ratio": (
|
| 1657 |
+
cache_read / input_tokens if input_tokens else None
|
| 1658 |
+
),
|
| 1659 |
+
}
|
| 1660 |
+
|
| 1661 |
+
|
| 1662 |
+
PREFIX_PRESERVING_COMPACTION_PROMPT = """Create a durable checkpoint summary of the older conversation prefix above.
|
| 1663 |
+
|
| 1664 |
+
The first {selected_messages} conversation messages will be replaced by your checkpoint. The final {retained_messages} messages (approximately {retained_tokens} tokens) will remain verbatim. Summarize only the older prefix; do not duplicate the retained tail except where a short reference is necessary to explain a dependency between old and recent work.
|
| 1665 |
+
|
| 1666 |
+
Preserve concrete facts, current values after corrections, decisions, constraints, unresolved work, prior tool evidence needed later, and enough causal detail to continue without re-running tools. Omit transient chatter and redundant wording. Never call a tool and do not answer the conversation's latest question.
|
| 1667 |
+
|
| 1668 |
+
Return only the checkpoint summary."""
|
| 1669 |
+
|
| 1670 |
+
|
| 1671 |
+
class PrefixPreservingCompactionMiddleware(InstrumentedSummarizationMiddleware):
|
| 1672 |
+
"""Compact through a structured prefix-extension request.
|
| 1673 |
+
|
| 1674 |
+
LangChain's stock summarizer serializes selected messages into a new XML
|
| 1675 |
+
prompt, so even the summary-generation call loses the provider's cached
|
| 1676 |
+
prefix. This middleware instead runs after request-shaping middleware and
|
| 1677 |
+
binds the exact same model, tool schemas, tool choice, model settings, and
|
| 1678 |
+
system message. It then sends the unchanged *entire* current request prefix
|
| 1679 |
+
followed by one checkpoint instruction. The resulting checkpoint summary
|
| 1680 |
+
may overlap the retained tail; that conservative duplication matches the
|
| 1681 |
+
cache-friendly local Codex pattern and reduces the chance of lost evidence.
|
| 1682 |
+
|
| 1683 |
+
Installing the resulting summary necessarily changes the *next* agent
|
| 1684 |
+
prefix; no client-side middleware can avoid that boundary without a
|
| 1685 |
+
provider-native opaque continuation/compaction primitive.
|
| 1686 |
+
"""
|
| 1687 |
+
|
| 1688 |
+
def before_model(self, state, runtime):
|
| 1689 |
+
# Planning must happen after all request-shaping middleware has produced
|
| 1690 |
+
# the final ModelRequest, so wrap_model_call owns the operation.
|
| 1691 |
+
return None
|
| 1692 |
+
|
| 1693 |
+
async def abefore_model(self, state, runtime):
|
| 1694 |
+
return None
|
| 1695 |
+
|
| 1696 |
+
def _checkpoint_instruction(self, plan: dict[str, Any]) -> HumanMessage:
|
| 1697 |
+
return HumanMessage(
|
| 1698 |
+
content=PREFIX_PRESERVING_COMPACTION_PROMPT.format(
|
| 1699 |
+
selected_messages=len(plan["selected"]),
|
| 1700 |
+
retained_messages=len(plan["preserved"]),
|
| 1701 |
+
retained_tokens=int(self._partial_token_counter(plan["preserved"])),
|
| 1702 |
+
)
|
| 1703 |
+
)
|
| 1704 |
+
|
| 1705 |
+
def _summary_messages(
|
| 1706 |
+
self, request: Any, plan: dict[str, Any]
|
| 1707 |
+
) -> list[BaseMessage]:
|
| 1708 |
+
messages: list[BaseMessage] = []
|
| 1709 |
+
if request.system_message is not None:
|
| 1710 |
+
messages.append(request.system_message)
|
| 1711 |
+
messages.extend(request.messages)
|
| 1712 |
+
messages.append(self._checkpoint_instruction(plan))
|
| 1713 |
+
return messages
|
| 1714 |
+
|
| 1715 |
+
def _prepare_summary_request(
|
| 1716 |
+
self, request: Any, plan: dict[str, Any]
|
| 1717 |
+
) -> tuple[Any, list[BaseMessage]]:
|
| 1718 |
+
if request.response_format is not None:
|
| 1719 |
+
raise RuntimeError(
|
| 1720 |
+
"Structured-prefix compaction does not support a structured "
|
| 1721 |
+
"agent response format."
|
| 1722 |
+
)
|
| 1723 |
+
if list(request.messages) != list(plan["messages"]):
|
| 1724 |
+
raise RuntimeError(
|
| 1725 |
+
"Structured-prefix compaction requires the finalized model view "
|
| 1726 |
+
"to match checkpoint history exactly; an earlier middleware "
|
| 1727 |
+
"changed request.messages without persisting that change."
|
| 1728 |
+
)
|
| 1729 |
+
messages = self._summary_messages(request, plan)
|
| 1730 |
+
request_tokens = int(count_tokens_approximately(messages))
|
| 1731 |
+
if (
|
| 1732 |
+
self.summary_input_guard_tokens is not None
|
| 1733 |
+
and request_tokens > self.summary_input_guard_tokens
|
| 1734 |
+
):
|
| 1735 |
+
raise RuntimeError(
|
| 1736 |
+
"Structured-prefix summary request exceeds the experiment safety "
|
| 1737 |
+
f"guard: {request_tokens:,} > "
|
| 1738 |
+
f"{self.summary_input_guard_tokens:,} approximate tokens."
|
| 1739 |
+
)
|
| 1740 |
+
|
| 1741 |
+
settings = dict(request.model_settings or {})
|
| 1742 |
+
if request.tools:
|
| 1743 |
+
summary_model = request.model.bind_tools(
|
| 1744 |
+
request.tools,
|
| 1745 |
+
tool_choice=request.tool_choice,
|
| 1746 |
+
**settings,
|
| 1747 |
+
)
|
| 1748 |
+
else:
|
| 1749 |
+
summary_model = request.model.bind(**settings)
|
| 1750 |
+
|
| 1751 |
+
system_tokens = int(
|
| 1752 |
+
count_tokens_approximately([request.system_message])
|
| 1753 |
+
if request.system_message is not None
|
| 1754 |
+
else 0
|
| 1755 |
+
)
|
| 1756 |
+
instruction_tokens = int(count_tokens_approximately([messages[-1]]))
|
| 1757 |
+
plan["summary_strategy"] = "structured_prefix"
|
| 1758 |
+
plan["summary_request_telemetry"] = {
|
| 1759 |
+
"summary_prefix_messages": len(request.messages),
|
| 1760 |
+
"summary_prefix_tokens_approx": int(
|
| 1761 |
+
count_tokens_approximately(request.messages)
|
| 1762 |
+
),
|
| 1763 |
+
"summary_selected_messages": len(plan["trimmed"]),
|
| 1764 |
+
"summary_selected_tokens_approx": int(plan["summary_input_tokens"]),
|
| 1765 |
+
"summary_request_messages": len(messages),
|
| 1766 |
+
"summary_request_tokens_approx": request_tokens,
|
| 1767 |
+
"summary_request_is_strict_extension": True,
|
| 1768 |
+
"summary_instruction_selected_messages": len(plan["selected"]),
|
| 1769 |
+
"summary_instruction_retained_messages": len(plan["preserved"]),
|
| 1770 |
+
"summary_instruction_retained_tokens_approx": int(
|
| 1771 |
+
self._partial_token_counter(plan["preserved"])
|
| 1772 |
+
),
|
| 1773 |
+
"summary_system_tokens_approx": system_tokens,
|
| 1774 |
+
"summary_instruction_tokens_approx": instruction_tokens,
|
| 1775 |
+
"summary_system_message_present": request.system_message is not None,
|
| 1776 |
+
"summary_tools_bound": len(request.tools or []),
|
| 1777 |
+
"summary_tool_choice_preserved": True,
|
| 1778 |
+
"summary_tool_choice_value": (
|
| 1779 |
+
str(request.tool_choice) if request.tool_choice is not None else None
|
| 1780 |
+
),
|
| 1781 |
+
"summary_model_settings_keys": sorted(settings),
|
| 1782 |
+
"summary_cache_user_id_preserved": bool(
|
| 1783 |
+
(settings.get("extra_body") or {}).get("user_id")
|
| 1784 |
+
),
|
| 1785 |
+
}
|
| 1786 |
+
return summary_model, messages
|
| 1787 |
+
|
| 1788 |
+
def _invoke_prefix_summary(self, request: Any, plan: dict[str, Any]) -> str:
|
| 1789 |
+
model, messages = self._prepare_summary_request(request, plan)
|
| 1790 |
+
retry_reasons: list[str] = []
|
| 1791 |
+
for attempt in range(1, self.MAX_SUMMARY_ATTEMPTS + 1):
|
| 1792 |
+
try:
|
| 1793 |
+
response = model.invoke(
|
| 1794 |
+
messages,
|
| 1795 |
+
config={
|
| 1796 |
+
"metadata": {
|
| 1797 |
+
"lc_source": "summarization",
|
| 1798 |
+
"compaction_strategy": "structured_prefix",
|
| 1799 |
+
}
|
| 1800 |
+
},
|
| 1801 |
+
)
|
| 1802 |
+
plan["summary_provider_telemetry"] = self._summary_provider_telemetry(
|
| 1803 |
+
response
|
| 1804 |
+
)
|
| 1805 |
+
summary = message_content_to_text(response.content).strip()
|
| 1806 |
+
except Exception as exc:
|
| 1807 |
+
if (
|
| 1808 |
+
attempt >= self.MAX_SUMMARY_ATTEMPTS
|
| 1809 |
+
or not self._is_retryable_exception(exc)
|
| 1810 |
+
):
|
| 1811 |
+
raise
|
| 1812 |
+
reason = self._retry_reason(exc)
|
| 1813 |
+
else:
|
| 1814 |
+
if summary:
|
| 1815 |
+
plan["summary_attempts"] = attempt
|
| 1816 |
+
plan["summary_retry_reasons"] = retry_reasons
|
| 1817 |
+
return summary
|
| 1818 |
+
if attempt >= self.MAX_SUMMARY_ATTEMPTS:
|
| 1819 |
+
raise RuntimeError(
|
| 1820 |
+
"Structured-prefix summarization returned an empty summary "
|
| 1821 |
+
f"after {attempt} attempts."
|
| 1822 |
+
)
|
| 1823 |
+
reason = "empty response"
|
| 1824 |
+
retry_reasons.append(reason)
|
| 1825 |
+
delay = self._retry_delay(attempt)
|
| 1826 |
+
self._log_summary_retry(request.runtime, attempt, reason, delay)
|
| 1827 |
+
time.sleep(delay)
|
| 1828 |
+
raise AssertionError("unreachable")
|
| 1829 |
+
|
| 1830 |
+
async def _ainvoke_prefix_summary(self, request: Any, plan: dict[str, Any]) -> str:
|
| 1831 |
+
model, messages = self._prepare_summary_request(request, plan)
|
| 1832 |
+
retry_reasons: list[str] = []
|
| 1833 |
+
for attempt in range(1, self.MAX_SUMMARY_ATTEMPTS + 1):
|
| 1834 |
+
try:
|
| 1835 |
+
response = await model.ainvoke(
|
| 1836 |
+
messages,
|
| 1837 |
+
config={
|
| 1838 |
+
"metadata": {
|
| 1839 |
+
"lc_source": "summarization",
|
| 1840 |
+
"compaction_strategy": "structured_prefix",
|
| 1841 |
+
}
|
| 1842 |
+
},
|
| 1843 |
+
)
|
| 1844 |
+
plan["summary_provider_telemetry"] = self._summary_provider_telemetry(
|
| 1845 |
+
response
|
| 1846 |
+
)
|
| 1847 |
+
summary = message_content_to_text(response.content).strip()
|
| 1848 |
+
except Exception as exc:
|
| 1849 |
+
if (
|
| 1850 |
+
attempt >= self.MAX_SUMMARY_ATTEMPTS
|
| 1851 |
+
or not self._is_retryable_exception(exc)
|
| 1852 |
+
):
|
| 1853 |
+
raise
|
| 1854 |
+
reason = self._retry_reason(exc)
|
| 1855 |
+
else:
|
| 1856 |
+
if summary:
|
| 1857 |
+
plan["summary_attempts"] = attempt
|
| 1858 |
+
plan["summary_retry_reasons"] = retry_reasons
|
| 1859 |
+
return summary
|
| 1860 |
+
if attempt >= self.MAX_SUMMARY_ATTEMPTS:
|
| 1861 |
+
raise RuntimeError(
|
| 1862 |
+
"Structured-prefix summarization returned an empty summary "
|
| 1863 |
+
f"after {attempt} attempts."
|
| 1864 |
+
)
|
| 1865 |
+
reason = "empty response"
|
| 1866 |
+
retry_reasons.append(reason)
|
| 1867 |
+
delay = self._retry_delay(attempt)
|
| 1868 |
+
self._log_summary_retry(request.runtime, attempt, reason, delay)
|
| 1869 |
+
await asyncio.sleep(delay)
|
| 1870 |
+
raise AssertionError("unreachable")
|
| 1871 |
+
|
| 1872 |
+
def _compacted_request_and_command(
|
| 1873 |
+
self, request: Any, plan: dict[str, Any], summary: str
|
| 1874 |
+
) -> tuple[Any, list[BaseMessage]]:
|
| 1875 |
+
update = self._record_compaction(request.runtime, plan, summary)
|
| 1876 |
+
compacted = list(update["messages"])[1:]
|
| 1877 |
+
compacted_state = dict(request.state)
|
| 1878 |
+
compacted_state["messages"] = compacted
|
| 1879 |
+
return request.override(messages=compacted, state=compacted_state), compacted
|
| 1880 |
+
|
| 1881 |
+
@staticmethod
|
| 1882 |
+
def _with_checkpoint_command(
|
| 1883 |
+
response: Any, compacted: list[BaseMessage]
|
| 1884 |
+
) -> ExtendedModelResponse:
|
| 1885 |
+
# The model response command is applied first by create_agent. The
|
| 1886 |
+
# additional command then replaces the old checkpoint and re-adds that
|
| 1887 |
+
# same response, so it survives REMOVE_ALL_MESSAGES exactly once.
|
| 1888 |
+
return ExtendedModelResponse(
|
| 1889 |
+
model_response=response,
|
| 1890 |
+
command=Command(
|
| 1891 |
+
update={
|
| 1892 |
+
"messages": [
|
| 1893 |
+
RemoveMessage(id=REMOVE_ALL_MESSAGES),
|
| 1894 |
+
*compacted,
|
| 1895 |
+
*response.result,
|
| 1896 |
+
]
|
| 1897 |
+
}
|
| 1898 |
+
),
|
| 1899 |
+
)
|
| 1900 |
+
|
| 1901 |
+
def wrap_model_call(self, request, handler):
|
| 1902 |
+
plan = self._plan_compaction(request.state)
|
| 1903 |
+
if plan is None:
|
| 1904 |
+
return handler(request)
|
| 1905 |
+
summary = self._invoke_prefix_summary(request, plan)
|
| 1906 |
+
compacted_request, compacted = self._compacted_request_and_command(
|
| 1907 |
+
request, plan, summary
|
| 1908 |
+
)
|
| 1909 |
+
response = handler(compacted_request)
|
| 1910 |
+
return self._with_checkpoint_command(response, compacted)
|
| 1911 |
+
|
| 1912 |
+
async def awrap_model_call(self, request, handler):
|
| 1913 |
+
plan = self._plan_compaction(request.state)
|
| 1914 |
+
if plan is None:
|
| 1915 |
+
return await handler(request)
|
| 1916 |
+
summary = await self._ainvoke_prefix_summary(request, plan)
|
| 1917 |
+
compacted_request, compacted = self._compacted_request_and_command(
|
| 1918 |
+
request, plan, summary
|
| 1919 |
+
)
|
| 1920 |
+
response = await handler(compacted_request)
|
| 1921 |
+
return self._with_checkpoint_command(response, compacted)
|
| 1922 |
+
|
| 1923 |
+
|
| 1924 |
class SlidingWindowMiddleware(AgentMiddleware):
|
| 1925 |
"""Keep only the last N messages in the model's view; drop older ones.
|
| 1926 |
|
|
|
|
| 2293 |
model: Any, memory_config: MemoryConfig
|
| 2294 |
) -> list[AgentMiddleware]:
|
| 2295 |
"""Assemble the compaction/memory middleware stack for one preset."""
|
| 2296 |
+
if memory_config.summarization_strategy not in {"xml", "structured_prefix"}:
|
| 2297 |
+
raise ValueError(
|
| 2298 |
+
f"Unknown summarization strategy: {memory_config.summarization_strategy!r}"
|
| 2299 |
+
)
|
| 2300 |
+
if (
|
| 2301 |
+
memory_config.summarization_strategy == "structured_prefix"
|
| 2302 |
+
and not memory_config.summarization
|
| 2303 |
+
):
|
| 2304 |
+
raise ValueError(
|
| 2305 |
+
"structured_prefix summarization strategy requires summarization=True"
|
| 2306 |
+
)
|
| 2307 |
middleware: list[AgentMiddleware] = []
|
| 2308 |
+
prefix_compactor: PrefixPreservingCompactionMiddleware | None = None
|
| 2309 |
+
if memory_config.experiment_mode:
|
| 2310 |
+
middleware.append(
|
| 2311 |
+
DeepSeekCacheIsolationMiddleware(
|
| 2312 |
+
memory_config.experiment_request_guard_tokens
|
| 2313 |
+
)
|
| 2314 |
+
)
|
| 2315 |
+
if memory_config.tool_output_cap_bytes is not None:
|
| 2316 |
+
middleware.append(
|
| 2317 |
+
StableToolOutputCapMiddleware(memory_config.tool_output_cap_bytes)
|
| 2318 |
+
)
|
| 2319 |
if memory_config.context_editing:
|
| 2320 |
middleware.append(
|
| 2321 |
ContextEditingMiddleware(
|
|
|
|
| 2339 |
)
|
| 2340 |
)
|
| 2341 |
if memory_config.summarization:
|
| 2342 |
+
keep: tuple[str, int]
|
| 2343 |
+
if memory_config.summarization_keep_tokens is not None:
|
| 2344 |
+
keep = ("tokens", memory_config.summarization_keep_tokens)
|
| 2345 |
+
else:
|
| 2346 |
+
keep = ("messages", memory_config.summarization_keep_messages)
|
| 2347 |
summarization_kwargs: dict[str, Any] = {
|
| 2348 |
"model": model,
|
| 2349 |
"trigger": ("tokens", memory_config.summarization_trigger_tokens),
|
| 2350 |
+
"keep": keep,
|
| 2351 |
+
"trim_tokens_to_summarize": memory_config.summarization_trim_tokens,
|
| 2352 |
}
|
| 2353 |
# A custom summary prompt (selective_retention / context_reset) overrides
|
| 2354 |
# the library default; None keeps it.
|
| 2355 |
if memory_config.summary_prompt:
|
| 2356 |
summarization_kwargs["summary_prompt"] = memory_config.summary_prompt
|
| 2357 |
+
if memory_config.experiment_mode:
|
| 2358 |
+
summarization_kwargs["summary_input_guard_tokens"] = (
|
| 2359 |
+
memory_config.summarization_input_guard_tokens
|
| 2360 |
+
)
|
| 2361 |
+
if memory_config.summarization_strategy == "structured_prefix":
|
| 2362 |
+
if not memory_config.experiment_mode:
|
| 2363 |
+
raise ValueError(
|
| 2364 |
+
"structured_prefix summarization is restricted to "
|
| 2365 |
+
"experiment-mode presets"
|
| 2366 |
+
)
|
| 2367 |
+
prefix_compactor = PrefixPreservingCompactionMiddleware(
|
| 2368 |
+
**summarization_kwargs
|
| 2369 |
+
)
|
| 2370 |
+
else:
|
| 2371 |
+
summary_middleware = (
|
| 2372 |
+
InstrumentedSummarizationMiddleware
|
| 2373 |
+
if memory_config.experiment_mode
|
| 2374 |
+
else SummarizationMiddleware
|
| 2375 |
+
)
|
| 2376 |
+
middleware.append(summary_middleware(**summarization_kwargs))
|
| 2377 |
# Part C per-call-view mechanisms (each preset enables at most one). They
|
| 2378 |
# reshape only the request, not the checkpoint, and report via the
|
| 2379 |
# turn-signal registry.
|
|
|
|
| 2408 |
if memory_config.longterm_memory:
|
| 2409 |
middleware.append(StudentProfileMiddleware())
|
| 2410 |
middleware.append(SourcePreferenceMiddleware())
|
| 2411 |
+
if prefix_compactor is not None:
|
| 2412 |
+
# wrap_model_call middleware compose left-to-right (first is outermost).
|
| 2413 |
+
# Keep this last so it observes the finalized system message from every
|
| 2414 |
+
# preceding request-shaping middleware before issuing the prefix call.
|
| 2415 |
+
middleware.append(prefix_compactor)
|
| 2416 |
return middleware
|
| 2417 |
|
| 2418 |
|
|
|
|
| 2608 |
"include_reasoning": bool(request.include_reasoning),
|
| 2609 |
"memory_preset": preset,
|
| 2610 |
"student_id": request.student_id,
|
| 2611 |
+
"cache_user_id": request.cache_user_id,
|
| 2612 |
},
|
| 2613 |
}
|
| 2614 |
)
|
|
|
|
| 2758 |
kb_session_id=message_id,
|
| 2759 |
kb_command_limit=DEFAULT_KB_COMMAND_LIMIT,
|
| 2760 |
student_id=request.student_id,
|
| 2761 |
+
cache_user_id=request.cache_user_id,
|
| 2762 |
retrieval_budget=request.retrieval_budget,
|
| 2763 |
retriever_kind=request.retriever,
|
| 2764 |
),
|
|
|
|
| 2853 |
# ToolMessage and must not re-emit it as new tool activity.
|
| 2854 |
if step == "tools" and getattr(message, "type", None) == "tool":
|
| 2855 |
payload = message_content_to_text(message.content)
|
| 2856 |
+
cap_metadata = dict(
|
| 2857 |
+
(getattr(message, "additional_kwargs", None) or {}).get(
|
| 2858 |
+
"stable_tool_cap"
|
| 2859 |
+
)
|
| 2860 |
+
or {}
|
| 2861 |
+
)
|
| 2862 |
tool_call_id = str(
|
| 2863 |
getattr(message, "tool_call_id", "") or uuid4().hex
|
| 2864 |
)
|
|
|
|
| 2895 |
"args": tool_call.get("args"),
|
| 2896 |
"args_text": format_tool_args(tool_call.get("args")),
|
| 2897 |
"output_text": payload,
|
| 2898 |
+
"output_was_capped": bool(cap_metadata),
|
| 2899 |
+
"output_original_bytes": cap_metadata.get("original_bytes"),
|
| 2900 |
+
"output_original_chars": cap_metadata.get("original_chars"),
|
| 2901 |
+
"output_retained_bytes": cap_metadata.get("retained_bytes"),
|
| 2902 |
+
"output_sha256": cap_metadata.get("sha256"),
|
| 2903 |
"matches": [
|
| 2904 |
source_match_payload(
|
| 2905 |
match,
|
|
|
|
| 3068 |
"messages", []
|
| 3069 |
)
|
| 3070 |
totals = usage_totals(usage_handler.usage_metadata)
|
| 3071 |
+
cost_breakdown = aggregate_cost_breakdown(usage_handler.usage_metadata)
|
| 3072 |
+
model_calls = list(usage_handler.model_calls)
|
| 3073 |
+
summarization_cost = sum(
|
| 3074 |
+
float((call.get("cost") or {}).get("total_usd") or 0)
|
| 3075 |
+
for call in model_calls
|
| 3076 |
+
if call.get("source") == "summarization"
|
| 3077 |
+
)
|
| 3078 |
+
turn_events = pop_turn_events(message_id)
|
| 3079 |
+
turn_signals = pop_turn_signals(message_id)
|
| 3080 |
total_ms = int((time.monotonic() - turn_started) * 1000)
|
| 3081 |
yield ChatEvent(
|
| 3082 |
"context_stats",
|
|
|
|
| 3091 |
model_key: dict(usage)
|
| 3092 |
for model_key, usage in usage_handler.usage_metadata.items()
|
| 3093 |
},
|
| 3094 |
+
"model_calls": model_calls,
|
| 3095 |
**totals,
|
| 3096 |
+
"cache_miss_tokens": max(
|
| 3097 |
+
0,
|
| 3098 |
+
totals["input_tokens"]
|
| 3099 |
+
- totals["cache_read_tokens"]
|
| 3100 |
+
- totals["cache_creation_tokens"],
|
| 3101 |
+
),
|
| 3102 |
"est_cost_usd": estimate_cost_usd(usage_handler.usage_metadata),
|
| 3103 |
+
"cost_breakdown": cost_breakdown,
|
| 3104 |
+
"summarization_cost_usd": summarization_cost,
|
| 3105 |
+
"max_request_context_tokens_approx": max(
|
| 3106 |
+
(
|
| 3107 |
+
int(call.get("request_context_tokens_approx") or 0)
|
| 3108 |
+
for call in model_calls
|
| 3109 |
+
),
|
| 3110 |
+
default=0,
|
| 3111 |
+
),
|
| 3112 |
+
"compaction_events": turn_events,
|
| 3113 |
"ttft_ms": (
|
| 3114 |
int((first_text_at - turn_started) * 1000)
|
| 3115 |
if first_text_at is not None
|
|
|
|
| 3127 |
**context_window_stats(state_messages, CLEARED_TOOL_OUTPUT_PLACEHOLDER),
|
| 3128 |
# Signals from per-call-view middlewares (this turn only); absent
|
| 3129 |
# keys simply mean that mechanism did not fire.
|
| 3130 |
+
**turn_signals,
|
| 3131 |
},
|
| 3132 |
)
|
| 3133 |
yield ChatEvent(
|
|
@@ -39,6 +39,9 @@ class ChatRequest:
|
|
| 39 |
# Long-term memory key: profile-memory presets read and update the stored
|
| 40 |
# student profile under this id. Empty disables profile memory I/O.
|
| 41 |
student_id: str = ""
|
|
|
|
|
|
|
|
|
|
| 42 |
# Part C / Axis B ablation: drop the run_kb_command tool (and its prompt
|
| 43 |
# section) while keeping retrieval, to measure whether KB browsing helps.
|
| 44 |
disable_kb: bool = False
|
|
|
|
| 39 |
# Long-term memory key: profile-memory presets read and update the stored
|
| 40 |
# student profile under this id. Empty disables profile memory I/O.
|
| 41 |
student_id: str = ""
|
| 42 |
+
# Experiment-only DeepSeek KV-cache namespace. The eval runner generates a
|
| 43 |
+
# stable opaque id per arm/session/trial to prevent cross-arm cache warming.
|
| 44 |
+
cache_user_id: str = ""
|
| 45 |
# Part C / Axis B ablation: drop the run_kb_command tool (and its prompt
|
| 46 |
# section) while keeping retrieval, to measure whether KB browsing helps.
|
| 47 |
disable_kb: bool = False
|
|
@@ -92,9 +92,26 @@ class MemoryConfig:
|
|
| 92 |
summarization: bool = True
|
| 93 |
summarization_trigger_tokens: int = 30_000
|
| 94 |
summarization_keep_messages: int = 20
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 95 |
# Custom SummarizationMiddleware prompt (None = the library default). Used by
|
| 96 |
# the selective_retention / context_reset arms; must template {messages}.
|
| 97 |
summary_prompt: str | None = None
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 98 |
context_editing: bool = True
|
| 99 |
context_editing_trigger_tokens: int = 5_000
|
| 100 |
context_editing_keep: int = 5
|
|
@@ -112,6 +129,19 @@ class MemoryConfig:
|
|
| 112 |
truncate_head_chars: int = 2_000
|
| 113 |
truncate_tail_chars: int = 500
|
| 114 |
truncate_trigger_chars: int = 4_000
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 115 |
compress_prompt: bool = False # deterministic per-call text compaction
|
| 116 |
# In-context history retrieval (Axis A subsystem): keep the last N turn-blocks
|
| 117 |
# and retrieve only the top-k most relevant older blocks. None disables it.
|
|
@@ -132,6 +162,78 @@ MEMORY_PRESETS: dict[str, MemoryConfig] = {
|
|
| 132 |
"full_history": MemoryConfig(
|
| 133 |
name="full_history", summarization=False, context_editing=False
|
| 134 |
),
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 135 |
# What production runs today.
|
| 136 |
"prod": MemoryConfig(name="prod"),
|
| 137 |
"summarization_only": MemoryConfig(
|
|
|
|
| 92 |
summarization: bool = True
|
| 93 |
summarization_trigger_tokens: int = 30_000
|
| 94 |
summarization_keep_messages: int = 20
|
| 95 |
+
# Experiment arms use token-based retention so a single large tool message
|
| 96 |
+
# cannot make the post-compaction window vary by hundreds of thousands of
|
| 97 |
+
# tokens. None preserves the historical message-count behavior.
|
| 98 |
+
summarization_keep_tokens: int | None = None
|
| 99 |
+
# LangChain defaults this to 4k. None deliberately sends the entire selected
|
| 100 |
+
# older history to the summarizer (the corrected long-context experiment).
|
| 101 |
+
summarization_trim_tokens: int | None = 4_000
|
| 102 |
+
# Fail rather than silently trim if a full-input experimental summary would
|
| 103 |
+
# approach the provider's context ceiling. None disables the guard.
|
| 104 |
+
summarization_input_guard_tokens: int | None = None
|
| 105 |
# Custom SummarizationMiddleware prompt (None = the library default). Used by
|
| 106 |
# the selective_retention / context_reset arms; must template {messages}.
|
| 107 |
summary_prompt: str | None = None
|
| 108 |
+
# ``xml`` is LangChain's historical behavior: serialize selected messages
|
| 109 |
+
# into one new prompt string. ``structured_prefix`` keeps the original
|
| 110 |
+
# system message, tool schemas, model settings, and selected message prefix
|
| 111 |
+
# byte-for-byte at the message boundary, then appends one checkpoint
|
| 112 |
+
# instruction. The latter is experiment-only because it changes request
|
| 113 |
+
# shape and checkpoint installation semantics.
|
| 114 |
+
summarization_strategy: str = "xml"
|
| 115 |
context_editing: bool = True
|
| 116 |
context_editing_trigger_tokens: int = 5_000
|
| 117 |
context_editing_keep: int = 5
|
|
|
|
| 129 |
truncate_head_chars: int = 2_000
|
| 130 |
truncate_tail_chars: int = 500
|
| 131 |
truncate_trigger_chars: int = 4_000
|
| 132 |
+
# Persistent insertion-time cap: unlike truncate_tool_outputs, this changes
|
| 133 |
+
# the checkpoint itself, so every later model call and the summarizer see the
|
| 134 |
+
# same stable text. The experiment uses 40k UTF-8 bytes as a nominal 10k-token
|
| 135 |
+
# cap, matching the reproducible approximation used by the Codex harness.
|
| 136 |
+
tool_output_cap_bytes: int | None = None
|
| 137 |
+
# Enables explanatory compaction telemetry and DeepSeek user_id isolation.
|
| 138 |
+
# Kept off for production and historical presets so this study is additive.
|
| 139 |
+
experiment_mode: bool = False
|
| 140 |
+
# Fail before an experimental agent call exceeds this approximate input
|
| 141 |
+
# size. The approximation over-counted the provider-reported input by about
|
| 142 |
+
# 29k near 870k, so 990k preserves real headroom inside DeepSeek's 1M window
|
| 143 |
+
# without prematurely truncating the full-history control.
|
| 144 |
+
experiment_request_guard_tokens: int = 990_000
|
| 145 |
compress_prompt: bool = False # deterministic per-call text compaction
|
| 146 |
# In-context history retrieval (Axis A subsystem): keep the last N turn-blocks
|
| 147 |
# and retrieve only the top-k most relevant older blocks. None disables it.
|
|
|
|
| 162 |
"full_history": MemoryConfig(
|
| 163 |
name="full_history", summarization=False, context_editing=False
|
| 164 |
),
|
| 165 |
+
# --- DeepSeek long-context compaction experiment ----------------------
|
| 166 |
+
# Four mechanism-isolation arms. All disable age-based context editing;
|
| 167 |
+
# the capped arms instead perform one stable rewrite when tool output first
|
| 168 |
+
# enters history. C200 arms summarize the complete selected prefix at 200k
|
| 169 |
+
# and retain a controlled 50k-token recent tail.
|
| 170 |
+
"exp_fh_raw": MemoryConfig(
|
| 171 |
+
name="exp_fh_raw",
|
| 172 |
+
summarization=False,
|
| 173 |
+
context_editing=False,
|
| 174 |
+
experiment_mode=True,
|
| 175 |
+
),
|
| 176 |
+
"exp_fh_cap10k": MemoryConfig(
|
| 177 |
+
name="exp_fh_cap10k",
|
| 178 |
+
summarization=False,
|
| 179 |
+
context_editing=False,
|
| 180 |
+
tool_output_cap_bytes=40_000,
|
| 181 |
+
experiment_mode=True,
|
| 182 |
+
),
|
| 183 |
+
"exp_c200_raw": MemoryConfig(
|
| 184 |
+
name="exp_c200_raw",
|
| 185 |
+
summarization_trigger_tokens=200_000,
|
| 186 |
+
summarization_keep_tokens=50_000,
|
| 187 |
+
summarization_trim_tokens=None,
|
| 188 |
+
summarization_input_guard_tokens=900_000,
|
| 189 |
+
context_editing=False,
|
| 190 |
+
experiment_mode=True,
|
| 191 |
+
),
|
| 192 |
+
"exp_c200_cap10k": MemoryConfig(
|
| 193 |
+
name="exp_c200_cap10k",
|
| 194 |
+
summarization_trigger_tokens=200_000,
|
| 195 |
+
summarization_keep_tokens=50_000,
|
| 196 |
+
summarization_trim_tokens=None,
|
| 197 |
+
summarization_input_guard_tokens=900_000,
|
| 198 |
+
context_editing=False,
|
| 199 |
+
tool_output_cap_bytes=40_000,
|
| 200 |
+
experiment_mode=True,
|
| 201 |
+
),
|
| 202 |
+
# Cache-friendly version of exp_c200_cap10k. It deliberately remains a
|
| 203 |
+
# separate arm so the completed XML run stays reproducible and comparable.
|
| 204 |
+
"exp_c200_cap10k_structured": MemoryConfig(
|
| 205 |
+
name="exp_c200_cap10k_structured",
|
| 206 |
+
summarization_trigger_tokens=200_000,
|
| 207 |
+
summarization_keep_tokens=50_000,
|
| 208 |
+
summarization_trim_tokens=None,
|
| 209 |
+
summarization_input_guard_tokens=900_000,
|
| 210 |
+
summarization_strategy="structured_prefix",
|
| 211 |
+
context_editing=False,
|
| 212 |
+
tool_output_cap_bytes=40_000,
|
| 213 |
+
experiment_mode=True,
|
| 214 |
+
),
|
| 215 |
+
# Stage-2 threshold sensitivity arms share the exact same cap, summary
|
| 216 |
+
# input, and post-compaction retention. Only the trigger changes.
|
| 217 |
+
"exp_c400_cap10k": MemoryConfig(
|
| 218 |
+
name="exp_c400_cap10k",
|
| 219 |
+
summarization_trigger_tokens=400_000,
|
| 220 |
+
summarization_keep_tokens=50_000,
|
| 221 |
+
summarization_trim_tokens=None,
|
| 222 |
+
summarization_input_guard_tokens=900_000,
|
| 223 |
+
context_editing=False,
|
| 224 |
+
tool_output_cap_bytes=40_000,
|
| 225 |
+
experiment_mode=True,
|
| 226 |
+
),
|
| 227 |
+
"exp_c800_cap10k": MemoryConfig(
|
| 228 |
+
name="exp_c800_cap10k",
|
| 229 |
+
summarization_trigger_tokens=800_000,
|
| 230 |
+
summarization_keep_tokens=50_000,
|
| 231 |
+
summarization_trim_tokens=None,
|
| 232 |
+
summarization_input_guard_tokens=900_000,
|
| 233 |
+
context_editing=False,
|
| 234 |
+
tool_output_cap_bytes=40_000,
|
| 235 |
+
experiment_mode=True,
|
| 236 |
+
),
|
| 237 |
# What production runs today.
|
| 238 |
"prod": MemoryConfig(name="prod"),
|
| 239 |
"summarization_only": MemoryConfig(
|
|
@@ -16,13 +16,15 @@ number.
|
|
| 16 |
from __future__ import annotations
|
| 17 |
|
| 18 |
import threading
|
|
|
|
| 19 |
from dataclasses import dataclass
|
| 20 |
from typing import Any
|
|
|
|
| 21 |
|
| 22 |
from langchain_core.callbacks.usage import UsageMetadataCallbackHandler
|
| 23 |
-
from langchain_core.messages import BaseMessage
|
| 24 |
from langchain_core.messages.utils import count_tokens_approximately
|
| 25 |
-
from langchain_core.outputs import LLMResult
|
| 26 |
|
| 27 |
|
| 28 |
# --- Turn-scoped middleware signals ------------------------------------------
|
|
@@ -39,6 +41,7 @@ from langchain_core.outputs import LLMResult
|
|
| 39 |
# wrap_model_call in a worker thread, where a ContextVar update lands on a copy
|
| 40 |
# and is lost. A plain global + lock is visible from any thread.
|
| 41 |
_TURN_SIGNALS: dict[str, dict[str, int]] = {}
|
|
|
|
| 42 |
_TURN_SIGNALS_LOCK = threading.Lock()
|
| 43 |
_MAX_TRACKED_TURNS = 256
|
| 44 |
|
|
@@ -67,8 +70,11 @@ def reset_turn_signals(turn_id: str) -> None:
|
|
| 67 |
# Never clear() the whole dict — that would wipe concurrent in-flight
|
| 68 |
# turns' signals and could mis-grade a probe as "compaction never fired".
|
| 69 |
while len(_TURN_SIGNALS) >= _MAX_TRACKED_TURNS:
|
| 70 |
-
|
|
|
|
|
|
|
| 71 |
_TURN_SIGNALS[turn_id] = {}
|
|
|
|
| 72 |
|
| 73 |
|
| 74 |
def record_turn_signal(turn_id: str, name: str, amount: int = 1) -> None:
|
|
@@ -103,16 +109,107 @@ def pop_turn_signals(turn_id: str) -> dict[str, int]:
|
|
| 103 |
return _TURN_SIGNALS.pop(turn_id, {})
|
| 104 |
|
| 105 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 106 |
class TurnUsageHandler(UsageMetadataCallbackHandler):
|
| 107 |
-
"""Aggregate
|
| 108 |
|
| 109 |
def __init__(self) -> None:
|
| 110 |
super().__init__()
|
| 111 |
self.llm_calls = 0
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 112 |
|
| 113 |
def on_llm_end(self, response: LLMResult, **kwargs: Any) -> None:
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 114 |
with self._lock:
|
| 115 |
self.llm_calls += 1
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 116 |
super().on_llm_end(response, **kwargs)
|
| 117 |
|
| 118 |
|
|
@@ -165,6 +262,52 @@ def pricing_for_model(model_key: str) -> ModelPricing | None:
|
|
| 165 |
return best[1] if best else None
|
| 166 |
|
| 167 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 168 |
def usage_totals(usage_by_model: dict[str, Any]) -> dict[str, int]:
|
| 169 |
"""Sum usage across models into the fields the stats event reports."""
|
| 170 |
totals = {
|
|
@@ -191,27 +334,8 @@ def estimate_cost_usd(usage_by_model: dict[str, Any]) -> float | None:
|
|
| 191 |
includes them (LangChain's UsageMetadata convention), so they are carved
|
| 192 |
out of the plain-input bucket rather than added on top.
|
| 193 |
"""
|
| 194 |
-
|
| 195 |
-
|
| 196 |
-
pricing = pricing_for_model(model_key)
|
| 197 |
-
if pricing is None:
|
| 198 |
-
return None
|
| 199 |
-
input_tokens = int(usage.get("input_tokens", 0) or 0)
|
| 200 |
-
output_tokens = int(usage.get("output_tokens", 0) or 0)
|
| 201 |
-
details = usage.get("input_token_details") or {}
|
| 202 |
-
cache_read = int(details.get("cache_read", 0) or 0)
|
| 203 |
-
cache_creation = int(details.get("cache_creation", 0) or 0)
|
| 204 |
-
plain_input = max(0, input_tokens - cache_read - cache_creation)
|
| 205 |
-
write_rate = (
|
| 206 |
-
pricing.cache_write if pricing.cache_write is not None else pricing.input
|
| 207 |
-
)
|
| 208 |
-
total += (
|
| 209 |
-
plain_input * pricing.input
|
| 210 |
-
+ cache_read * pricing.cache_read
|
| 211 |
-
+ cache_creation * write_rate
|
| 212 |
-
+ output_tokens * pricing.output
|
| 213 |
-
) / 1_000_000
|
| 214 |
-
return total
|
| 215 |
|
| 216 |
|
| 217 |
def context_window_stats(
|
|
|
|
| 16 |
from __future__ import annotations
|
| 17 |
|
| 18 |
import threading
|
| 19 |
+
import time
|
| 20 |
from dataclasses import dataclass
|
| 21 |
from typing import Any
|
| 22 |
+
from uuid import UUID
|
| 23 |
|
| 24 |
from langchain_core.callbacks.usage import UsageMetadataCallbackHandler
|
| 25 |
+
from langchain_core.messages import AIMessage, BaseMessage
|
| 26 |
from langchain_core.messages.utils import count_tokens_approximately
|
| 27 |
+
from langchain_core.outputs import ChatGeneration, LLMResult
|
| 28 |
|
| 29 |
|
| 30 |
# --- Turn-scoped middleware signals ------------------------------------------
|
|
|
|
| 41 |
# wrap_model_call in a worker thread, where a ContextVar update lands on a copy
|
| 42 |
# and is lost. A plain global + lock is visible from any thread.
|
| 43 |
_TURN_SIGNALS: dict[str, dict[str, int]] = {}
|
| 44 |
+
_TURN_EVENTS: dict[str, list[dict[str, Any]]] = {}
|
| 45 |
_TURN_SIGNALS_LOCK = threading.Lock()
|
| 46 |
_MAX_TRACKED_TURNS = 256
|
| 47 |
|
|
|
|
| 70 |
# Never clear() the whole dict — that would wipe concurrent in-flight
|
| 71 |
# turns' signals and could mis-grade a probe as "compaction never fired".
|
| 72 |
while len(_TURN_SIGNALS) >= _MAX_TRACKED_TURNS:
|
| 73 |
+
oldest = next(iter(_TURN_SIGNALS))
|
| 74 |
+
_TURN_SIGNALS.pop(oldest, None)
|
| 75 |
+
_TURN_EVENTS.pop(oldest, None)
|
| 76 |
_TURN_SIGNALS[turn_id] = {}
|
| 77 |
+
_TURN_EVENTS[turn_id] = []
|
| 78 |
|
| 79 |
|
| 80 |
def record_turn_signal(turn_id: str, name: str, amount: int = 1) -> None:
|
|
|
|
| 109 |
return _TURN_SIGNALS.pop(turn_id, {})
|
| 110 |
|
| 111 |
|
| 112 |
+
def record_turn_event(turn_id: str, event: dict[str, Any]) -> None:
|
| 113 |
+
"""Append structured middleware telemetry for one turn (thread-safe)."""
|
| 114 |
+
if not turn_id:
|
| 115 |
+
return
|
| 116 |
+
with _TURN_SIGNALS_LOCK:
|
| 117 |
+
_TURN_EVENTS.setdefault(turn_id, []).append(dict(event))
|
| 118 |
+
|
| 119 |
+
|
| 120 |
+
def pop_turn_events(turn_id: str) -> list[dict[str, Any]]:
|
| 121 |
+
"""Take and clear a turn's structured middleware events."""
|
| 122 |
+
if not turn_id:
|
| 123 |
+
return []
|
| 124 |
+
with _TURN_SIGNALS_LOCK:
|
| 125 |
+
return _TURN_EVENTS.pop(turn_id, [])
|
| 126 |
+
|
| 127 |
+
|
| 128 |
class TurnUsageHandler(UsageMetadataCallbackHandler):
|
| 129 |
+
"""Aggregate usage and retain one explanatory record per model call."""
|
| 130 |
|
| 131 |
def __init__(self) -> None:
|
| 132 |
super().__init__()
|
| 133 |
self.llm_calls = 0
|
| 134 |
+
self.model_calls: list[dict[str, Any]] = []
|
| 135 |
+
self._call_starts: dict[UUID, dict[str, Any]] = {}
|
| 136 |
+
|
| 137 |
+
def on_chat_model_start(
|
| 138 |
+
self,
|
| 139 |
+
serialized: dict[str, Any],
|
| 140 |
+
messages: list[list[BaseMessage]],
|
| 141 |
+
*,
|
| 142 |
+
run_id: UUID,
|
| 143 |
+
parent_run_id: UUID | None = None,
|
| 144 |
+
tags: list[str] | None = None,
|
| 145 |
+
metadata: dict[str, Any] | None = None,
|
| 146 |
+
**kwargs: Any,
|
| 147 |
+
) -> None:
|
| 148 |
+
del serialized, parent_run_id, tags, kwargs
|
| 149 |
+
request_messages = messages[0] if messages else []
|
| 150 |
+
try:
|
| 151 |
+
request_tokens = int(count_tokens_approximately(request_messages))
|
| 152 |
+
except Exception: # telemetry must never break a model call
|
| 153 |
+
request_tokens = 0
|
| 154 |
+
call_metadata = metadata or {}
|
| 155 |
+
with self._lock:
|
| 156 |
+
self._call_starts[run_id] = {
|
| 157 |
+
"started_at": time.monotonic(),
|
| 158 |
+
"source": str(call_metadata.get("lc_source") or "agent"),
|
| 159 |
+
"request_context_tokens_approx": request_tokens,
|
| 160 |
+
}
|
| 161 |
|
| 162 |
def on_llm_end(self, response: LLMResult, **kwargs: Any) -> None:
|
| 163 |
+
run_id = kwargs.get("run_id")
|
| 164 |
+
generation = None
|
| 165 |
+
try:
|
| 166 |
+
generation = response.generations[0][0]
|
| 167 |
+
except IndexError:
|
| 168 |
+
pass
|
| 169 |
+
|
| 170 |
+
message = generation.message if isinstance(generation, ChatGeneration) else None
|
| 171 |
+
usage = (
|
| 172 |
+
dict(message.usage_metadata or {}) if isinstance(message, AIMessage) else {}
|
| 173 |
+
)
|
| 174 |
+
model_name = (
|
| 175 |
+
str(message.response_metadata.get("model_name") or "")
|
| 176 |
+
if isinstance(message, AIMessage)
|
| 177 |
+
else ""
|
| 178 |
+
)
|
| 179 |
+
details = dict(usage.get("input_token_details") or {})
|
| 180 |
+
breakdown = usage_cost_breakdown(model_name, usage) if model_name else None
|
| 181 |
with self._lock:
|
| 182 |
self.llm_calls += 1
|
| 183 |
+
start = self._call_starts.pop(run_id, {}) if run_id else {}
|
| 184 |
+
self.model_calls.append(
|
| 185 |
+
{
|
| 186 |
+
"sequence": self.llm_calls,
|
| 187 |
+
"source": start.get("source", "agent"),
|
| 188 |
+
"model": model_name,
|
| 189 |
+
"input_tokens": int(usage.get("input_tokens", 0) or 0),
|
| 190 |
+
"cache_read_tokens": int(details.get("cache_read", 0) or 0),
|
| 191 |
+
"cache_miss_tokens": max(
|
| 192 |
+
0,
|
| 193 |
+
int(usage.get("input_tokens", 0) or 0)
|
| 194 |
+
- int(details.get("cache_read", 0) or 0)
|
| 195 |
+
- int(details.get("cache_creation", 0) or 0),
|
| 196 |
+
),
|
| 197 |
+
"cache_creation_tokens": int(details.get("cache_creation", 0) or 0),
|
| 198 |
+
"output_tokens": int(usage.get("output_tokens", 0) or 0),
|
| 199 |
+
"total_tokens": int(usage.get("total_tokens", 0) or 0),
|
| 200 |
+
"usage_reported": bool(usage),
|
| 201 |
+
"cache_details_reported": "cache_read" in details,
|
| 202 |
+
"request_context_tokens_approx": int(
|
| 203 |
+
start.get("request_context_tokens_approx", 0) or 0
|
| 204 |
+
),
|
| 205 |
+
"duration_ms": (
|
| 206 |
+
int((time.monotonic() - start["started_at"]) * 1000)
|
| 207 |
+
if start.get("started_at") is not None
|
| 208 |
+
else None
|
| 209 |
+
),
|
| 210 |
+
"cost": breakdown,
|
| 211 |
+
}
|
| 212 |
+
)
|
| 213 |
super().on_llm_end(response, **kwargs)
|
| 214 |
|
| 215 |
|
|
|
|
| 262 |
return best[1] if best else None
|
| 263 |
|
| 264 |
|
| 265 |
+
def usage_cost_breakdown(
|
| 266 |
+
model_key: str, usage: dict[str, Any]
|
| 267 |
+
) -> dict[str, float] | None:
|
| 268 |
+
"""Return mutually exclusive USD components for one model's usage."""
|
| 269 |
+
pricing = pricing_for_model(model_key)
|
| 270 |
+
if pricing is None:
|
| 271 |
+
return None
|
| 272 |
+
input_tokens = int(usage.get("input_tokens", 0) or 0)
|
| 273 |
+
output_tokens = int(usage.get("output_tokens", 0) or 0)
|
| 274 |
+
details = usage.get("input_token_details") or {}
|
| 275 |
+
cache_read = int(details.get("cache_read", 0) or 0)
|
| 276 |
+
cache_creation = int(details.get("cache_creation", 0) or 0)
|
| 277 |
+
plain_input = max(0, input_tokens - cache_read - cache_creation)
|
| 278 |
+
write_rate = (
|
| 279 |
+
pricing.cache_write if pricing.cache_write is not None else pricing.input
|
| 280 |
+
)
|
| 281 |
+
result = {
|
| 282 |
+
"cache_miss_input_usd": plain_input * pricing.input / 1_000_000,
|
| 283 |
+
"cache_read_input_usd": cache_read * pricing.cache_read / 1_000_000,
|
| 284 |
+
"cache_creation_input_usd": cache_creation * write_rate / 1_000_000,
|
| 285 |
+
"output_usd": output_tokens * pricing.output / 1_000_000,
|
| 286 |
+
}
|
| 287 |
+
result["total_usd"] = sum(result.values())
|
| 288 |
+
return result
|
| 289 |
+
|
| 290 |
+
|
| 291 |
+
def aggregate_cost_breakdown(
|
| 292 |
+
usage_by_model: dict[str, Any],
|
| 293 |
+
) -> dict[str, float] | None:
|
| 294 |
+
"""Sum cost components across every model used in a turn."""
|
| 295 |
+
totals = {
|
| 296 |
+
"cache_miss_input_usd": 0.0,
|
| 297 |
+
"cache_read_input_usd": 0.0,
|
| 298 |
+
"cache_creation_input_usd": 0.0,
|
| 299 |
+
"output_usd": 0.0,
|
| 300 |
+
"total_usd": 0.0,
|
| 301 |
+
}
|
| 302 |
+
for model_key, usage in usage_by_model.items():
|
| 303 |
+
breakdown = usage_cost_breakdown(model_key, usage)
|
| 304 |
+
if breakdown is None:
|
| 305 |
+
return None
|
| 306 |
+
for key in totals:
|
| 307 |
+
totals[key] += breakdown[key]
|
| 308 |
+
return totals
|
| 309 |
+
|
| 310 |
+
|
| 311 |
def usage_totals(usage_by_model: dict[str, Any]) -> dict[str, int]:
|
| 312 |
"""Sum usage across models into the fields the stats event reports."""
|
| 313 |
totals = {
|
|
|
|
| 334 |
includes them (LangChain's UsageMetadata convention), so they are carved
|
| 335 |
out of the plain-input bucket rather than added on top.
|
| 336 |
"""
|
| 337 |
+
breakdown = aggregate_cost_breakdown(usage_by_model)
|
| 338 |
+
return breakdown["total_usd"] if breakdown is not None else None
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 339 |
|
| 340 |
|
| 341 |
def context_window_stats(
|
|
@@ -32,6 +32,7 @@ from app.chat_service import (
|
|
| 32 |
extract_shell_source_matches,
|
| 33 |
resolve_answer_citations,
|
| 34 |
retrieve_tutor_context,
|
|
|
|
| 35 |
supports_gemini_tool_combination,
|
| 36 |
sync_thread_with_history,
|
| 37 |
stream_chat,
|
|
@@ -466,6 +467,59 @@ class ChatServiceTestCase(unittest.TestCase):
|
|
| 466 |
finally:
|
| 467 |
_clear_kb_command_budget(session_id)
|
| 468 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 469 |
def test_retrieve_tutor_context_degrades_on_retriever_failure(self) -> None:
|
| 470 |
runtime = types.SimpleNamespace(
|
| 471 |
context=types.SimpleNamespace(allowed_sources=("transformers",))
|
|
@@ -480,6 +534,21 @@ class ChatServiceTestCase(unittest.TestCase):
|
|
| 480 |
# ...and the raw provider error is not exposed in the tool output.
|
| 481 |
self.assertNotIn("cohere 500 boom", result)
|
| 482 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 483 |
def test_resolve_answer_citations_uses_current_turn_evidence(self) -> None:
|
| 484 |
retrieval = SourceMatch(
|
| 485 |
doc_id="peft:lora",
|
|
|
|
| 32 |
extract_shell_source_matches,
|
| 33 |
resolve_answer_citations,
|
| 34 |
retrieve_tutor_context,
|
| 35 |
+
run_kb_command,
|
| 36 |
supports_gemini_tool_combination,
|
| 37 |
sync_thread_with_history,
|
| 38 |
stream_chat,
|
|
|
|
| 467 |
finally:
|
| 468 |
_clear_kb_command_budget(session_id)
|
| 469 |
|
| 470 |
+
def test_kb_command_accepts_timeout_alias_with_runtime_cap(self) -> None:
|
| 471 |
+
session_id = "test_timeout_alias"
|
| 472 |
+
_clear_kb_command_budget(session_id)
|
| 473 |
+
runtime = types.SimpleNamespace(
|
| 474 |
+
context=types.SimpleNamespace(
|
| 475 |
+
kb_session_id=session_id,
|
| 476 |
+
kb_command_limit=3,
|
| 477 |
+
)
|
| 478 |
+
)
|
| 479 |
+
try:
|
| 480 |
+
with (
|
| 481 |
+
patch("app.chat_service.ensure_local_vector_db"),
|
| 482 |
+
patch("app.chat_service.execute_kb_command") as execute,
|
| 483 |
+
patch("app.chat_service.format_command_payload", return_value="$ ls"),
|
| 484 |
+
):
|
| 485 |
+
result = run_kb_command.func(
|
| 486 |
+
command="ls",
|
| 487 |
+
runtime=runtime,
|
| 488 |
+
timeout=5,
|
| 489 |
+
)
|
| 490 |
+
finally:
|
| 491 |
+
_clear_kb_command_budget(session_id)
|
| 492 |
+
self.assertEqual(result, "$ ls")
|
| 493 |
+
execute.assert_called_once_with(
|
| 494 |
+
"ls",
|
| 495 |
+
timeout_seconds=5,
|
| 496 |
+
max_output_chars=40000,
|
| 497 |
+
)
|
| 498 |
+
|
| 499 |
+
def test_kb_command_unknown_arguments_return_soft_tool_error(self) -> None:
|
| 500 |
+
runtime = types.SimpleNamespace(
|
| 501 |
+
context=types.SimpleNamespace(
|
| 502 |
+
kb_session_id="test_unknown_tool_arg",
|
| 503 |
+
kb_command_limit=3,
|
| 504 |
+
)
|
| 505 |
+
)
|
| 506 |
+
with patch("app.chat_service.execute_kb_command") as execute:
|
| 507 |
+
result = run_kb_command.func(
|
| 508 |
+
command="ls",
|
| 509 |
+
runtime=runtime,
|
| 510 |
+
working_directory="raw",
|
| 511 |
+
)
|
| 512 |
+
self.assertIn("unsupported run_kb_command argument", result)
|
| 513 |
+
self.assertIn("working_directory", result)
|
| 514 |
+
execute.assert_not_called()
|
| 515 |
+
|
| 516 |
+
def test_kb_command_schema_disallows_unpublished_arguments(self) -> None:
|
| 517 |
+
schema = run_kb_command.args_schema
|
| 518 |
+
self.assertIsInstance(schema, dict)
|
| 519 |
+
self.assertFalse(schema["additionalProperties"])
|
| 520 |
+
self.assertIn("timeout", schema["properties"])
|
| 521 |
+
self.assertEqual(schema["properties"]["timeout"]["maximum"], 30)
|
| 522 |
+
|
| 523 |
def test_retrieve_tutor_context_degrades_on_retriever_failure(self) -> None:
|
| 524 |
runtime = types.SimpleNamespace(
|
| 525 |
context=types.SimpleNamespace(allowed_sources=("transformers",))
|
|
|
|
| 534 |
# ...and the raw provider error is not exposed in the tool output.
|
| 535 |
self.assertNotIn("cohere 500 boom", result)
|
| 536 |
|
| 537 |
+
def test_retrieve_tutor_context_unknown_arguments_return_soft_error(self) -> None:
|
| 538 |
+
runtime = types.SimpleNamespace(
|
| 539 |
+
context=types.SimpleNamespace(allowed_sources=("transformers",))
|
| 540 |
+
)
|
| 541 |
+
with patch("app.chat_service.select_retriever") as select:
|
| 542 |
+
result = retrieve_tutor_context.func(
|
| 543 |
+
query="What is RAG?",
|
| 544 |
+
runtime=runtime,
|
| 545 |
+
top_k=20,
|
| 546 |
+
)
|
| 547 |
+
self.assertIn("unsupported argument", result)
|
| 548 |
+
self.assertIn("top_k", result)
|
| 549 |
+
select.assert_not_called()
|
| 550 |
+
self.assertFalse(retrieve_tutor_context.args_schema["additionalProperties"])
|
| 551 |
+
|
| 552 |
def test_resolve_answer_citations_uses_current_turn_evidence(self) -> None:
|
| 553 |
retrieval = SourceMatch(
|
| 554 |
doc_id="peft:lora",
|
|
@@ -5,7 +5,11 @@ import unittest
|
|
| 5 |
from unittest.mock import patch
|
| 6 |
|
| 7 |
from app.chat_service import (
|
|
|
|
|
|
|
|
|
|
| 8 |
SourcePreferenceMiddleware,
|
|
|
|
| 9 |
StudentProfileMiddleware,
|
| 10 |
build_agent,
|
| 11 |
build_agent_middleware,
|
|
@@ -58,6 +62,26 @@ class MemoryPresetResolutionTests(unittest.TestCase):
|
|
| 58 |
with self.assertRaises(ValueError):
|
| 59 |
resolve_memory_preset("")
|
| 60 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 61 |
|
| 62 |
class MiddlewareAssemblyTests(unittest.TestCase):
|
| 63 |
def test_full_history_disables_compaction(self) -> None:
|
|
@@ -102,6 +126,40 @@ class MiddlewareAssemblyTests(unittest.TestCase):
|
|
| 102 |
any(isinstance(m, StudentProfileMiddleware) for m in middleware)
|
| 103 |
)
|
| 104 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 105 |
def test_build_agent_cache_keys_include_memory_config(self) -> None:
|
| 106 |
build_agent.cache_clear()
|
| 107 |
created = []
|
|
|
|
| 5 |
from unittest.mock import patch
|
| 6 |
|
| 7 |
from app.chat_service import (
|
| 8 |
+
DeepSeekCacheIsolationMiddleware,
|
| 9 |
+
InstrumentedSummarizationMiddleware,
|
| 10 |
+
PrefixPreservingCompactionMiddleware,
|
| 11 |
SourcePreferenceMiddleware,
|
| 12 |
+
StableToolOutputCapMiddleware,
|
| 13 |
StudentProfileMiddleware,
|
| 14 |
build_agent,
|
| 15 |
build_agent_middleware,
|
|
|
|
| 62 |
with self.assertRaises(ValueError):
|
| 63 |
resolve_memory_preset("")
|
| 64 |
|
| 65 |
+
def test_deepseek_experiment_arms_are_single_axis_configs(self) -> None:
|
| 66 |
+
raw = MEMORY_PRESETS["exp_fh_raw"]
|
| 67 |
+
capped = MEMORY_PRESETS["exp_fh_cap10k"]
|
| 68 |
+
compact = MEMORY_PRESETS["exp_c200_cap10k"]
|
| 69 |
+
self.assertTrue(raw.experiment_mode)
|
| 70 |
+
self.assertFalse(raw.summarization)
|
| 71 |
+
self.assertFalse(raw.context_editing)
|
| 72 |
+
self.assertEqual(capped.tool_output_cap_bytes, 40_000)
|
| 73 |
+
self.assertEqual(compact.summarization_trigger_tokens, 200_000)
|
| 74 |
+
self.assertEqual(compact.summarization_keep_tokens, 50_000)
|
| 75 |
+
self.assertIsNone(compact.summarization_trim_tokens)
|
| 76 |
+
self.assertEqual(compact.summarization_input_guard_tokens, 900_000)
|
| 77 |
+
self.assertEqual(compact.experiment_request_guard_tokens, 990_000)
|
| 78 |
+
self.assertFalse(compact.context_editing)
|
| 79 |
+
structured = MEMORY_PRESETS["exp_c200_cap10k_structured"]
|
| 80 |
+
self.assertEqual(structured.summarization_strategy, "structured_prefix")
|
| 81 |
+
self.assertEqual(structured.summarization_trigger_tokens, 200_000)
|
| 82 |
+
self.assertEqual(structured.summarization_keep_tokens, 50_000)
|
| 83 |
+
self.assertEqual(structured.tool_output_cap_bytes, 40_000)
|
| 84 |
+
|
| 85 |
|
| 86 |
class MiddlewareAssemblyTests(unittest.TestCase):
|
| 87 |
def test_full_history_disables_compaction(self) -> None:
|
|
|
|
| 126 |
any(isinstance(m, StudentProfileMiddleware) for m in middleware)
|
| 127 |
)
|
| 128 |
|
| 129 |
+
def test_experiment_stack_has_isolation_cap_and_full_input_summary(self) -> None:
|
| 130 |
+
middleware = build_agent_middleware(
|
| 131 |
+
_stub_model(), MEMORY_PRESETS["exp_c200_cap10k"]
|
| 132 |
+
)
|
| 133 |
+
self.assertTrue(
|
| 134 |
+
any(isinstance(m, DeepSeekCacheIsolationMiddleware) for m in middleware)
|
| 135 |
+
)
|
| 136 |
+
self.assertTrue(
|
| 137 |
+
any(isinstance(m, StableToolOutputCapMiddleware) for m in middleware)
|
| 138 |
+
)
|
| 139 |
+
summary = next(
|
| 140 |
+
m for m in middleware if isinstance(m, InstrumentedSummarizationMiddleware)
|
| 141 |
+
)
|
| 142 |
+
self.assertEqual(summary.keep, ("tokens", 50_000))
|
| 143 |
+
self.assertIsNone(summary.trim_tokens_to_summarize)
|
| 144 |
+
self.assertEqual(summary.summary_input_guard_tokens, 900_000)
|
| 145 |
+
self.assertFalse(
|
| 146 |
+
any(isinstance(m, ContextEditingMiddleware) for m in middleware)
|
| 147 |
+
)
|
| 148 |
+
|
| 149 |
+
def test_structured_compactor_runs_after_request_shaping_middleware(self) -> None:
|
| 150 |
+
middleware = build_agent_middleware(
|
| 151 |
+
_stub_model(), MEMORY_PRESETS["exp_c200_cap10k_structured"]
|
| 152 |
+
)
|
| 153 |
+
self.assertIsInstance(middleware[0], DeepSeekCacheIsolationMiddleware)
|
| 154 |
+
self.assertTrue(
|
| 155 |
+
any(isinstance(m, StableToolOutputCapMiddleware) for m in middleware)
|
| 156 |
+
)
|
| 157 |
+
self.assertIsInstance(middleware[-2], SourcePreferenceMiddleware)
|
| 158 |
+
self.assertIsInstance(middleware[-1], PrefixPreservingCompactionMiddleware)
|
| 159 |
+
self.assertFalse(
|
| 160 |
+
any(type(m) is InstrumentedSummarizationMiddleware for m in middleware)
|
| 161 |
+
)
|
| 162 |
+
|
| 163 |
def test_build_agent_cache_keys_include_memory_config(self) -> None:
|
| 164 |
build_agent.cache_clear()
|
| 165 |
created = []
|
|
@@ -8,22 +8,49 @@ no model client, no API keys, no vector DB.
|
|
| 8 |
|
| 9 |
from __future__ import annotations
|
| 10 |
|
|
|
|
| 11 |
import unittest
|
| 12 |
from types import SimpleNamespace
|
|
|
|
| 13 |
|
| 14 |
import tiktoken
|
| 15 |
-
from
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 16 |
|
| 17 |
from app.chat_service import (
|
|
|
|
|
|
|
| 18 |
InContextHistoryRetrievalMiddleware,
|
|
|
|
| 19 |
ObservationTruncationMiddleware,
|
| 20 |
PromptCompressionMiddleware,
|
|
|
|
| 21 |
SlidingWindowMiddleware,
|
|
|
|
| 22 |
build_agent_middleware,
|
| 23 |
)
|
| 24 |
from app.chroma_rag import LocalChromaRetriever
|
| 25 |
from app.memory_presets import resolve_memory_preset
|
| 26 |
-
from app.telemetry import
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 27 |
from evals.common import COMPACTION_SIGNAL_KEYS, compaction_active
|
| 28 |
|
| 29 |
|
|
@@ -140,6 +167,1171 @@ class ObservationTruncationTests(unittest.TestCase):
|
|
| 140 |
self.assertEqual(pop_turn_signals("t1"), {})
|
| 141 |
|
| 142 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 143 |
class PromptCompressionTests(unittest.TestCase):
|
| 144 |
def test_collapses_whitespace_and_reports(self) -> None:
|
| 145 |
reset_turn_signals("t1")
|
|
|
|
| 8 |
|
| 9 |
from __future__ import annotations
|
| 10 |
|
| 11 |
+
import hashlib
|
| 12 |
import unittest
|
| 13 |
from types import SimpleNamespace
|
| 14 |
+
from unittest import mock
|
| 15 |
|
| 16 |
import tiktoken
|
| 17 |
+
from langchain.agents import create_agent
|
| 18 |
+
from langchain.agents.middleware import ModelRequest, ModelResponse
|
| 19 |
+
from langchain.tools import tool
|
| 20 |
+
from langchain_core.language_models.fake_chat_models import FakeMessagesListChatModel
|
| 21 |
+
from langchain_core.messages import (
|
| 22 |
+
AIMessage,
|
| 23 |
+
HumanMessage,
|
| 24 |
+
RemoveMessage,
|
| 25 |
+
SystemMessage,
|
| 26 |
+
ToolMessage,
|
| 27 |
+
)
|
| 28 |
+
from langgraph.checkpoint.memory import InMemorySaver
|
| 29 |
+
from langgraph.graph.message import REMOVE_ALL_MESSAGES
|
| 30 |
|
| 31 |
from app.chat_service import (
|
| 32 |
+
AppContext,
|
| 33 |
+
DeepSeekCacheIsolationMiddleware,
|
| 34 |
InContextHistoryRetrievalMiddleware,
|
| 35 |
+
InstrumentedSummarizationMiddleware,
|
| 36 |
ObservationTruncationMiddleware,
|
| 37 |
PromptCompressionMiddleware,
|
| 38 |
+
PrefixPreservingCompactionMiddleware,
|
| 39 |
SlidingWindowMiddleware,
|
| 40 |
+
StableToolOutputCapMiddleware,
|
| 41 |
build_agent_middleware,
|
| 42 |
)
|
| 43 |
from app.chroma_rag import LocalChromaRetriever
|
| 44 |
from app.memory_presets import resolve_memory_preset
|
| 45 |
+
from app.telemetry import (
|
| 46 |
+
COMPACTION_SIGNAL_NAMES,
|
| 47 |
+
TurnUsageHandler,
|
| 48 |
+
estimate_cost_usd,
|
| 49 |
+
pop_turn_events,
|
| 50 |
+
pop_turn_signals,
|
| 51 |
+
reset_turn_signals,
|
| 52 |
+
usage_totals,
|
| 53 |
+
)
|
| 54 |
from evals.common import COMPACTION_SIGNAL_KEYS, compaction_active
|
| 55 |
|
| 56 |
|
|
|
|
| 167 |
self.assertEqual(pop_turn_signals("t1"), {})
|
| 168 |
|
| 169 |
|
| 170 |
+
OVERSIZED_TOOL_OUTPUT = "HEAD" + "x" * 20_000 + "TAIL"
|
| 171 |
+
|
| 172 |
+
|
| 173 |
+
@tool
|
| 174 |
+
def big_lookup(query: str) -> str:
|
| 175 |
+
"""Return deliberately oversized evidence."""
|
| 176 |
+
del query
|
| 177 |
+
return OVERSIZED_TOOL_OUTPUT
|
| 178 |
+
|
| 179 |
+
|
| 180 |
+
class ToolCallingFakeModel(FakeMessagesListChatModel):
|
| 181 |
+
"""Scripted model that accepts tool binding (the base class raises)."""
|
| 182 |
+
|
| 183 |
+
def bind_tools(self, tools, **kwargs):
|
| 184 |
+
return self
|
| 185 |
+
|
| 186 |
+
|
| 187 |
+
class StableToolOutputCapTests(unittest.TestCase):
|
| 188 |
+
def test_cap_persists_to_checkpoint_and_summarizer_never_sees_raw(self) -> None:
|
| 189 |
+
model = ToolCallingFakeModel(
|
| 190 |
+
responses=[
|
| 191 |
+
AIMessage(
|
| 192 |
+
content="",
|
| 193 |
+
tool_calls=[
|
| 194 |
+
{
|
| 195 |
+
"name": "big_lookup",
|
| 196 |
+
"args": {"query": "q"},
|
| 197 |
+
"id": "call-big",
|
| 198 |
+
}
|
| 199 |
+
],
|
| 200 |
+
),
|
| 201 |
+
AIMessage(content="final answer"),
|
| 202 |
+
]
|
| 203 |
+
)
|
| 204 |
+
agent = create_agent(
|
| 205 |
+
model=model,
|
| 206 |
+
tools=[big_lookup],
|
| 207 |
+
middleware=[StableToolOutputCapMiddleware(2_048)],
|
| 208 |
+
checkpointer=InMemorySaver(),
|
| 209 |
+
)
|
| 210 |
+
config = {"configurable": {"thread_id": "cap-thread"}}
|
| 211 |
+
reset_turn_signals("cap-e2e")
|
| 212 |
+
agent.invoke(
|
| 213 |
+
{"messages": [HumanMessage(content="look this up")]},
|
| 214 |
+
config=config,
|
| 215 |
+
context=AppContext(allowed_sources=(), kb_session_id="cap-e2e"),
|
| 216 |
+
)
|
| 217 |
+
|
| 218 |
+
checkpointed = agent.get_state(config).values["messages"]
|
| 219 |
+
tool_messages = [m for m in checkpointed if isinstance(m, ToolMessage)]
|
| 220 |
+
self.assertEqual(len(tool_messages), 1)
|
| 221 |
+
capped = tool_messages[0]
|
| 222 |
+
self.assertLessEqual(len(capped.content.encode("utf-8")), 2_048)
|
| 223 |
+
self.assertIn("truncated at stable 2048-byte cap", capped.content)
|
| 224 |
+
self.assertNotIn(OVERSIZED_TOOL_OUTPUT, capped.content)
|
| 225 |
+
metadata = capped.additional_kwargs["stable_tool_cap"]
|
| 226 |
+
self.assertEqual(
|
| 227 |
+
metadata["sha256"],
|
| 228 |
+
hashlib.sha256(OVERSIZED_TOOL_OUTPUT.encode("utf-8")).hexdigest(),
|
| 229 |
+
)
|
| 230 |
+
self.assertEqual(
|
| 231 |
+
metadata["original_bytes"], len(OVERSIZED_TOOL_OUTPUT.encode("utf-8"))
|
| 232 |
+
)
|
| 233 |
+
self.assertEqual(pop_turn_signals("cap-e2e")["tool_outputs_capped"], 1)
|
| 234 |
+
|
| 235 |
+
# XML path: the summarizer's prompt is built from the checkpointed
|
| 236 |
+
# (capped) ToolMessage, never the raw oversized output.
|
| 237 |
+
summarizer = InstrumentedSummarizationMiddleware(
|
| 238 |
+
model=ExperimentCompactionMiddlewareTests.FakeModel(),
|
| 239 |
+
trigger=("tokens", 100),
|
| 240 |
+
keep=("tokens", 30),
|
| 241 |
+
trim_tokens_to_summarize=None,
|
| 242 |
+
)
|
| 243 |
+
plan = summarizer._plan_compaction({"messages": checkpointed})
|
| 244 |
+
self.assertIsNotNone(plan)
|
| 245 |
+
planned_tool = next(m for m in plan["trimmed"] if isinstance(m, ToolMessage))
|
| 246 |
+
self.assertEqual(planned_tool.content, capped.content)
|
| 247 |
+
prompt = summarizer._summary_prompt_text(plan["trimmed"])
|
| 248 |
+
self.assertIn("truncated at stable 2048-byte cap", prompt)
|
| 249 |
+
self.assertNotIn(OVERSIZED_TOOL_OUTPUT, prompt)
|
| 250 |
+
|
| 251 |
+
# Structured path: the summary request extends the same checkpointed
|
| 252 |
+
# prefix, so it carries the identical capped ToolMessage.
|
| 253 |
+
structured = PrefixPreservingCompactionMiddleware(
|
| 254 |
+
model=ExperimentCompactionMiddlewareTests.FakeModel(),
|
| 255 |
+
trigger=("tokens", 100),
|
| 256 |
+
keep=("tokens", 30),
|
| 257 |
+
trim_tokens_to_summarize=None,
|
| 258 |
+
)
|
| 259 |
+
request = ModelRequest(
|
| 260 |
+
model=ExperimentCompactionMiddlewareTests.FakeModel(),
|
| 261 |
+
messages=list(checkpointed),
|
| 262 |
+
system_message=SystemMessage(content="system"),
|
| 263 |
+
tools=[],
|
| 264 |
+
tool_choice=None,
|
| 265 |
+
response_format=None,
|
| 266 |
+
model_settings={},
|
| 267 |
+
state={"messages": list(checkpointed)},
|
| 268 |
+
runtime=SimpleNamespace(
|
| 269 |
+
context=SimpleNamespace(kb_session_id="cap-e2e", cache_user_id="")
|
| 270 |
+
),
|
| 271 |
+
)
|
| 272 |
+
structured_plan = structured._plan_compaction(request.state)
|
| 273 |
+
self.assertIsNotNone(structured_plan)
|
| 274 |
+
_, summary_request_messages = structured._prepare_summary_request(
|
| 275 |
+
request, structured_plan
|
| 276 |
+
)
|
| 277 |
+
request_tools = [
|
| 278 |
+
m for m in summary_request_messages if isinstance(m, ToolMessage)
|
| 279 |
+
]
|
| 280 |
+
self.assertEqual([m.content for m in request_tools], [capped.content])
|
| 281 |
+
|
| 282 |
+
def test_cap_is_persistent_bounded_and_auditable(self) -> None:
|
| 283 |
+
raw = "HEAD" + ("é" * 30_000) + "TAIL"
|
| 284 |
+
request = make_request([], "stable-cap")
|
| 285 |
+
reset_turn_signals("stable-cap")
|
| 286 |
+
result = StableToolOutputCapMiddleware(40_000)._cap(
|
| 287 |
+
request,
|
| 288 |
+
ToolMessage(content=raw, tool_call_id="call-cap"),
|
| 289 |
+
)
|
| 290 |
+
self.assertLessEqual(len(result.content.encode("utf-8")), 40_000)
|
| 291 |
+
self.assertTrue(result.content.startswith("HEAD"))
|
| 292 |
+
self.assertTrue(result.content.endswith("TAIL"))
|
| 293 |
+
metadata = result.additional_kwargs["stable_tool_cap"]
|
| 294 |
+
self.assertEqual(metadata["original_bytes"], len(raw.encode("utf-8")))
|
| 295 |
+
self.assertEqual(len(metadata["sha256"]), 64)
|
| 296 |
+
signals = pop_turn_signals("stable-cap")
|
| 297 |
+
self.assertEqual(signals["tool_outputs_capped"], 1)
|
| 298 |
+
self.assertGreater(
|
| 299 |
+
signals["tool_output_original_bytes"],
|
| 300 |
+
signals["tool_output_retained_bytes"],
|
| 301 |
+
)
|
| 302 |
+
|
| 303 |
+
|
| 304 |
+
class ExperimentCompactionMiddlewareTests(unittest.TestCase):
|
| 305 |
+
class FakeModel:
|
| 306 |
+
_llm_type = "fake-chat-model"
|
| 307 |
+
|
| 308 |
+
def __init__(self, responses: list[str] | None = None) -> None:
|
| 309 |
+
self.bound: list[dict] = []
|
| 310 |
+
self.prompts: list[str] = []
|
| 311 |
+
self.responses = list(responses or ["durable full-input summary"])
|
| 312 |
+
|
| 313 |
+
def bind(self, **kwargs):
|
| 314 |
+
self.bound.append(kwargs)
|
| 315 |
+
return self
|
| 316 |
+
|
| 317 |
+
def invoke(self, prompt, config=None):
|
| 318 |
+
self.prompts.append(prompt)
|
| 319 |
+
return AIMessage(content=self.responses.pop(0))
|
| 320 |
+
|
| 321 |
+
async def ainvoke(self, prompt, config=None):
|
| 322 |
+
return self.invoke(prompt, config=config)
|
| 323 |
+
|
| 324 |
+
def _get_ls_params(self):
|
| 325 |
+
return {"ls_provider": "deepseek"}
|
| 326 |
+
|
| 327 |
+
class StructuredFakeModel(FakeModel):
|
| 328 |
+
def __init__(self, responses: list[str] | None = None) -> None:
|
| 329 |
+
super().__init__(responses)
|
| 330 |
+
self.bound_tools: list[tuple[list, dict]] = []
|
| 331 |
+
self.invocations: list[tuple[list, dict | None]] = []
|
| 332 |
+
|
| 333 |
+
def bind_tools(self, tools, **kwargs):
|
| 334 |
+
self.bound_tools.append((list(tools), dict(kwargs)))
|
| 335 |
+
return self
|
| 336 |
+
|
| 337 |
+
def invoke(self, prompt, config=None):
|
| 338 |
+
self.invocations.append((list(prompt), config))
|
| 339 |
+
return AIMessage(
|
| 340 |
+
content=self.responses.pop(0),
|
| 341 |
+
usage_metadata={
|
| 342 |
+
"input_tokens": 10_000,
|
| 343 |
+
"output_tokens": 100,
|
| 344 |
+
"total_tokens": 10_100,
|
| 345 |
+
"input_token_details": {"cache_read": 9_000},
|
| 346 |
+
},
|
| 347 |
+
response_metadata={"model_name": "deepseek-v4-flash"},
|
| 348 |
+
)
|
| 349 |
+
|
| 350 |
+
class ScriptedMessageModel(FakeModel):
|
| 351 |
+
"""Structured-path fake returning prebuilt AIMessage responses verbatim."""
|
| 352 |
+
|
| 353 |
+
def __init__(self, responses: list[AIMessage]) -> None:
|
| 354 |
+
super().__init__([])
|
| 355 |
+
self.message_responses = list(responses)
|
| 356 |
+
self.invocations: list[tuple[list, dict | None]] = []
|
| 357 |
+
|
| 358 |
+
def bind_tools(self, tools, **kwargs):
|
| 359 |
+
return self
|
| 360 |
+
|
| 361 |
+
def invoke(self, prompt, config=None):
|
| 362 |
+
self.invocations.append((list(prompt), config))
|
| 363 |
+
return self.message_responses.pop(0)
|
| 364 |
+
|
| 365 |
+
@staticmethod
|
| 366 |
+
def _structured_request(model, messages: list, turn_id: str) -> ModelRequest:
|
| 367 |
+
return ModelRequest(
|
| 368 |
+
model=model,
|
| 369 |
+
messages=messages,
|
| 370 |
+
system_message=SystemMessage(content="system"),
|
| 371 |
+
tools=[],
|
| 372 |
+
tool_choice=None,
|
| 373 |
+
response_format=None,
|
| 374 |
+
model_settings={},
|
| 375 |
+
state={"messages": messages},
|
| 376 |
+
runtime=SimpleNamespace(
|
| 377 |
+
context=SimpleNamespace(kb_session_id=turn_id, cache_user_id="")
|
| 378 |
+
),
|
| 379 |
+
)
|
| 380 |
+
|
| 381 |
+
def test_cache_user_id_is_added_to_agent_model_settings(self) -> None:
|
| 382 |
+
runtime = SimpleNamespace(
|
| 383 |
+
context=SimpleNamespace(cache_user_id="eval_abc", kb_session_id="turn")
|
| 384 |
+
)
|
| 385 |
+
request = SimpleNamespace(
|
| 386 |
+
runtime=runtime, model_settings={}, messages=[], system_message=None
|
| 387 |
+
)
|
| 388 |
+
request.override = lambda **updates: SimpleNamespace(
|
| 389 |
+
runtime=runtime,
|
| 390 |
+
model_settings=updates.get("model_settings", request.model_settings),
|
| 391 |
+
)
|
| 392 |
+
isolated = DeepSeekCacheIsolationMiddleware()._isolate(request)
|
| 393 |
+
self.assertEqual(isolated.model_settings["extra_body"], {"user_id": "eval_abc"})
|
| 394 |
+
|
| 395 |
+
def test_agent_request_guard_fails_before_model_handler(self) -> None:
|
| 396 |
+
runtime = SimpleNamespace(
|
| 397 |
+
context=SimpleNamespace(cache_user_id="eval_guard", kb_session_id="turn")
|
| 398 |
+
)
|
| 399 |
+
request = SimpleNamespace(
|
| 400 |
+
runtime=runtime,
|
| 401 |
+
model_settings={},
|
| 402 |
+
messages=[HumanMessage(content="x" * 4_000)],
|
| 403 |
+
system_message=None,
|
| 404 |
+
)
|
| 405 |
+
with self.assertRaisesRegex(RuntimeError, "Agent request exceeds"):
|
| 406 |
+
DeepSeekCacheIsolationMiddleware(100)._isolate(request)
|
| 407 |
+
|
| 408 |
+
def test_full_selected_history_reaches_summarizer_and_records_event(self) -> None:
|
| 409 |
+
model = self.FakeModel()
|
| 410 |
+
middleware = InstrumentedSummarizationMiddleware(
|
| 411 |
+
model=model,
|
| 412 |
+
trigger=("tokens", 1_000),
|
| 413 |
+
keep=("tokens", 500),
|
| 414 |
+
trim_tokens_to_summarize=None,
|
| 415 |
+
)
|
| 416 |
+
messages = [
|
| 417 |
+
(HumanMessage if index % 2 == 0 else AIMessage)(content="x" * 1_000)
|
| 418 |
+
for index in range(28)
|
| 419 |
+
]
|
| 420 |
+
runtime = SimpleNamespace(
|
| 421 |
+
context=SimpleNamespace(
|
| 422 |
+
kb_session_id="summary-turn", cache_user_id="eval_summary"
|
| 423 |
+
)
|
| 424 |
+
)
|
| 425 |
+
reset_turn_signals("summary-turn")
|
| 426 |
+
update = middleware.before_model({"messages": messages}, runtime)
|
| 427 |
+
self.assertIsNotNone(update)
|
| 428 |
+
events = pop_turn_events("summary-turn")
|
| 429 |
+
self.assertEqual(len(events), 1)
|
| 430 |
+
event = events[0]
|
| 431 |
+
self.assertTrue(event["summary_input_untrimmed"])
|
| 432 |
+
self.assertEqual(event["configured_trigger_tokens"], 1_000)
|
| 433 |
+
self.assertGreater(event["summary_input_tokens_approx"], 4_000)
|
| 434 |
+
self.assertLessEqual(event["retained_tail_tokens_approx"], 500)
|
| 435 |
+
self.assertEqual(pop_turn_signals("summary-turn")["compactions_this_turn"], 1)
|
| 436 |
+
self.assertEqual(model.bound[-1]["extra_body"], {"user_id": "eval_summary"})
|
| 437 |
+
self.assertGreater(len(model.prompts[-1]), 20_000)
|
| 438 |
+
|
| 439 |
+
def test_provider_reported_tokens_can_trigger_below_approximation(self) -> None:
|
| 440 |
+
model = self.FakeModel()
|
| 441 |
+
middleware = InstrumentedSummarizationMiddleware(
|
| 442 |
+
model=model,
|
| 443 |
+
trigger=("tokens", 200_000),
|
| 444 |
+
keep=("tokens", 500),
|
| 445 |
+
trim_tokens_to_summarize=None,
|
| 446 |
+
)
|
| 447 |
+
messages = [HumanMessage(content="x" * 2_000) for _ in range(6)]
|
| 448 |
+
messages.append(
|
| 449 |
+
AIMessage(
|
| 450 |
+
content="previous answer",
|
| 451 |
+
usage_metadata={
|
| 452 |
+
"input_tokens": 205_664,
|
| 453 |
+
"output_tokens": 1_672,
|
| 454 |
+
"total_tokens": 207_336,
|
| 455 |
+
},
|
| 456 |
+
response_metadata={"model_provider": "deepseek"},
|
| 457 |
+
)
|
| 458 |
+
)
|
| 459 |
+
runtime = SimpleNamespace(
|
| 460 |
+
context=SimpleNamespace(
|
| 461 |
+
kb_session_id="reported-trigger", cache_user_id="eval_reported"
|
| 462 |
+
)
|
| 463 |
+
)
|
| 464 |
+
reset_turn_signals("reported-trigger")
|
| 465 |
+
with mock.patch.object(middleware, "token_counter", return_value=199_567):
|
| 466 |
+
update = middleware.before_model({"messages": messages}, runtime)
|
| 467 |
+
self.assertIsNotNone(update)
|
| 468 |
+
event = pop_turn_events("reported-trigger")[0]
|
| 469 |
+
self.assertEqual(event["pre_compaction_tokens_approx"], 199_567)
|
| 470 |
+
self.assertEqual(event["trigger_reported_tokens"], 207_336)
|
| 471 |
+
self.assertEqual(event["trigger_source"], "provider_reported")
|
| 472 |
+
|
| 473 |
+
def test_summary_input_guard_fails_before_provider_call(self) -> None:
|
| 474 |
+
model = self.FakeModel()
|
| 475 |
+
middleware = InstrumentedSummarizationMiddleware(
|
| 476 |
+
model=model,
|
| 477 |
+
trigger=("tokens", 100),
|
| 478 |
+
keep=("tokens", 100),
|
| 479 |
+
trim_tokens_to_summarize=None,
|
| 480 |
+
summary_input_guard_tokens=200,
|
| 481 |
+
)
|
| 482 |
+
messages = [HumanMessage(content="x" * 2_000) for _ in range(4)]
|
| 483 |
+
runtime = SimpleNamespace(
|
| 484 |
+
context=SimpleNamespace(kb_session_id="guard", cache_user_id="eval_guard")
|
| 485 |
+
)
|
| 486 |
+
with self.assertRaisesRegex(RuntimeError, "safety guard"):
|
| 487 |
+
middleware.before_model({"messages": messages}, runtime)
|
| 488 |
+
self.assertEqual(model.prompts, [])
|
| 489 |
+
|
| 490 |
+
def test_empty_summary_is_retried_and_recorded(self) -> None:
|
| 491 |
+
model = self.FakeModel(["", "durable retry summary"])
|
| 492 |
+
middleware = InstrumentedSummarizationMiddleware(
|
| 493 |
+
model=model,
|
| 494 |
+
trigger=("tokens", 100),
|
| 495 |
+
keep=("tokens", 100),
|
| 496 |
+
trim_tokens_to_summarize=None,
|
| 497 |
+
)
|
| 498 |
+
messages = [HumanMessage(content="x" * 2_000) for _ in range(4)]
|
| 499 |
+
runtime = SimpleNamespace(
|
| 500 |
+
context=SimpleNamespace(kb_session_id="retry", cache_user_id="eval_retry")
|
| 501 |
+
)
|
| 502 |
+
reset_turn_signals("retry")
|
| 503 |
+
with mock.patch("app.chat_service.time.sleep") as sleep:
|
| 504 |
+
update = middleware.before_model({"messages": messages}, runtime)
|
| 505 |
+
self.assertIsNotNone(update)
|
| 506 |
+
self.assertEqual(len(model.prompts), 2)
|
| 507 |
+
sleep.assert_called_once_with(1.0)
|
| 508 |
+
event = pop_turn_events("retry")[0]
|
| 509 |
+
self.assertEqual(event["summary_attempts"], 2)
|
| 510 |
+
self.assertEqual(event["summary_retry_reasons"], ["empty response"])
|
| 511 |
+
|
| 512 |
+
def test_non_retryable_summary_failure_is_not_retried(self) -> None:
|
| 513 |
+
model = self.FakeModel()
|
| 514 |
+
model.invoke = mock.Mock(side_effect=ValueError("invalid request"))
|
| 515 |
+
middleware = InstrumentedSummarizationMiddleware(
|
| 516 |
+
model=model,
|
| 517 |
+
trigger=("tokens", 100),
|
| 518 |
+
keep=("tokens", 100),
|
| 519 |
+
trim_tokens_to_summarize=None,
|
| 520 |
+
)
|
| 521 |
+
messages = [HumanMessage(content="x" * 2_000) for _ in range(4)]
|
| 522 |
+
runtime = SimpleNamespace(
|
| 523 |
+
context=SimpleNamespace(kb_session_id="no-retry", cache_user_id="eval")
|
| 524 |
+
)
|
| 525 |
+
with self.assertRaisesRegex(ValueError, "invalid request"):
|
| 526 |
+
middleware.before_model({"messages": messages}, runtime)
|
| 527 |
+
self.assertEqual(model.invoke.call_count, 1)
|
| 528 |
+
|
| 529 |
+
def test_real_agent_attributes_summary_and_agent_calls_separately(self) -> None:
|
| 530 |
+
def response(text: str, input_tokens: int, output_tokens: int) -> AIMessage:
|
| 531 |
+
return AIMessage(
|
| 532 |
+
content=text,
|
| 533 |
+
usage_metadata={
|
| 534 |
+
"input_tokens": input_tokens,
|
| 535 |
+
"output_tokens": output_tokens,
|
| 536 |
+
"total_tokens": input_tokens + output_tokens,
|
| 537 |
+
"input_token_details": {"cache_read": 0},
|
| 538 |
+
},
|
| 539 |
+
response_metadata={"model_name": "deepseek-v4-flash"},
|
| 540 |
+
)
|
| 541 |
+
|
| 542 |
+
model = FakeMessagesListChatModel(
|
| 543 |
+
responses=[response("summary", 6_000, 10), response("answer", 700, 20)]
|
| 544 |
+
)
|
| 545 |
+
summary = InstrumentedSummarizationMiddleware(
|
| 546 |
+
model=model,
|
| 547 |
+
trigger=("tokens", 1_000),
|
| 548 |
+
keep=("tokens", 500),
|
| 549 |
+
trim_tokens_to_summarize=None,
|
| 550 |
+
summary_input_guard_tokens=900_000,
|
| 551 |
+
)
|
| 552 |
+
agent = create_agent(
|
| 553 |
+
model=model,
|
| 554 |
+
tools=[],
|
| 555 |
+
middleware=[DeepSeekCacheIsolationMiddleware(900_000), summary],
|
| 556 |
+
)
|
| 557 |
+
messages = [
|
| 558 |
+
(HumanMessage if index % 2 == 0 else AIMessage)(content="x" * 1_000)
|
| 559 |
+
for index in range(28)
|
| 560 |
+
]
|
| 561 |
+
handler = TurnUsageHandler()
|
| 562 |
+
reset_turn_signals("integration-turn")
|
| 563 |
+
result = agent.invoke(
|
| 564 |
+
{"messages": messages},
|
| 565 |
+
config={"callbacks": [handler]},
|
| 566 |
+
context=AppContext(
|
| 567 |
+
allowed_sources=(),
|
| 568 |
+
kb_session_id="integration-turn",
|
| 569 |
+
cache_user_id="eval_integration",
|
| 570 |
+
),
|
| 571 |
+
)
|
| 572 |
+
self.assertEqual(handler.llm_calls, 2)
|
| 573 |
+
self.assertEqual(
|
| 574 |
+
[call["source"] for call in handler.model_calls],
|
| 575 |
+
["summarization", "agent"],
|
| 576 |
+
)
|
| 577 |
+
self.assertEqual(result["messages"][-1].content, "answer")
|
| 578 |
+
event = pop_turn_events("integration-turn")[0]
|
| 579 |
+
self.assertGreater(event["summary_input_tokens_approx"], 4_000)
|
| 580 |
+
self.assertEqual(
|
| 581 |
+
pop_turn_signals("integration-turn")["compactions_this_turn"], 1
|
| 582 |
+
)
|
| 583 |
+
|
| 584 |
+
def test_structured_prefix_preserves_request_shape_and_persists_checkpoint(
|
| 585 |
+
self,
|
| 586 |
+
) -> None:
|
| 587 |
+
model = self.StructuredFakeModel(["durable structured summary"])
|
| 588 |
+
middleware = PrefixPreservingCompactionMiddleware(
|
| 589 |
+
model=model,
|
| 590 |
+
trigger=("tokens", 1_000),
|
| 591 |
+
keep=("tokens", 500),
|
| 592 |
+
trim_tokens_to_summarize=None,
|
| 593 |
+
summary_input_guard_tokens=900_000,
|
| 594 |
+
)
|
| 595 |
+
messages = [
|
| 596 |
+
(HumanMessage if index % 2 == 0 else AIMessage)(
|
| 597 |
+
content=f"m{index}:" + "x" * 1_000
|
| 598 |
+
)
|
| 599 |
+
for index in range(28)
|
| 600 |
+
]
|
| 601 |
+
system = SystemMessage(content="stable system prompt")
|
| 602 |
+
tools = [
|
| 603 |
+
{
|
| 604 |
+
"type": "function",
|
| 605 |
+
"function": {
|
| 606 |
+
"name": "lookup",
|
| 607 |
+
"description": "lookup evidence",
|
| 608 |
+
"parameters": {"type": "object", "properties": {}},
|
| 609 |
+
},
|
| 610 |
+
}
|
| 611 |
+
]
|
| 612 |
+
runtime = SimpleNamespace(
|
| 613 |
+
context=SimpleNamespace(
|
| 614 |
+
kb_session_id="prefix-turn", cache_user_id="stable-prefix-user"
|
| 615 |
+
)
|
| 616 |
+
)
|
| 617 |
+
request = ModelRequest(
|
| 618 |
+
model=model,
|
| 619 |
+
messages=messages,
|
| 620 |
+
system_message=system,
|
| 621 |
+
tools=tools,
|
| 622 |
+
tool_choice=None,
|
| 623 |
+
response_format=None,
|
| 624 |
+
model_settings={"extra_body": {"user_id": "stable-prefix-user"}},
|
| 625 |
+
state={"messages": messages},
|
| 626 |
+
runtime=runtime,
|
| 627 |
+
)
|
| 628 |
+
expected_plan = middleware._plan_compaction(request.state)
|
| 629 |
+
self.assertIsNotNone(expected_plan)
|
| 630 |
+
handled: list[ModelRequest] = []
|
| 631 |
+
|
| 632 |
+
def handler(compacted_request):
|
| 633 |
+
handled.append(compacted_request)
|
| 634 |
+
return ModelResponse(
|
| 635 |
+
result=[AIMessage(content="final answer", id="answer")]
|
| 636 |
+
)
|
| 637 |
+
|
| 638 |
+
reset_turn_signals("prefix-turn")
|
| 639 |
+
result = middleware.wrap_model_call(request, handler)
|
| 640 |
+
|
| 641 |
+
self.assertEqual(model.bound_tools[0][0], tools)
|
| 642 |
+
self.assertEqual(
|
| 643 |
+
model.bound_tools[0][1]["extra_body"],
|
| 644 |
+
{"user_id": "stable-prefix-user"},
|
| 645 |
+
)
|
| 646 |
+
summary_messages, summary_config = model.invocations[0]
|
| 647 |
+
self.assertIs(summary_messages[0], system)
|
| 648 |
+
self.assertEqual(summary_messages[1:-1], messages)
|
| 649 |
+
self.assertNotIn("<messages>", summary_messages[-1].content)
|
| 650 |
+
self.assertIn(
|
| 651 |
+
f"final {len(expected_plan['preserved'])} messages",
|
| 652 |
+
summary_messages[-1].content,
|
| 653 |
+
)
|
| 654 |
+
self.assertIn(
|
| 655 |
+
f"approximately {middleware._partial_token_counter(expected_plan['preserved'])} tokens",
|
| 656 |
+
summary_messages[-1].content,
|
| 657 |
+
)
|
| 658 |
+
self.assertEqual(summary_config["metadata"]["lc_source"], "summarization")
|
| 659 |
+
self.assertEqual(
|
| 660 |
+
summary_config["metadata"]["compaction_strategy"],
|
| 661 |
+
"structured_prefix",
|
| 662 |
+
)
|
| 663 |
+
|
| 664 |
+
compacted = handled[0].messages
|
| 665 |
+
self.assertEqual(
|
| 666 |
+
compacted[0].additional_kwargs.get("lc_source"), "summarization"
|
| 667 |
+
)
|
| 668 |
+
self.assertEqual(handled[0].state["messages"], compacted)
|
| 669 |
+
command_messages = result.command.update["messages"]
|
| 670 |
+
self.assertIsInstance(command_messages[0], RemoveMessage)
|
| 671 |
+
self.assertEqual(command_messages[0].id, REMOVE_ALL_MESSAGES)
|
| 672 |
+
self.assertEqual(command_messages[-1].content, "final answer")
|
| 673 |
+
|
| 674 |
+
event = pop_turn_events("prefix-turn")[0]
|
| 675 |
+
self.assertEqual(event["summary_strategy"], "structured_prefix")
|
| 676 |
+
self.assertTrue(event["summary_request_is_strict_extension"])
|
| 677 |
+
self.assertEqual(event["summary_prefix_messages"], len(messages))
|
| 678 |
+
self.assertLess(event["summary_selected_messages"], len(messages))
|
| 679 |
+
self.assertEqual(
|
| 680 |
+
event["summary_instruction_retained_messages"],
|
| 681 |
+
len(expected_plan["preserved"]),
|
| 682 |
+
)
|
| 683 |
+
self.assertTrue(event["summary_system_message_present"])
|
| 684 |
+
self.assertEqual(event["summary_tools_bound"], 1)
|
| 685 |
+
self.assertTrue(event["summary_cache_user_id_preserved"])
|
| 686 |
+
self.assertEqual(event["summary_provider_input_tokens"], 10_000)
|
| 687 |
+
self.assertEqual(event["summary_provider_cache_read_tokens"], 9_000)
|
| 688 |
+
self.assertEqual(event["summary_provider_cache_miss_tokens"], 1_000)
|
| 689 |
+
self.assertEqual(event["summary_provider_cache_hit_ratio"], 0.9)
|
| 690 |
+
|
| 691 |
+
def test_structured_prefix_safe_tail_never_starts_with_orphaned_tool(self) -> None:
|
| 692 |
+
model = self.StructuredFakeModel(["summary"])
|
| 693 |
+
middleware = PrefixPreservingCompactionMiddleware(
|
| 694 |
+
model=model,
|
| 695 |
+
trigger=("tokens", 100),
|
| 696 |
+
keep=("tokens", 120),
|
| 697 |
+
trim_tokens_to_summarize=None,
|
| 698 |
+
)
|
| 699 |
+
messages = [
|
| 700 |
+
HumanMessage(content="old " + "x" * 2_000),
|
| 701 |
+
AIMessage(content="old answer " + "x" * 2_000),
|
| 702 |
+
HumanMessage(content="tool turn"),
|
| 703 |
+
AIMessage(
|
| 704 |
+
content="",
|
| 705 |
+
tool_calls=[{"name": "lookup", "args": {}, "id": "call-1"}],
|
| 706 |
+
),
|
| 707 |
+
ToolMessage(content="evidence", tool_call_id="call-1"),
|
| 708 |
+
AIMessage(content="tool answer"),
|
| 709 |
+
HumanMessage(content="current"),
|
| 710 |
+
]
|
| 711 |
+
plan = middleware._plan_compaction({"messages": messages})
|
| 712 |
+
self.assertIsNotNone(plan)
|
| 713 |
+
self.assertTrue(plan["preserved"])
|
| 714 |
+
self.assertNotIsInstance(plan["preserved"][0], ToolMessage)
|
| 715 |
+
|
| 716 |
+
def test_structured_prefix_rejects_an_unpersisted_message_view(self) -> None:
|
| 717 |
+
model = self.StructuredFakeModel(["summary"])
|
| 718 |
+
middleware = PrefixPreservingCompactionMiddleware(
|
| 719 |
+
model=model,
|
| 720 |
+
trigger=("tokens", 100),
|
| 721 |
+
keep=("tokens", 100),
|
| 722 |
+
trim_tokens_to_summarize=None,
|
| 723 |
+
)
|
| 724 |
+
state_messages = [HumanMessage(content="x" * 2_000) for _ in range(4)]
|
| 725 |
+
request = ModelRequest(
|
| 726 |
+
model=model,
|
| 727 |
+
messages=state_messages[1:],
|
| 728 |
+
system_message=SystemMessage(content="system"),
|
| 729 |
+
tools=[],
|
| 730 |
+
response_format=None,
|
| 731 |
+
state={"messages": state_messages},
|
| 732 |
+
runtime=SimpleNamespace(
|
| 733 |
+
context=SimpleNamespace(kb_session_id="mismatched-view")
|
| 734 |
+
),
|
| 735 |
+
)
|
| 736 |
+
plan = middleware._plan_compaction(request.state)
|
| 737 |
+
self.assertIsNotNone(plan)
|
| 738 |
+
with self.assertRaisesRegex(RuntimeError, "match checkpoint history"):
|
| 739 |
+
middleware._prepare_summary_request(request, plan)
|
| 740 |
+
|
| 741 |
+
def test_structured_prefix_empty_summary_is_retried_and_recorded(self) -> None:
|
| 742 |
+
# Empty text and tool_calls-with-empty-text are both "empty" responses:
|
| 743 |
+
# each is retried and shows up in summary_retry_reasons.
|
| 744 |
+
model = self.ScriptedMessageModel(
|
| 745 |
+
[
|
| 746 |
+
AIMessage(content=""),
|
| 747 |
+
AIMessage(
|
| 748 |
+
content="",
|
| 749 |
+
tool_calls=[{"name": "lookup", "args": {}, "id": "call-empty"}],
|
| 750 |
+
),
|
| 751 |
+
AIMessage(content="structured retry checkpoint"),
|
| 752 |
+
]
|
| 753 |
+
)
|
| 754 |
+
middleware = PrefixPreservingCompactionMiddleware(
|
| 755 |
+
model=model,
|
| 756 |
+
trigger=("tokens", 100),
|
| 757 |
+
keep=("tokens", 100),
|
| 758 |
+
trim_tokens_to_summarize=None,
|
| 759 |
+
)
|
| 760 |
+
messages = [HumanMessage(content="x" * 2_000, id=f"h{i}") for i in range(4)]
|
| 761 |
+
request = self._structured_request(model, messages, "structured-retry")
|
| 762 |
+
handled: list[ModelRequest] = []
|
| 763 |
+
|
| 764 |
+
def handler(compacted_request):
|
| 765 |
+
handled.append(compacted_request)
|
| 766 |
+
return ModelResponse(result=[AIMessage(content="answer", id="a")])
|
| 767 |
+
|
| 768 |
+
reset_turn_signals("structured-retry")
|
| 769 |
+
with mock.patch("app.chat_service.time.sleep") as sleep:
|
| 770 |
+
middleware.wrap_model_call(request, handler)
|
| 771 |
+
self.assertEqual(len(model.invocations), 3)
|
| 772 |
+
self.assertEqual(sleep.call_args_list, [mock.call(1.0), mock.call(2.0)])
|
| 773 |
+
self.assertEqual(len(handled), 1)
|
| 774 |
+
self.assertIn("structured retry checkpoint", handled[0].messages[0].content)
|
| 775 |
+
event = pop_turn_events("structured-retry")[0]
|
| 776 |
+
self.assertEqual(event["summary_attempts"], 3)
|
| 777 |
+
self.assertEqual(
|
| 778 |
+
event["summary_retry_reasons"], ["empty response", "empty response"]
|
| 779 |
+
)
|
| 780 |
+
self.assertEqual(
|
| 781 |
+
pop_turn_signals("structured-retry")["compactions_this_turn"], 1
|
| 782 |
+
)
|
| 783 |
+
|
| 784 |
+
def test_structured_prefix_empty_summary_raises_after_max_attempts(self) -> None:
|
| 785 |
+
model = self.ScriptedMessageModel([AIMessage(content="")] * 3)
|
| 786 |
+
middleware = PrefixPreservingCompactionMiddleware(
|
| 787 |
+
model=model,
|
| 788 |
+
trigger=("tokens", 100),
|
| 789 |
+
keep=("tokens", 100),
|
| 790 |
+
trim_tokens_to_summarize=None,
|
| 791 |
+
)
|
| 792 |
+
messages = [HumanMessage(content="x" * 2_000, id=f"h{i}") for i in range(4)]
|
| 793 |
+
request = self._structured_request(model, messages, "structured-exhausted")
|
| 794 |
+
handled: list[ModelRequest] = []
|
| 795 |
+
reset_turn_signals("structured-exhausted")
|
| 796 |
+
with mock.patch("app.chat_service.time.sleep"):
|
| 797 |
+
with self.assertRaisesRegex(RuntimeError, "empty summary after 3 attempts"):
|
| 798 |
+
middleware.wrap_model_call(request, handled.append)
|
| 799 |
+
self.assertEqual(len(model.invocations), 3)
|
| 800 |
+
# The agent call never runs on a failed checkpoint, and nothing is
|
| 801 |
+
# recorded as a successful compaction.
|
| 802 |
+
self.assertEqual(handled, [])
|
| 803 |
+
self.assertEqual(pop_turn_events("structured-exhausted"), [])
|
| 804 |
+
self.assertNotIn(
|
| 805 |
+
"compactions_this_turn", pop_turn_signals("structured-exhausted")
|
| 806 |
+
)
|
| 807 |
+
|
| 808 |
+
def test_structured_prefix_uses_text_and_ignores_summary_tool_calls(self) -> None:
|
| 809 |
+
model = self.ScriptedMessageModel(
|
| 810 |
+
[
|
| 811 |
+
AIMessage(
|
| 812 |
+
content="checkpoint despite tool call",
|
| 813 |
+
tool_calls=[{"name": "lookup", "args": {}, "id": "call-x"}],
|
| 814 |
+
)
|
| 815 |
+
]
|
| 816 |
+
)
|
| 817 |
+
middleware = PrefixPreservingCompactionMiddleware(
|
| 818 |
+
model=model,
|
| 819 |
+
trigger=("tokens", 100),
|
| 820 |
+
keep=("tokens", 100),
|
| 821 |
+
trim_tokens_to_summarize=None,
|
| 822 |
+
)
|
| 823 |
+
messages = [HumanMessage(content="x" * 2_000, id=f"h{i}") for i in range(4)]
|
| 824 |
+
request = self._structured_request(model, messages, "structured-toolcall")
|
| 825 |
+
handled: list[ModelRequest] = []
|
| 826 |
+
|
| 827 |
+
def handler(compacted_request):
|
| 828 |
+
handled.append(compacted_request)
|
| 829 |
+
return ModelResponse(result=[AIMessage(content="answer", id="a")])
|
| 830 |
+
|
| 831 |
+
reset_turn_signals("structured-toolcall")
|
| 832 |
+
result = middleware.wrap_model_call(request, handler)
|
| 833 |
+
self.assertEqual(len(model.invocations), 1)
|
| 834 |
+
self.assertIn("checkpoint despite tool call", handled[0].messages[0].content)
|
| 835 |
+
# The summarizer's AIMessage never enters state, so its tool calls can
|
| 836 |
+
# never be executed: no message anywhere carries call-x.
|
| 837 |
+
command_messages = result.command.update["messages"]
|
| 838 |
+
self.assertFalse(
|
| 839 |
+
any(
|
| 840 |
+
call["id"] == "call-x"
|
| 841 |
+
for message in [*handled[0].messages, *command_messages]
|
| 842 |
+
if isinstance(message, AIMessage)
|
| 843 |
+
for call in (message.tool_calls or [])
|
| 844 |
+
)
|
| 845 |
+
)
|
| 846 |
+
event = pop_turn_events("structured-toolcall")[0]
|
| 847 |
+
self.assertEqual(event["summary_attempts"], 1)
|
| 848 |
+
self.assertEqual(event["summary_retry_reasons"], [])
|
| 849 |
+
self.assertEqual(
|
| 850 |
+
pop_turn_signals("structured-toolcall")["compactions_this_turn"], 1
|
| 851 |
+
)
|
| 852 |
+
|
| 853 |
+
def test_real_agent_structured_compaction_rewrites_checkpoint_once(self) -> None:
|
| 854 |
+
def response(text: str, input_tokens: int, cache_read: int) -> AIMessage:
|
| 855 |
+
return AIMessage(
|
| 856 |
+
content=text,
|
| 857 |
+
usage_metadata={
|
| 858 |
+
"input_tokens": input_tokens,
|
| 859 |
+
"output_tokens": 10,
|
| 860 |
+
"total_tokens": input_tokens + 10,
|
| 861 |
+
"input_token_details": {"cache_read": cache_read},
|
| 862 |
+
},
|
| 863 |
+
response_metadata={"model_name": "deepseek-v4-flash"},
|
| 864 |
+
)
|
| 865 |
+
|
| 866 |
+
model = FakeMessagesListChatModel(
|
| 867 |
+
responses=[
|
| 868 |
+
response("structured checkpoint", 6_000, 5_500),
|
| 869 |
+
response("answer after checkpoint", 700, 0),
|
| 870 |
+
]
|
| 871 |
+
)
|
| 872 |
+
compactor = PrefixPreservingCompactionMiddleware(
|
| 873 |
+
model=model,
|
| 874 |
+
trigger=("tokens", 1_000),
|
| 875 |
+
keep=("tokens", 500),
|
| 876 |
+
trim_tokens_to_summarize=None,
|
| 877 |
+
summary_input_guard_tokens=900_000,
|
| 878 |
+
)
|
| 879 |
+
agent = create_agent(
|
| 880 |
+
model=model,
|
| 881 |
+
tools=[],
|
| 882 |
+
system_prompt="stable system",
|
| 883 |
+
middleware=[DeepSeekCacheIsolationMiddleware(900_000), compactor],
|
| 884 |
+
)
|
| 885 |
+
messages = [
|
| 886 |
+
(HumanMessage if index % 2 == 0 else AIMessage)(
|
| 887 |
+
content=f"old-{index}:" + "x" * 1_000
|
| 888 |
+
)
|
| 889 |
+
for index in range(28)
|
| 890 |
+
]
|
| 891 |
+
handler = TurnUsageHandler()
|
| 892 |
+
reset_turn_signals("structured-integration")
|
| 893 |
+
result = agent.invoke(
|
| 894 |
+
{"messages": messages},
|
| 895 |
+
config={"callbacks": [handler]},
|
| 896 |
+
context=AppContext(
|
| 897 |
+
allowed_sources=(),
|
| 898 |
+
kb_session_id="structured-integration",
|
| 899 |
+
cache_user_id="eval-structured-integration",
|
| 900 |
+
),
|
| 901 |
+
)
|
| 902 |
+
|
| 903 |
+
self.assertEqual(handler.llm_calls, 2)
|
| 904 |
+
self.assertEqual(
|
| 905 |
+
[call["source"] for call in handler.model_calls],
|
| 906 |
+
["summarization", "agent"],
|
| 907 |
+
)
|
| 908 |
+
summaries = [
|
| 909 |
+
message
|
| 910 |
+
for message in result["messages"]
|
| 911 |
+
if message.additional_kwargs.get("lc_source") == "summarization"
|
| 912 |
+
]
|
| 913 |
+
self.assertEqual(len(summaries), 1)
|
| 914 |
+
self.assertIn("structured checkpoint", summaries[0].content)
|
| 915 |
+
answers = [
|
| 916 |
+
message
|
| 917 |
+
for message in result["messages"]
|
| 918 |
+
if isinstance(message, AIMessage)
|
| 919 |
+
and message.content == "answer after checkpoint"
|
| 920 |
+
]
|
| 921 |
+
self.assertEqual(len(answers), 1)
|
| 922 |
+
self.assertFalse(
|
| 923 |
+
any(
|
| 924 |
+
message.content == messages[0].content for message in result["messages"]
|
| 925 |
+
)
|
| 926 |
+
)
|
| 927 |
+
event = pop_turn_events("structured-integration")[0]
|
| 928 |
+
self.assertEqual(event["summary_provider_cache_read_tokens"], 5_500)
|
| 929 |
+
self.assertEqual(event["summary_provider_cache_miss_tokens"], 500)
|
| 930 |
+
self.assertEqual(
|
| 931 |
+
pop_turn_signals("structured-integration")["compactions_this_turn"],
|
| 932 |
+
1,
|
| 933 |
+
)
|
| 934 |
+
|
| 935 |
+
|
| 936 |
+
class ExperimentCompactionMiddlewareAsyncTests(unittest.IsolatedAsyncioTestCase):
|
| 937 |
+
async def test_async_empty_summary_is_retried(self) -> None:
|
| 938 |
+
model = ExperimentCompactionMiddlewareTests.FakeModel(["", "async summary"])
|
| 939 |
+
middleware = InstrumentedSummarizationMiddleware(
|
| 940 |
+
model=model,
|
| 941 |
+
trigger=("tokens", 100),
|
| 942 |
+
keep=("tokens", 100),
|
| 943 |
+
trim_tokens_to_summarize=None,
|
| 944 |
+
)
|
| 945 |
+
messages = [HumanMessage(content="x" * 2_000) for _ in range(4)]
|
| 946 |
+
runtime = SimpleNamespace(
|
| 947 |
+
context=SimpleNamespace(
|
| 948 |
+
kb_session_id="async-retry", cache_user_id="eval_async_retry"
|
| 949 |
+
)
|
| 950 |
+
)
|
| 951 |
+
reset_turn_signals("async-retry")
|
| 952 |
+
with mock.patch("app.chat_service.asyncio.sleep") as sleep:
|
| 953 |
+
update = await middleware.abefore_model({"messages": messages}, runtime)
|
| 954 |
+
self.assertIsNotNone(update)
|
| 955 |
+
self.assertEqual(len(model.prompts), 2)
|
| 956 |
+
sleep.assert_awaited_once_with(1.0)
|
| 957 |
+
event = pop_turn_events("async-retry")[0]
|
| 958 |
+
self.assertEqual(event["summary_attempts"], 2)
|
| 959 |
+
|
| 960 |
+
|
| 961 |
+
class CompactionPathEquivalenceTests(unittest.TestCase):
|
| 962 |
+
"""Same history + config: both paths must agree on the compaction boundary."""
|
| 963 |
+
|
| 964 |
+
TRIGGER = ("tokens", 1_000)
|
| 965 |
+
KEEP = ("tokens", 500)
|
| 966 |
+
|
| 967 |
+
def _history(self) -> list:
|
| 968 |
+
# Pre-assigned ids let boundary selection be compared across paths.
|
| 969 |
+
messages = [
|
| 970 |
+
(HumanMessage if index % 2 == 0 else AIMessage)(
|
| 971 |
+
content=f"m{index}:" + "x" * 1_000, id=f"m{index}"
|
| 972 |
+
)
|
| 973 |
+
for index in range(24)
|
| 974 |
+
]
|
| 975 |
+
messages += [
|
| 976 |
+
AIMessage(
|
| 977 |
+
content="",
|
| 978 |
+
id="m-toolcall",
|
| 979 |
+
tool_calls=[{"name": "lookup", "args": {}, "id": "call-1"}],
|
| 980 |
+
),
|
| 981 |
+
ToolMessage(content="evidence", tool_call_id="call-1", id="m-tool"),
|
| 982 |
+
AIMessage(content="tool answer", id="m-tool-answer"),
|
| 983 |
+
HumanMessage(content="current question", id="m-current"),
|
| 984 |
+
]
|
| 985 |
+
return messages
|
| 986 |
+
|
| 987 |
+
def test_both_paths_select_the_identical_boundary(self) -> None:
|
| 988 |
+
history = self._history()
|
| 989 |
+
xml_plan = InstrumentedSummarizationMiddleware(
|
| 990 |
+
model=ExperimentCompactionMiddlewareTests.FakeModel(),
|
| 991 |
+
trigger=self.TRIGGER,
|
| 992 |
+
keep=self.KEEP,
|
| 993 |
+
trim_tokens_to_summarize=None,
|
| 994 |
+
)._plan_compaction({"messages": history})
|
| 995 |
+
structured_plan = PrefixPreservingCompactionMiddleware(
|
| 996 |
+
model=ExperimentCompactionMiddlewareTests.StructuredFakeModel(),
|
| 997 |
+
trigger=self.TRIGGER,
|
| 998 |
+
keep=self.KEEP,
|
| 999 |
+
trim_tokens_to_summarize=None,
|
| 1000 |
+
)._plan_compaction({"messages": history})
|
| 1001 |
+
self.assertIsNotNone(xml_plan)
|
| 1002 |
+
self.assertIsNotNone(structured_plan)
|
| 1003 |
+
self.assertEqual(
|
| 1004 |
+
[m.id for m in xml_plan["selected"]],
|
| 1005 |
+
[m.id for m in structured_plan["selected"]],
|
| 1006 |
+
)
|
| 1007 |
+
self.assertEqual(
|
| 1008 |
+
[m.id for m in xml_plan["preserved"]],
|
| 1009 |
+
[m.id for m in structured_plan["preserved"]],
|
| 1010 |
+
)
|
| 1011 |
+
# The boundary is a clean partition of the full history.
|
| 1012 |
+
self.assertEqual(
|
| 1013 |
+
[m.id for m in [*xml_plan["selected"], *xml_plan["preserved"]]],
|
| 1014 |
+
[m.id for m in history],
|
| 1015 |
+
)
|
| 1016 |
+
self.assertGreater(len(xml_plan["selected"]), 0)
|
| 1017 |
+
self.assertGreater(len(xml_plan["preserved"]), 0)
|
| 1018 |
+
|
| 1019 |
+
def test_both_paths_install_the_same_post_compaction_structure(self) -> None:
|
| 1020 |
+
history = self._history()
|
| 1021 |
+
|
| 1022 |
+
xml_middleware = InstrumentedSummarizationMiddleware(
|
| 1023 |
+
model=ExperimentCompactionMiddlewareTests.FakeModel(["xml summary"]),
|
| 1024 |
+
trigger=self.TRIGGER,
|
| 1025 |
+
keep=self.KEEP,
|
| 1026 |
+
trim_tokens_to_summarize=None,
|
| 1027 |
+
)
|
| 1028 |
+
reset_turn_signals("eq-xml")
|
| 1029 |
+
xml_update = xml_middleware.before_model(
|
| 1030 |
+
{"messages": list(history)},
|
| 1031 |
+
SimpleNamespace(
|
| 1032 |
+
context=SimpleNamespace(kb_session_id="eq-xml", cache_user_id="")
|
| 1033 |
+
),
|
| 1034 |
+
)
|
| 1035 |
+
pop_turn_events("eq-xml")
|
| 1036 |
+
pop_turn_signals("eq-xml")
|
| 1037 |
+
|
| 1038 |
+
structured_model = ExperimentCompactionMiddlewareTests.StructuredFakeModel(
|
| 1039 |
+
["structured summary"]
|
| 1040 |
+
)
|
| 1041 |
+
structured_middleware = PrefixPreservingCompactionMiddleware(
|
| 1042 |
+
model=structured_model,
|
| 1043 |
+
trigger=self.TRIGGER,
|
| 1044 |
+
keep=self.KEEP,
|
| 1045 |
+
trim_tokens_to_summarize=None,
|
| 1046 |
+
)
|
| 1047 |
+
request = ExperimentCompactionMiddlewareTests._structured_request(
|
| 1048 |
+
structured_model, list(history), "eq-structured"
|
| 1049 |
+
)
|
| 1050 |
+
handled: list[ModelRequest] = []
|
| 1051 |
+
|
| 1052 |
+
def handler(compacted_request):
|
| 1053 |
+
handled.append(compacted_request)
|
| 1054 |
+
return ModelResponse(result=[AIMessage(content="answer", id="a")])
|
| 1055 |
+
|
| 1056 |
+
reset_turn_signals("eq-structured")
|
| 1057 |
+
result = structured_middleware.wrap_model_call(request, handler)
|
| 1058 |
+
pop_turn_events("eq-structured")
|
| 1059 |
+
pop_turn_signals("eq-structured")
|
| 1060 |
+
|
| 1061 |
+
self.assertIsInstance(xml_update["messages"][0], RemoveMessage)
|
| 1062 |
+
self.assertEqual(xml_update["messages"][0].id, REMOVE_ALL_MESSAGES)
|
| 1063 |
+
xml_summary, xml_tail = xml_update["messages"][1], xml_update["messages"][2:]
|
| 1064 |
+
compacted = handled[0].messages
|
| 1065 |
+
structured_summary, structured_tail = compacted[0], compacted[1:]
|
| 1066 |
+
|
| 1067 |
+
for summary in (xml_summary, structured_summary):
|
| 1068 |
+
self.assertIsInstance(summary, HumanMessage)
|
| 1069 |
+
self.assertEqual(
|
| 1070 |
+
summary.additional_kwargs.get("lc_source"), "summarization"
|
| 1071 |
+
)
|
| 1072 |
+
self.assertTrue(
|
| 1073 |
+
summary.content.startswith(
|
| 1074 |
+
"Here is a summary of the conversation to date:"
|
| 1075 |
+
)
|
| 1076 |
+
)
|
| 1077 |
+
self.assertEqual([m.id for m in xml_tail], [m.id for m in structured_tail])
|
| 1078 |
+
self.assertEqual(
|
| 1079 |
+
[m.content for m in xml_tail], [m.content for m in structured_tail]
|
| 1080 |
+
)
|
| 1081 |
+
# Only the summary text differs between the two paths.
|
| 1082 |
+
self.assertIn("xml summary", xml_summary.content)
|
| 1083 |
+
self.assertIn("structured summary", structured_summary.content)
|
| 1084 |
+
# The structured path's checkpoint command installs the same structure.
|
| 1085 |
+
command_messages = result.command.update["messages"]
|
| 1086 |
+
self.assertIsInstance(command_messages[0], RemoveMessage)
|
| 1087 |
+
self.assertIs(command_messages[1], structured_summary)
|
| 1088 |
+
self.assertEqual(
|
| 1089 |
+
[m.id for m in command_messages[2:-1]],
|
| 1090 |
+
[m.id for m in structured_tail],
|
| 1091 |
+
)
|
| 1092 |
+
self.assertEqual(command_messages[-1].content, "answer")
|
| 1093 |
+
|
| 1094 |
+
|
| 1095 |
+
class MultiCompactionTests(unittest.TestCase):
|
| 1096 |
+
def test_xml_second_compaction_replaces_prior_summary_without_orphans(
|
| 1097 |
+
self,
|
| 1098 |
+
) -> None:
|
| 1099 |
+
model = ExperimentCompactionMiddlewareTests.FakeModel(
|
| 1100 |
+
["first summary", "second summary"]
|
| 1101 |
+
)
|
| 1102 |
+
middleware = InstrumentedSummarizationMiddleware(
|
| 1103 |
+
model=model,
|
| 1104 |
+
trigger=("tokens", 1_000),
|
| 1105 |
+
keep=("tokens", 500),
|
| 1106 |
+
trim_tokens_to_summarize=None,
|
| 1107 |
+
)
|
| 1108 |
+
runtime = SimpleNamespace(
|
| 1109 |
+
context=SimpleNamespace(kb_session_id="multi-xml", cache_user_id="")
|
| 1110 |
+
)
|
| 1111 |
+
history = [
|
| 1112 |
+
(HumanMessage if index % 2 == 0 else AIMessage)(
|
| 1113 |
+
content=f"m{index}:" + "x" * 1_000, id=f"m{index}"
|
| 1114 |
+
)
|
| 1115 |
+
for index in range(24)
|
| 1116 |
+
]
|
| 1117 |
+
reset_turn_signals("multi-xml")
|
| 1118 |
+
first_update = middleware.before_model({"messages": history}, runtime)
|
| 1119 |
+
self.assertIsNotNone(first_update)
|
| 1120 |
+
summarized_state = list(first_update["messages"][1:])
|
| 1121 |
+
|
| 1122 |
+
second_turn = [
|
| 1123 |
+
HumanMessage(content="new question " + "y" * 3_000, id="n0"),
|
| 1124 |
+
AIMessage(
|
| 1125 |
+
content="",
|
| 1126 |
+
id="n1",
|
| 1127 |
+
tool_calls=[{"name": "lookup", "args": {}, "id": "call-2"}],
|
| 1128 |
+
),
|
| 1129 |
+
ToolMessage(
|
| 1130 |
+
content="evidence " + "y" * 1_000, tool_call_id="call-2", id="n2"
|
| 1131 |
+
),
|
| 1132 |
+
AIMessage(content="answer two", id="n3"),
|
| 1133 |
+
]
|
| 1134 |
+
second_update = middleware.before_model(
|
| 1135 |
+
{"messages": [*summarized_state, *second_turn]}, runtime
|
| 1136 |
+
)
|
| 1137 |
+
self.assertIsNotNone(second_update)
|
| 1138 |
+
|
| 1139 |
+
final_state = list(second_update["messages"][1:])
|
| 1140 |
+
summaries = [
|
| 1141 |
+
message
|
| 1142 |
+
for message in final_state
|
| 1143 |
+
if message.additional_kwargs.get("lc_source") == "summarization"
|
| 1144 |
+
]
|
| 1145 |
+
self.assertEqual(len(summaries), 1)
|
| 1146 |
+
self.assertIn("second summary", summaries[0].content)
|
| 1147 |
+
self.assertNotIn("first summary", summaries[0].content)
|
| 1148 |
+
# The first summary fed the second summarization instead of surviving.
|
| 1149 |
+
self.assertIn("first summary", model.prompts[1])
|
| 1150 |
+
for index, message in enumerate(final_state):
|
| 1151 |
+
if isinstance(message, ToolMessage):
|
| 1152 |
+
self.assertGreater(index, 0)
|
| 1153 |
+
previous = final_state[index - 1]
|
| 1154 |
+
self.assertIsInstance(previous, AIMessage)
|
| 1155 |
+
self.assertIn(
|
| 1156 |
+
message.tool_call_id,
|
| 1157 |
+
[call["id"] for call in previous.tool_calls],
|
| 1158 |
+
)
|
| 1159 |
+
self.assertLessEqual(middleware._partial_token_counter(final_state[1:]), 500)
|
| 1160 |
+
self.assertEqual(len(pop_turn_events("multi-xml")), 2)
|
| 1161 |
+
self.assertEqual(pop_turn_signals("multi-xml")["compactions_this_turn"], 2)
|
| 1162 |
+
|
| 1163 |
+
def test_real_agent_structured_second_compaction_on_summarized_thread(
|
| 1164 |
+
self,
|
| 1165 |
+
) -> None:
|
| 1166 |
+
def response(text: str) -> AIMessage:
|
| 1167 |
+
return AIMessage(
|
| 1168 |
+
content=text,
|
| 1169 |
+
usage_metadata={
|
| 1170 |
+
"input_tokens": 1_000,
|
| 1171 |
+
"output_tokens": 10,
|
| 1172 |
+
"total_tokens": 1_010,
|
| 1173 |
+
"input_token_details": {"cache_read": 0},
|
| 1174 |
+
},
|
| 1175 |
+
response_metadata={"model_name": "deepseek-v4-flash"},
|
| 1176 |
+
)
|
| 1177 |
+
|
| 1178 |
+
model = FakeMessagesListChatModel(
|
| 1179 |
+
responses=[
|
| 1180 |
+
response("checkpoint one"),
|
| 1181 |
+
response("answer one"),
|
| 1182 |
+
response("checkpoint two"),
|
| 1183 |
+
response("answer two"),
|
| 1184 |
+
]
|
| 1185 |
+
)
|
| 1186 |
+
compactor = PrefixPreservingCompactionMiddleware(
|
| 1187 |
+
model=model,
|
| 1188 |
+
trigger=("tokens", 1_000),
|
| 1189 |
+
keep=("tokens", 500),
|
| 1190 |
+
trim_tokens_to_summarize=None,
|
| 1191 |
+
summary_input_guard_tokens=900_000,
|
| 1192 |
+
)
|
| 1193 |
+
agent = create_agent(
|
| 1194 |
+
model=model,
|
| 1195 |
+
tools=[],
|
| 1196 |
+
system_prompt="stable system",
|
| 1197 |
+
middleware=[compactor],
|
| 1198 |
+
checkpointer=InMemorySaver(),
|
| 1199 |
+
)
|
| 1200 |
+
config = {"configurable": {"thread_id": "structured-multi"}}
|
| 1201 |
+
first_turn = [
|
| 1202 |
+
(HumanMessage if index % 2 == 0 else AIMessage)(
|
| 1203 |
+
content=f"old-{index}:" + "x" * 1_000
|
| 1204 |
+
)
|
| 1205 |
+
for index in range(28)
|
| 1206 |
+
]
|
| 1207 |
+
reset_turn_signals("structured-multi-t1")
|
| 1208 |
+
agent.invoke(
|
| 1209 |
+
{"messages": first_turn},
|
| 1210 |
+
config=config,
|
| 1211 |
+
context=AppContext(allowed_sources=(), kb_session_id="structured-multi-t1"),
|
| 1212 |
+
)
|
| 1213 |
+
state = agent.get_state(config).values["messages"]
|
| 1214 |
+
first_summaries = [
|
| 1215 |
+
message
|
| 1216 |
+
for message in state
|
| 1217 |
+
if message.additional_kwargs.get("lc_source") == "summarization"
|
| 1218 |
+
]
|
| 1219 |
+
self.assertEqual(len(first_summaries), 1)
|
| 1220 |
+
self.assertIn("checkpoint one", first_summaries[0].content)
|
| 1221 |
+
self.assertEqual(len(pop_turn_events("structured-multi-t1")), 1)
|
| 1222 |
+
self.assertEqual(
|
| 1223 |
+
pop_turn_signals("structured-multi-t1")["compactions_this_turn"], 1
|
| 1224 |
+
)
|
| 1225 |
+
|
| 1226 |
+
reset_turn_signals("structured-multi-t2")
|
| 1227 |
+
agent.invoke(
|
| 1228 |
+
{"messages": [HumanMessage(content="second wave " + "y" * 6_000)]},
|
| 1229 |
+
config=config,
|
| 1230 |
+
context=AppContext(allowed_sources=(), kb_session_id="structured-multi-t2"),
|
| 1231 |
+
)
|
| 1232 |
+
state = agent.get_state(config).values["messages"]
|
| 1233 |
+
summaries = [
|
| 1234 |
+
message
|
| 1235 |
+
for message in state
|
| 1236 |
+
if message.additional_kwargs.get("lc_source") == "summarization"
|
| 1237 |
+
]
|
| 1238 |
+
self.assertEqual(len(summaries), 1)
|
| 1239 |
+
self.assertIn("checkpoint two", summaries[0].content)
|
| 1240 |
+
contents = [str(message.content) for message in state]
|
| 1241 |
+
self.assertFalse(any("checkpoint one" in content for content in contents))
|
| 1242 |
+
self.assertFalse(any("old-0:" in content for content in contents))
|
| 1243 |
+
self.assertFalse(any(isinstance(m, ToolMessage) for m in state))
|
| 1244 |
+
# Retained tail survives verbatim, followed by the new answer.
|
| 1245 |
+
self.assertTrue(any(content.startswith("second wave") for content in contents))
|
| 1246 |
+
self.assertEqual(contents.count("answer two"), 1)
|
| 1247 |
+
event = pop_turn_events("structured-multi-t2")[0]
|
| 1248 |
+
self.assertEqual(event["summary_strategy"], "structured_prefix")
|
| 1249 |
+
self.assertEqual(event["summary_instruction_retained_messages"], 1)
|
| 1250 |
+
self.assertEqual(
|
| 1251 |
+
pop_turn_signals("structured-multi-t2")["compactions_this_turn"], 1
|
| 1252 |
+
)
|
| 1253 |
+
|
| 1254 |
+
|
| 1255 |
+
class TurnUsageAccountingInvariantTests(unittest.TestCase):
|
| 1256 |
+
def test_model_call_rows_sum_to_the_billed_usage_totals(self) -> None:
|
| 1257 |
+
# est_cost_usd is computed from usage_by_model; the per-call rows are
|
| 1258 |
+
# the explanation. If a call were double-counted (or dropped) on either
|
| 1259 |
+
# side, the two aggregates would disagree.
|
| 1260 |
+
def response(
|
| 1261 |
+
text: str, input_tokens: int, cache_read: int, cache_creation: int
|
| 1262 |
+
) -> AIMessage:
|
| 1263 |
+
return AIMessage(
|
| 1264 |
+
content=text,
|
| 1265 |
+
usage_metadata={
|
| 1266 |
+
"input_tokens": input_tokens,
|
| 1267 |
+
"output_tokens": 40,
|
| 1268 |
+
"total_tokens": input_tokens + 40,
|
| 1269 |
+
"input_token_details": {
|
| 1270 |
+
"cache_read": cache_read,
|
| 1271 |
+
"cache_creation": cache_creation,
|
| 1272 |
+
},
|
| 1273 |
+
},
|
| 1274 |
+
response_metadata={"model_name": "deepseek-v4-flash"},
|
| 1275 |
+
)
|
| 1276 |
+
|
| 1277 |
+
model = FakeMessagesListChatModel(
|
| 1278 |
+
responses=[
|
| 1279 |
+
response("summary", 6_000, 5_000, 500),
|
| 1280 |
+
response("answer", 700, 100, 50),
|
| 1281 |
+
]
|
| 1282 |
+
)
|
| 1283 |
+
summary = InstrumentedSummarizationMiddleware(
|
| 1284 |
+
model=model,
|
| 1285 |
+
trigger=("tokens", 1_000),
|
| 1286 |
+
keep=("tokens", 500),
|
| 1287 |
+
trim_tokens_to_summarize=None,
|
| 1288 |
+
)
|
| 1289 |
+
agent = create_agent(model=model, tools=[], middleware=[summary])
|
| 1290 |
+
messages = [
|
| 1291 |
+
(HumanMessage if index % 2 == 0 else AIMessage)(content="x" * 1_000)
|
| 1292 |
+
for index in range(28)
|
| 1293 |
+
]
|
| 1294 |
+
handler = TurnUsageHandler()
|
| 1295 |
+
reset_turn_signals("usage-invariant")
|
| 1296 |
+
agent.invoke(
|
| 1297 |
+
{"messages": messages},
|
| 1298 |
+
config={"callbacks": [handler]},
|
| 1299 |
+
context=AppContext(allowed_sources=(), kb_session_id="usage-invariant"),
|
| 1300 |
+
)
|
| 1301 |
+
pop_turn_events("usage-invariant")
|
| 1302 |
+
pop_turn_signals("usage-invariant")
|
| 1303 |
+
|
| 1304 |
+
self.assertEqual(handler.llm_calls, 2)
|
| 1305 |
+
self.assertEqual(len(handler.model_calls), 2)
|
| 1306 |
+
self.assertEqual(
|
| 1307 |
+
sorted(call["source"] for call in handler.model_calls),
|
| 1308 |
+
["agent", "summarization"],
|
| 1309 |
+
)
|
| 1310 |
+
totals = usage_totals(handler.usage_metadata)
|
| 1311 |
+
summed = {
|
| 1312 |
+
field: sum(call[field] for call in handler.model_calls)
|
| 1313 |
+
for field in (
|
| 1314 |
+
"input_tokens",
|
| 1315 |
+
"output_tokens",
|
| 1316 |
+
"total_tokens",
|
| 1317 |
+
"cache_read_tokens",
|
| 1318 |
+
"cache_creation_tokens",
|
| 1319 |
+
)
|
| 1320 |
+
}
|
| 1321 |
+
self.assertEqual(summed, totals)
|
| 1322 |
+
# Known scripted usage pins the absolute numbers, not just consistency.
|
| 1323 |
+
self.assertEqual(totals["input_tokens"], 6_700)
|
| 1324 |
+
self.assertEqual(totals["output_tokens"], 80)
|
| 1325 |
+
self.assertEqual(totals["cache_read_tokens"], 5_100)
|
| 1326 |
+
self.assertEqual(totals["cache_creation_tokens"], 550)
|
| 1327 |
+
estimated = estimate_cost_usd(handler.usage_metadata)
|
| 1328 |
+
self.assertIsNotNone(estimated)
|
| 1329 |
+
self.assertAlmostEqual(
|
| 1330 |
+
estimated,
|
| 1331 |
+
sum(call["cost"]["total_usd"] for call in handler.model_calls),
|
| 1332 |
+
)
|
| 1333 |
+
|
| 1334 |
+
|
| 1335 |
class PromptCompressionTests(unittest.TestCase):
|
| 1336 |
def test_collapses_whitespace_and_reports(self) -> None:
|
| 1337 |
reset_turn_signals("t1")
|
|
@@ -1,6 +1,7 @@
|
|
| 1 |
from __future__ import annotations
|
| 2 |
|
| 3 |
import unittest
|
|
|
|
| 4 |
|
| 5 |
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
|
| 6 |
from langchain_core.outputs import ChatGeneration, LLMResult
|
|
@@ -11,9 +12,12 @@ from app.telemetry import (
|
|
| 11 |
TurnUsageHandler,
|
| 12 |
context_window_stats,
|
| 13 |
estimate_cost_usd,
|
|
|
|
|
|
|
| 14 |
pop_turn_signals,
|
| 15 |
pricing_for_model,
|
| 16 |
record_turn_signal,
|
|
|
|
| 17 |
record_turn_signal_max,
|
| 18 |
reset_turn_signals,
|
| 19 |
usage_totals,
|
|
@@ -93,6 +97,27 @@ class CostEstimateTests(unittest.TestCase):
|
|
| 93 |
}
|
| 94 |
self.assertIsNone(estimate_cost_usd(mixed))
|
| 95 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 96 |
|
| 97 |
class ContextWindowStatsTests(unittest.TestCase):
|
| 98 |
def test_counts_summaries_and_cleared_tool_outputs(self) -> None:
|
|
@@ -145,6 +170,108 @@ class TurnUsageHandlerTests(unittest.TestCase):
|
|
| 145 |
self.assertEqual(usage["input_tokens"], 150)
|
| 146 |
self.assertEqual(usage["output_tokens"], 30)
|
| 147 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 148 |
|
| 149 |
class TurnSignalRegistryTests(unittest.TestCase):
|
| 150 |
def test_accumulates_and_pops_per_turn(self) -> None:
|
|
@@ -158,6 +285,15 @@ class TurnSignalRegistryTests(unittest.TestCase):
|
|
| 158 |
# Popping clears the entry: a second pop is empty.
|
| 159 |
self.assertEqual(pop_turn_signals("turn-a"), {})
|
| 160 |
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 161 |
def test_turns_are_isolated_and_noops_are_ignored(self) -> None:
|
| 162 |
reset_turn_signals("turn-x")
|
| 163 |
reset_turn_signals("turn-y")
|
|
|
|
| 1 |
from __future__ import annotations
|
| 2 |
|
| 3 |
import unittest
|
| 4 |
+
from uuid import uuid4
|
| 5 |
|
| 6 |
from langchain_core.messages import AIMessage, HumanMessage, ToolMessage
|
| 7 |
from langchain_core.outputs import ChatGeneration, LLMResult
|
|
|
|
| 12 |
TurnUsageHandler,
|
| 13 |
context_window_stats,
|
| 14 |
estimate_cost_usd,
|
| 15 |
+
aggregate_cost_breakdown,
|
| 16 |
+
pop_turn_events,
|
| 17 |
pop_turn_signals,
|
| 18 |
pricing_for_model,
|
| 19 |
record_turn_signal,
|
| 20 |
+
record_turn_event,
|
| 21 |
record_turn_signal_max,
|
| 22 |
reset_turn_signals,
|
| 23 |
usage_totals,
|
|
|
|
| 97 |
}
|
| 98 |
self.assertIsNone(estimate_cost_usd(mixed))
|
| 99 |
|
| 100 |
+
def test_deepseek_cost_breakdown_is_mutually_exclusive(self) -> None:
|
| 101 |
+
usage = {
|
| 102 |
+
"deepseek-v4-flash": {
|
| 103 |
+
"input_tokens": 1_000_000,
|
| 104 |
+
"output_tokens": 100_000,
|
| 105 |
+
"input_token_details": {"cache_read": 900_000},
|
| 106 |
+
}
|
| 107 |
+
}
|
| 108 |
+
breakdown = aggregate_cost_breakdown(usage)
|
| 109 |
+
self.assertIsNotNone(breakdown)
|
| 110 |
+
self.assertAlmostEqual(breakdown["cache_miss_input_usd"], 0.014)
|
| 111 |
+
self.assertAlmostEqual(breakdown["cache_read_input_usd"], 0.00252)
|
| 112 |
+
self.assertAlmostEqual(breakdown["output_usd"], 0.028)
|
| 113 |
+
self.assertAlmostEqual(
|
| 114 |
+
breakdown["total_usd"],
|
| 115 |
+
breakdown["cache_miss_input_usd"]
|
| 116 |
+
+ breakdown["cache_read_input_usd"]
|
| 117 |
+
+ breakdown["cache_creation_input_usd"]
|
| 118 |
+
+ breakdown["output_usd"],
|
| 119 |
+
)
|
| 120 |
+
|
| 121 |
|
| 122 |
class ContextWindowStatsTests(unittest.TestCase):
|
| 123 |
def test_counts_summaries_and_cleared_tool_outputs(self) -> None:
|
|
|
|
| 170 |
self.assertEqual(usage["input_tokens"], 150)
|
| 171 |
self.assertEqual(usage["output_tokens"], 30)
|
| 172 |
|
| 173 |
+
def test_records_one_explanatory_row_per_call(self) -> None:
|
| 174 |
+
handler = TurnUsageHandler()
|
| 175 |
+
run_id = uuid4()
|
| 176 |
+
handler.on_chat_model_start(
|
| 177 |
+
{},
|
| 178 |
+
[[HumanMessage(content="hello")]],
|
| 179 |
+
run_id=run_id,
|
| 180 |
+
metadata={"lc_source": "summarization"},
|
| 181 |
+
)
|
| 182 |
+
message = AIMessage(
|
| 183 |
+
content="summary",
|
| 184 |
+
usage_metadata={
|
| 185 |
+
"input_tokens": 100,
|
| 186 |
+
"output_tokens": 20,
|
| 187 |
+
"total_tokens": 120,
|
| 188 |
+
"input_token_details": {"cache_read": 80},
|
| 189 |
+
},
|
| 190 |
+
response_metadata={"model_name": "deepseek-v4-flash"},
|
| 191 |
+
)
|
| 192 |
+
handler.on_llm_end(
|
| 193 |
+
LLMResult(generations=[[ChatGeneration(message=message)]]),
|
| 194 |
+
run_id=run_id,
|
| 195 |
+
)
|
| 196 |
+
call = handler.model_calls[0]
|
| 197 |
+
self.assertEqual(call["source"], "summarization")
|
| 198 |
+
self.assertEqual(call["cache_read_tokens"], 80)
|
| 199 |
+
self.assertEqual(call["cache_miss_tokens"], 20)
|
| 200 |
+
self.assertTrue(call["cache_details_reported"])
|
| 201 |
+
self.assertGreater(call["request_context_tokens_approx"], 0)
|
| 202 |
+
self.assertAlmostEqual(
|
| 203 |
+
call["cost"]["total_usd"],
|
| 204 |
+
estimate_cost_usd({"deepseek-v4-flash": message.usage_metadata}),
|
| 205 |
+
)
|
| 206 |
+
|
| 207 |
+
|
| 208 |
+
class LangchainOpenAICacheFieldContractTests(unittest.TestCase):
|
| 209 |
+
"""Pin the langchain-openai usage conversion our cache accounting rides on.
|
| 210 |
+
|
| 211 |
+
DeepSeek (and OpenAI) report cached prompt tokens as
|
| 212 |
+
``prompt_tokens_details.cached_tokens``; langchain-openai must surface that
|
| 213 |
+
as ``usage_metadata.input_token_details.cache_read`` or TurnUsageHandler's
|
| 214 |
+
cache buckets (and the ~50x DeepSeek cache-read discount) silently read 0.
|
| 215 |
+
"""
|
| 216 |
+
|
| 217 |
+
def _convert(self, payload: dict) -> dict:
|
| 218 |
+
try:
|
| 219 |
+
from langchain_openai.chat_models.base import _create_usage_metadata
|
| 220 |
+
except ImportError as exc:
|
| 221 |
+
self.fail(
|
| 222 |
+
"langchain_openai.chat_models.base._create_usage_metadata is no "
|
| 223 |
+
f"longer importable ({exc}). A langchain-openai upgrade moved "
|
| 224 |
+
"the usage conversion; re-pin the prompt_tokens_details."
|
| 225 |
+
"cached_tokens -> input_token_details.cache_read mapping "
|
| 226 |
+
"against its new location."
|
| 227 |
+
)
|
| 228 |
+
return _create_usage_metadata(payload)
|
| 229 |
+
|
| 230 |
+
def test_cached_tokens_map_to_cache_read_details(self) -> None:
|
| 231 |
+
usage = self._convert(
|
| 232 |
+
{
|
| 233 |
+
"prompt_tokens": 1_000,
|
| 234 |
+
"completion_tokens": 100,
|
| 235 |
+
"total_tokens": 1_100,
|
| 236 |
+
"prompt_tokens_details": {"cached_tokens": 900},
|
| 237 |
+
}
|
| 238 |
+
)
|
| 239 |
+
# LangChain convention: input_tokens INCLUDES the cached bucket; the
|
| 240 |
+
# cost code carves cache_read out instead of adding it on top.
|
| 241 |
+
self.assertEqual(usage["input_tokens"], 1_000)
|
| 242 |
+
self.assertEqual(usage["output_tokens"], 100)
|
| 243 |
+
self.assertEqual(usage["input_token_details"]["cache_read"], 900)
|
| 244 |
+
|
| 245 |
+
def test_converted_usage_flows_through_turn_usage_handler(self) -> None:
|
| 246 |
+
usage = self._convert(
|
| 247 |
+
{
|
| 248 |
+
"prompt_tokens": 1_000,
|
| 249 |
+
"completion_tokens": 100,
|
| 250 |
+
"total_tokens": 1_100,
|
| 251 |
+
"prompt_tokens_details": {"cached_tokens": 900},
|
| 252 |
+
}
|
| 253 |
+
)
|
| 254 |
+
handler = TurnUsageHandler()
|
| 255 |
+
run_id = uuid4()
|
| 256 |
+
handler.on_chat_model_start({}, [[HumanMessage(content="q")]], run_id=run_id)
|
| 257 |
+
message = AIMessage(
|
| 258 |
+
content="answer",
|
| 259 |
+
usage_metadata=usage,
|
| 260 |
+
response_metadata={"model_name": "deepseek-v4-flash"},
|
| 261 |
+
)
|
| 262 |
+
handler.on_llm_end(
|
| 263 |
+
LLMResult(generations=[[ChatGeneration(message=message)]]),
|
| 264 |
+
run_id=run_id,
|
| 265 |
+
)
|
| 266 |
+
call = handler.model_calls[0]
|
| 267 |
+
self.assertEqual(call["cache_read_tokens"], 900)
|
| 268 |
+
self.assertEqual(call["cache_miss_tokens"], 100)
|
| 269 |
+
self.assertTrue(call["cache_details_reported"])
|
| 270 |
+
# deepseek-v4-flash: $0.14 miss / $0.0028 cache-read / $0.28 output
|
| 271 |
+
# per MTok, so the cache discount must show up in the billed cost.
|
| 272 |
+
expected = (100 * 0.14 + 900 * 0.0028 + 100 * 0.28) / 1_000_000
|
| 273 |
+
self.assertAlmostEqual(call["cost"]["total_usd"], expected)
|
| 274 |
+
|
| 275 |
|
| 276 |
class TurnSignalRegistryTests(unittest.TestCase):
|
| 277 |
def test_accumulates_and_pops_per_turn(self) -> None:
|
|
|
|
| 285 |
# Popping clears the entry: a second pop is empty.
|
| 286 |
self.assertEqual(pop_turn_signals("turn-a"), {})
|
| 287 |
|
| 288 |
+
def test_structured_events_are_isolated_and_popped(self) -> None:
|
| 289 |
+
reset_turn_signals("turn-events")
|
| 290 |
+
record_turn_event("turn-events", {"event": "summarization", "tokens": 9})
|
| 291 |
+
self.assertEqual(
|
| 292 |
+
pop_turn_events("turn-events"),
|
| 293 |
+
[{"event": "summarization", "tokens": 9}],
|
| 294 |
+
)
|
| 295 |
+
self.assertEqual(pop_turn_events("turn-events"), [])
|
| 296 |
+
|
| 297 |
def test_turns_are_isolated_and_noops_are_ignored(self) -> None:
|
| 298 |
reset_turn_signals("turn-x")
|
| 299 |
reset_turn_signals("turn-y")
|