返回 ViMax
test_agent_loop.py
根目录 / tests / test_agent_loop.py
1 import asyncio
2 import json
3 import tempfile
4 import unittest
5 from copy import deepcopy
6
7 from agent_runtime.context_compactor import ContextCompactor
8 from agent_runtime.llm import AssistantMessage
9 from agent_runtime.loop import AgentLoop
10 from agent_runtime.models import ToolCall, ToolResult
11 from agent_runtime.prompts import PromptBuilder
12 from agent_runtime.session_index import SessionIndex
13 from agent_runtime.tool_executor import ToolExecutor
14 from agent_runtime.tools import ToolArgumentSchema, ToolRegistry, ToolSpec
15
16
17 class FakeLLM:
18 def __init__(self, replies):
19 self.replies = list(replies)
20
21 async def complete(self, messages, tools):
22 return self.replies.pop(0)
23
24
25 class FailingLLM:
26 async def complete(self, messages, tools):
27 raise RuntimeError("provider returned invalid response shape")
28
29
30 class CapturingLLM:
31 def __init__(self, replies):
32 self.replies = list(replies)
33 self.calls = []
34
35 async def complete(self, messages, tools):
36 self.calls.append(deepcopy(messages))
37 return self.replies.pop(0)
38
39
40 class AgentLoopTests(unittest.IsolatedAsyncioTestCase):
41 async def test_no_tool_call_finishes(self):
42 with tempfile.TemporaryDirectory() as tmp:
43 index = SessionIndex(tmp)
44 registry = ToolRegistry([])
45 loop = AgentLoop(index, PromptBuilder(f"{tmp}/prompts", index, registry), registry, ToolExecutor(registry, index), FakeLLM([AssistantMessage(text="done")]))
46 events = [event async for event in loop.stream_events("hi")]
47 self.assertEqual(events[-2]["type"], "done")
48 turn_id = events[0]["turn_id"]
49 self.assertTrue(all(event.get("turn_id") == turn_id for event in events))
50 log_text = (index.logs_dir / "loop_history.jsonl").read_text(encoding="utf-8")
51 self.assertIn("assistant_finished_without_tools", log_text)
52
53
54 async def test_turn_record_follows_session_created_by_tool(self):
55 with tempfile.TemporaryDirectory() as tmp:
56 index = SessionIndex(tmp)
57 old = index.create(idea="old")
58
59 def create_actual(args):
60 record = index.create(idea="actual")
61 return ToolResult("create_actual", True, record["session_id"])
62
63 registry = ToolRegistry([ToolSpec("create_actual", "Create actual session", create_actual, schema={})])
64 llm = FakeLLM([AssistantMessage(tool_calls=[ToolCall(name="create_actual", arguments={})]), AssistantMessage(text="finished")])
65 loop = AgentLoop(index, PromptBuilder(f"{tmp}/prompts", index, registry), registry, ToolExecutor(registry, index), llm)
66 events = [event async for event in loop.stream_events("start new project")]
67 active = index.active()
68 self.assertNotEqual(active["session_id"], old["session_id"])
69 self.assertEqual(len(index.get(active["session_id"])["recent_turn_records"]), 1)
70 self.assertEqual(index.get(old["session_id"])["recent_turn_records"], [])
71 self.assertEqual(events[-1]["session"]["active_session_id"], active["session_id"])
72
73
74 async def test_tool_progress_streams_before_tool_result(self):
75 with tempfile.TemporaryDirectory() as tmp:
76 index = SessionIndex(tmp)
77 release = asyncio.Event()
78
79 async def slow_tool(args, runtime):
80 runtime.emit_progress("started", stage="running")
81 await release.wait()
82 return ToolResult("slow_tool", True, "done")
83
84 registry = ToolRegistry([ToolSpec("slow_tool", "Slow tool", slow_tool, schema={})])
85 llm = FakeLLM([AssistantMessage(tool_calls=[ToolCall(name="slow_tool", arguments={})]), AssistantMessage(text="finished")])
86 loop = AgentLoop(index, PromptBuilder(f"{tmp}/prompts", index, registry), registry, ToolExecutor(registry, index), llm)
87 agen = loop.stream_events("start")
88 seen = []
89 while True:
90 event = await asyncio.wait_for(anext(agen), timeout=1)
91 seen.append(event["type"])
92 if event["type"] == "tool_progress":
93 self.assertFalse(release.is_set())
94 break
95 release.set()
96 async for event in agen:
97 seen.append(event["type"])
98 self.assertLess(seen.index("tool_progress"), seen.index("tool_result"))
99
100
101 async def test_preflight_compact_summarizes_old_history(self):
102 with tempfile.TemporaryDirectory() as tmp:
103 index = SessionIndex(tmp)
104 registry = ToolRegistry([])
105 compactor = ContextCompactor(None, token_threshold=200, buffer_tokens=0, preserve_last_n=2, summary_max_chars=2000)
106 loop = AgentLoop(index, PromptBuilder(f"{tmp}/prompts", index, registry), registry, ToolExecutor(registry, index), FakeLLM([AssistantMessage(text="after compact")]), compactor)
107 loop.history = [
108 {"role": "user", "content": "old request " + "x" * 1200},
109 {"role": "assistant", "content": "old answer " + "y" * 1200},
110 {"role": "user", "content": "recent request"},
111 {"role": "assistant", "content": "recent answer"},
112 ]
113 events = [event async for event in loop.stream_events("continue")]
114 self.assertIn("compact", [event.get("phase") for event in events if event["type"] == "status"])
115 session = index.active()
116 self.assertIn("Reference Context Only", session["compacted_summary"])
117 self.assertGreaterEqual(session["compacted_turns"], 1)
118 self.assertTrue(session["compaction_snapshots"])
119 self.assertEqual(loop.history[0]["role"], "system")
120 self.assertIn("after compact", loop.history[-1]["content"])
121 self.assertNotIn("old request", index.memory_text())
122
123
124 async def test_llm_sampling_error_yields_error_without_crashing_loop(self):
125 with tempfile.TemporaryDirectory() as tmp:
126 index = SessionIndex(tmp)
127 registry = ToolRegistry([])
128 loop = AgentLoop(index, PromptBuilder(f"{tmp}/prompts", index, registry), registry, ToolExecutor(registry, index), FailingLLM())
129 events = [event async for event in loop.stream_events("start")]
130 self.assertTrue(any(event["type"] == "error" and event.get("metadata", {}).get("error_type") == "llm_sampling_failed" for event in events))
131 self.assertEqual(events[-2]["type"], "done")
132 self.assertEqual(events[-1]["type"], "session")
133 self.assertEqual(index.active()["recent_turn_records"][-1]["status"], "failed")
134
135 async def test_tool_call_continues_then_finishes(self):
136 with tempfile.TemporaryDirectory() as tmp:
137 index = SessionIndex(tmp)
138
139 def hello(args):
140 return ToolResult("hello", True, "hello result")
141
142 registry = ToolRegistry([ToolSpec("hello", "Say hello", hello, schema={"name": ToolArgumentSchema(str, False, "x")})])
143 llm = FakeLLM([AssistantMessage(tool_calls=[ToolCall(name="hello", arguments={})]), AssistantMessage(text="finished")])
144 loop = AgentLoop(index, PromptBuilder(f"{tmp}/prompts", index, registry), registry, ToolExecutor(registry, index), llm)
145 events = [event async for event in loop.stream_events("start")]
146 self.assertTrue(any(event["type"] == "tool_result" for event in events))
147 self.assertEqual(events[-2]["assistant"], "finished")
148
149 async def test_transient_tool_images_reach_next_llm_turn_but_not_events_or_history(self):
150 with tempfile.TemporaryDirectory() as tmp:
151 index = SessionIndex(tmp)
152 data_url = "data:image/jpeg;base64,ZmFrZS1pbWFnZQ=="
153
154 def view(args):
155 return ToolResult(
156 "view_image",
157 True,
158 "image loaded",
159 {"path": "idea2video/frame.png"},
160 model_content=[{"type": "image_url", "image_url": {"url": data_url, "detail": "high"}}],
161 )
162
163 registry = ToolRegistry([ToolSpec("view_image", "View image", view, schema={"path": ToolArgumentSchema(str, True)})])
164 llm = CapturingLLM(
165 [
166 AssistantMessage(tool_calls=[ToolCall(name="view_image", arguments={"path": "idea2video/frame.png"})]),
167 AssistantMessage(text="The frame is visible."),
168 ]
169 )
170 loop = AgentLoop(index, PromptBuilder(f"{tmp}/prompts", index, registry), registry, ToolExecutor(registry, index), llm)
171 events = [event async for event in loop.stream_events("inspect the frame")]
172
173 image_messages = [message for message in llm.calls[1] if message.get("role") == "user" and isinstance(message.get("content"), list)]
174 self.assertEqual(len(image_messages), 1)
175 self.assertEqual(image_messages[0]["content"][1]["image_url"]["url"], data_url)
176 self.assertNotIn(data_url, json.dumps(events))
177 self.assertNotIn(data_url, json.dumps(loop.history))
178 self.assertNotIn(data_url, (index.logs_dir / "tool_calls.jsonl").read_text(encoding="utf-8"))
179
179 lines PYTHON