| 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 |