test_compression_policy_persistence.py 9.8 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350
  1. from __future__ import annotations
  2. import pytest
  3. from agent.core.runner import AgentRunner, RunConfig
  4. from agent.tools.builtin.knowledge import KnowledgeConfig
  5. from agent.trace.models import Message, Trace
  6. from agent.trace.store import FileSystemTraceStore
  7. async def _add(store, *, trace_id, role, sequence, parent, content):
  8. message = Message.create(
  9. trace_id=trace_id,
  10. role=role,
  11. sequence=sequence,
  12. parent_sequence=parent,
  13. content=content,
  14. )
  15. await store.add_message(message)
  16. return message
  17. @pytest.mark.asyncio
  18. async def test_level2_compression_carries_active_system_policies_across_reload(
  19. tmp_path,
  20. ):
  21. store = FileSystemTraceStore(str(tmp_path))
  22. await store.create_trace(Trace(trace_id="trace", mode="agent"))
  23. initial = await _add(
  24. store,
  25. trace_id="trace",
  26. role="system",
  27. sequence=1,
  28. parent=None,
  29. content="initial framework policy",
  30. )
  31. first_user = await _add(
  32. store,
  33. trace_id="trace",
  34. role="user",
  35. sequence=2,
  36. parent=1,
  37. content="mission",
  38. )
  39. await _add(
  40. store,
  41. trace_id="trace",
  42. role="assistant",
  43. sequence=3,
  44. parent=2,
  45. content={"text": "old work"},
  46. )
  47. await _add(
  48. store,
  49. trace_id="trace",
  50. role="system",
  51. sequence=4,
  52. parent=3,
  53. content="phase two policy",
  54. )
  55. await _add(
  56. store,
  57. trace_id="trace",
  58. role="user",
  59. sequence=5,
  60. parent=4,
  61. content="continue",
  62. )
  63. await store.update_trace("trace", head_sequence=5)
  64. runner = AgentRunner(trace_store=store)
  65. summary, next_sequence, carried = await runner._persist_compressed_main_path(
  66. trace_id="trace",
  67. original_head_sequence=5,
  68. next_sequence=10,
  69. summary_content="summary one",
  70. )
  71. await store.update_trace("trace", head_sequence=summary.sequence)
  72. assert summary.sequence == 11
  73. assert next_sequence == 12
  74. in_memory = runner._rebuild_history_after_compression(
  75. [
  76. initial.to_llm_dict(),
  77. first_user.to_llm_dict(),
  78. {"role": "assistant", "content": "old work"},
  79. {"role": "system", "content": "phase two policy"},
  80. ],
  81. summary.to_llm_dict(),
  82. carried_system_messages=carried,
  83. )
  84. assert [item["role"] for item in in_memory] == [
  85. "system",
  86. "user",
  87. "system",
  88. "user",
  89. ]
  90. assert in_memory[2]["content"] == "phase two policy"
  91. reloaded = FileSystemTraceStore(str(tmp_path))
  92. trace = await reloaded.get_trace("trace")
  93. main_path = await reloaded.get_main_path_messages("trace", trace.head_sequence)
  94. assert [message.role for message in main_path] == [
  95. "system",
  96. "user",
  97. "system",
  98. "user",
  99. ]
  100. assert [message.content for message in main_path if message.role == "system"] == [
  101. "initial framework policy",
  102. "phase two policy",
  103. ]
  104. @pytest.mark.asyncio
  105. async def test_level2_compression_does_not_duplicate_systems_before_first_user(
  106. tmp_path,
  107. ):
  108. store = FileSystemTraceStore(str(tmp_path))
  109. await store.create_trace(Trace(trace_id="trace", mode="agent"))
  110. parent = None
  111. for sequence, role, content in (
  112. (1, "system", "initial"),
  113. (2, "system", "startup policy"),
  114. (3, "user", "mission"),
  115. (4, "system", "phase two policy"),
  116. (5, "user", "continue"),
  117. ):
  118. await _add(
  119. store,
  120. trace_id="trace",
  121. role=role,
  122. sequence=sequence,
  123. parent=parent,
  124. content=content,
  125. )
  126. parent = sequence
  127. runner = AgentRunner(trace_store=store)
  128. summary, _, carried = await runner._persist_compressed_main_path(
  129. trace_id="trace",
  130. original_head_sequence=5,
  131. next_sequence=10,
  132. summary_content="summary",
  133. )
  134. await store.update_trace("trace", head_sequence=summary.sequence)
  135. path = await store.get_main_path_messages("trace", summary.sequence)
  136. assert [item["content"] for item in carried] == ["phase two policy"]
  137. assert [message.content for message in path if message.role == "system"] == [
  138. "initial",
  139. "startup policy",
  140. "phase two policy",
  141. ]
  142. @pytest.mark.asyncio
  143. async def test_repeated_level2_compression_preserves_policy_order_without_growth(
  144. tmp_path,
  145. ):
  146. store = FileSystemTraceStore(str(tmp_path))
  147. await store.create_trace(Trace(trace_id="trace", mode="agent"))
  148. parent = None
  149. for sequence, role, content in (
  150. (1, "system", "initial"),
  151. (2, "user", "mission"),
  152. (3, "system", "phase one"),
  153. (4, "user", "work"),
  154. ):
  155. await _add(
  156. store,
  157. trace_id="trace",
  158. role=role,
  159. sequence=sequence,
  160. parent=parent,
  161. content=content,
  162. )
  163. parent = sequence
  164. runner = AgentRunner(trace_store=store)
  165. summary_one, _, _ = await runner._persist_compressed_main_path(
  166. trace_id="trace",
  167. original_head_sequence=4,
  168. next_sequence=10,
  169. summary_content="summary one",
  170. )
  171. await _add(
  172. store,
  173. trace_id="trace",
  174. role="system",
  175. sequence=12,
  176. parent=summary_one.sequence,
  177. content="phase two",
  178. )
  179. await _add(
  180. store,
  181. trace_id="trace",
  182. role="user",
  183. sequence=13,
  184. parent=12,
  185. content="more work",
  186. )
  187. summary_two, _, _ = await runner._persist_compressed_main_path(
  188. trace_id="trace",
  189. original_head_sequence=13,
  190. next_sequence=20,
  191. summary_content="summary two",
  192. )
  193. await store.update_trace("trace", head_sequence=summary_two.sequence)
  194. path = await store.get_main_path_messages("trace", summary_two.sequence)
  195. assert [message.content for message in path if message.role == "system"] == [
  196. "initial",
  197. "phase one",
  198. "phase two",
  199. ]
  200. assert [message.role for message in path] == [
  201. "system",
  202. "user",
  203. "system",
  204. "system",
  205. "user",
  206. ]
  207. @pytest.mark.asyncio
  208. async def test_real_runner_level2_branch_persists_policy_across_store_reload(tmp_path):
  209. calls = 0
  210. async def llm_call(**kwargs):
  211. nonlocal calls
  212. calls += 1
  213. messages = kwargs["messages"]
  214. if any(
  215. "[[SUMMARY]]" in str(message.get("content", "")) for message in messages
  216. ):
  217. return {
  218. "content": "[[SUMMARY]] durable compressed state",
  219. "tool_calls": None,
  220. "finish_reason": "stop",
  221. }
  222. return {"content": "done", "tool_calls": None, "finish_reason": "stop"}
  223. store = FileSystemTraceStore(str(tmp_path))
  224. runner = AgentRunner(trace_store=store, llm_call=llm_call)
  225. config = RunConfig(
  226. new_trace_id="real-compression-trace",
  227. name="compression test",
  228. max_iterations=4,
  229. side_branch_max_turns=1,
  230. force_side_branch=["compression"],
  231. goal_compression="none",
  232. knowledge=KnowledgeConfig(
  233. enable_extraction=False,
  234. enable_completion_extraction=False,
  235. enable_injection=False,
  236. ),
  237. )
  238. events = [
  239. event
  240. async for event in runner.run(
  241. [
  242. {"role": "system", "content": "initial framework policy"},
  243. {"role": "user", "content": "mission"},
  244. {"role": "system", "content": "phase two policy"},
  245. ],
  246. config,
  247. )
  248. ]
  249. assert calls >= 2
  250. assert events[-1].status == "completed"
  251. reloaded = FileSystemTraceStore(str(tmp_path))
  252. trace = await reloaded.get_trace("real-compression-trace")
  253. main_path = await reloaded.get_main_path_messages(
  254. trace.trace_id, trace.head_sequence
  255. )
  256. system_contents = [
  257. str(message.content) for message in main_path if message.role == "system"
  258. ]
  259. assert len(system_contents) == 2
  260. assert "initial framework policy" in system_contents[0]
  261. assert "phase two policy" in system_contents[1]
  262. assert any(
  263. message.role == "user" and "durable compressed state" in str(message.content)
  264. for message in main_path
  265. )
  266. @pytest.mark.asyncio
  267. async def test_compression_does_not_reactivate_policy_outside_current_main_path(
  268. tmp_path,
  269. ):
  270. store = FileSystemTraceStore(str(tmp_path))
  271. await store.create_trace(Trace(trace_id="trace", mode="agent"))
  272. await _add(
  273. store,
  274. trace_id="trace",
  275. role="system",
  276. sequence=1,
  277. parent=None,
  278. content="initial",
  279. )
  280. await _add(
  281. store,
  282. trace_id="trace",
  283. role="user",
  284. sequence=2,
  285. parent=1,
  286. content="mission",
  287. )
  288. await _add(
  289. store,
  290. trace_id="trace",
  291. role="system",
  292. sequence=3,
  293. parent=2,
  294. content="active phase policy",
  295. )
  296. side_policy = Message.create(
  297. trace_id="trace",
  298. role="system",
  299. sequence=4,
  300. parent_sequence=3,
  301. branch_type="compression",
  302. branch_id="side",
  303. content="side branch policy must stay inactive",
  304. )
  305. await store.add_message(side_policy)
  306. await _add(
  307. store,
  308. trace_id="trace",
  309. role="user",
  310. sequence=5,
  311. parent=3,
  312. content="main path continues",
  313. )
  314. runner = AgentRunner(trace_store=store)
  315. summary, _, _ = await runner._persist_compressed_main_path(
  316. trace_id="trace",
  317. original_head_sequence=5,
  318. next_sequence=10,
  319. summary_content="summary",
  320. )
  321. path = await store.get_main_path_messages("trace", summary.sequence)
  322. policies = [str(message.content) for message in path if message.role == "system"]
  323. assert policies == ["initial", "active phase policy"]