omarsol Claude Fable 5 commited on
Commit
6b117e7
·
1 Parent(s): 93642db

feat(experiments): DeepSeek stage-1 compaction arms + prefix-preserving summarization

Browse files

The 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 CHANGED
@@ -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(query: str, runtime: ToolRuntime[AppContext]) -> str:
 
 
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 = 8,
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=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": ("messages", memory_config.summarization_keep_messages),
 
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
- middleware.append(SummarizationMiddleware(**summarization_kwargs))
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
- **pop_turn_signals(message_id),
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(
app/chat_types.py CHANGED
@@ -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
app/memory_presets.py CHANGED
@@ -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(
app/telemetry.py CHANGED
@@ -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
- _TURN_SIGNALS.pop(next(iter(_TURN_SIGNALS)), None)
 
 
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 per-model usage for one turn and count chat-model calls."""
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
- total = 0.0
195
- for model_key, usage in usage_by_model.items():
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(
tests/test_chat_service.py CHANGED
@@ -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",
tests/test_memory_presets.py CHANGED
@@ -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 = []
tests/test_memory_variants.py CHANGED
@@ -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 langchain_core.messages import AIMessage, HumanMessage, ToolMessage
 
 
 
 
 
 
 
 
 
 
 
 
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 COMPACTION_SIGNAL_NAMES, pop_turn_signals, reset_turn_signals
 
 
 
 
 
 
 
 
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")
tests/test_telemetry.py CHANGED
@@ -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")