test_context_budget.py 6.7 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211
  1. import math
  2. import pytest
  3. from agent.core.prompts import build_summary_header
  4. from agent.core.runner import AgentRunner, RunConfig
  5. from agent.trace.compaction import (
  6. CompressionConfig,
  7. calibrate_prompt_estimate,
  8. measure_prompt_tokens,
  9. )
  10. from agent.trace.models import Trace
  11. from agent.trace.store import FileSystemTraceStore
  12. def _knowledge_off():
  13. from agent.tools.builtin.knowledge import KnowledgeConfig
  14. return KnowledgeConfig(
  15. enable_extraction=False,
  16. enable_completion_extraction=False,
  17. enable_injection=False,
  18. )
  19. def test_prompt_measurement_includes_tools_and_provider_calibration():
  20. config = CompressionConfig(max_tokens=100_000)
  21. messages = [{"role": "user", "content": "x" * 4_000}]
  22. schemas = [
  23. {
  24. "type": "function",
  25. "function": {
  26. "name": "large_tool",
  27. "description": "y" * 4_000,
  28. "parameters": {"type": "object"},
  29. },
  30. }
  31. ]
  32. raw = measure_prompt_tokens(messages, [], config, "unknown-model", 1.0)
  33. measured = measure_prompt_tokens(messages, schemas, config, "unknown-model", 1.0)
  34. factor = calibrate_prompt_estimate(
  35. actual_prompt_tokens=measured.estimated_tokens * 2,
  36. estimated_prompt_tokens=measured.estimated_tokens,
  37. previous_factor=1.25,
  38. )
  39. calibrated = measure_prompt_tokens(
  40. messages,
  41. schemas,
  42. config,
  43. "unknown-model",
  44. factor,
  45. )
  46. assert measured.tool_schema_tokens > 0
  47. assert measured.estimated_tokens > raw.estimated_tokens
  48. assert factor == 2
  49. assert calibrated.calibrated_tokens == measured.estimated_tokens * 2
  50. assert calibrated.trigger_tokens == 80_000
  51. assert calibrated.target_tokens == 30_000
  52. @pytest.mark.asyncio
  53. async def test_build_900036_growth_replay_compresses_before_hard_limit(tmp_path):
  54. store = FileSystemTraceStore(str(tmp_path))
  55. trace = Trace(trace_id="growth-replay", mode="agent", agent_role="legacy")
  56. await store.create_trace(trace)
  57. runner = AgentRunner(trace_store=store)
  58. config = RunConfig(
  59. model="unregistered-build-model",
  60. compression=CompressionConfig(max_tokens=100_000),
  61. knowledge=_knowledge_off(),
  62. )
  63. schemas = [
  64. {
  65. "type": "function",
  66. "function": {
  67. "name": "offline_step",
  68. "description": "append one deterministic fixture step",
  69. "parameters": {"type": "object", "properties": {}},
  70. },
  71. }
  72. ]
  73. history = [
  74. {"role": "system", "content": "keep system policy"},
  75. {"role": "user", "content": "offline 97-round context fixture"},
  76. ]
  77. total_prompt_tokens = 0
  78. max_prompt_tokens = 0
  79. for turn in range(97):
  80. history.extend(
  81. [
  82. {"role": "assistant", "content": f"turn {turn}", "tool_calls": []},
  83. {"role": "tool", "content": "x" * 4_000},
  84. ]
  85. )
  86. history, _, _, needs_compression = await runner._manage_context_usage(
  87. trace.trace_id,
  88. history,
  89. None,
  90. config,
  91. sequence=turn + 1,
  92. head_seq=0,
  93. tool_schemas=schemas,
  94. )
  95. current = await store.get_trace(trace.trace_id)
  96. measurement = measure_prompt_tokens(
  97. history,
  98. schemas,
  99. config.compression,
  100. config.model,
  101. current.runtime_state.get("context_budget", {}).get("calibration_factor"),
  102. )
  103. simulated_actual = math.ceil(measurement.estimated_tokens * 1.10)
  104. total_prompt_tokens += simulated_actual
  105. max_prompt_tokens = max(max_prompt_tokens, simulated_actual)
  106. await runner._record_prompt_measurement(
  107. current,
  108. config,
  109. history,
  110. schemas,
  111. simulated_actual,
  112. )
  113. if needs_compression:
  114. history = [
  115. history[0],
  116. history[1],
  117. {"role": "user", "content": "compact execution summary"},
  118. ]
  119. current = await store.get_trace(trace.trace_id)
  120. compacted = measure_prompt_tokens(
  121. history,
  122. schemas,
  123. config.compression,
  124. config.model,
  125. current.runtime_state["context_budget"]["calibration_factor"],
  126. )
  127. failure = await runner._finish_context_compression(
  128. current,
  129. compacted,
  130. before_message_count=turn * 2 + 4,
  131. after_message_count=len(history),
  132. )
  133. assert failure is None
  134. assert compacted.calibrated_tokens <= compacted.target_tokens
  135. events = await store.get_events(trace.trace_id)
  136. event_names = [event["event"] for event in events]
  137. assert "context_compression_started" in event_names
  138. assert "context_compression_completed" in event_names
  139. assert max_prompt_tokens < 100_000
  140. assert total_prompt_tokens < 6_327_689 * 0.50
  141. usage = runner.get_context_usage(trace.trace_id)
  142. assert usage.compression_count >= 1
  143. assert usage.actual_prompt_tokens is not None
  144. @pytest.mark.asyncio
  145. async def test_oversized_minimum_context_fails_before_model_request(tmp_path):
  146. calls = 0
  147. async def llm_call(**_kwargs):
  148. nonlocal calls
  149. calls += 1
  150. return {"content": "must not be called", "tool_calls": None}
  151. store = FileSystemTraceStore(str(tmp_path))
  152. result = await AgentRunner(
  153. trace_store=store,
  154. llm_call=llm_call,
  155. ).run_result(
  156. [{"role": "user", "content": "x" * 8_000}],
  157. RunConfig(
  158. max_iterations=4,
  159. tools=[],
  160. tool_groups=[],
  161. knowledge=_knowledge_off(),
  162. compression=CompressionConfig(max_tokens=1_000),
  163. ),
  164. )
  165. assert calls == 0
  166. assert result["status"] == "failed"
  167. assert result["failure"]["code"] == "CONTEXT_BUDGET_EXCEEDED"
  168. events = await store.get_events(result["trace_id"])
  169. assert "context_budget_exceeded" in [event["event"] for event in events]
  170. @pytest.mark.parametrize(
  171. "kwargs",
  172. [
  173. {"trigger_ratio": 0.5, "target_ratio": 0.5},
  174. {"trigger_ratio": 1.0},
  175. {"fallback_safety_factor": 0.9},
  176. {"max_tokens": -1},
  177. ],
  178. )
  179. def test_compression_config_rejects_unsafe_budgets(kwargs):
  180. with pytest.raises(ValueError):
  181. CompressionConfig(**kwargs)
  182. def test_compression_summary_forbids_repeating_successful_immutable_reads():
  183. summary = build_summary_header(
  184. "The accepted artifact was read; next save the candidate."
  185. )
  186. assert "不要因为原文已压缩而用相同参数重复读取" in summary
  187. assert "直接执行摘要所列的下一个" in summary