main
py 109 lines 3.69 KB
Raw
1 from types import SimpleNamespace
2
3 import pytest
4
5 import agent as agent_module
6 from agent import Agent
7 from helpers import extension
8 from helpers.llm_result import LLMResult
9
10
11 async def _run_monologue(monkeypatch, chunks, clock, *, advance=0.0):
12 agent = object.__new__(Agent)
13 agent.context = SimpleNamespace(streaming_agent=None, task=None)
14 agent.last_user_message = None
15 events = []
16 provider_callbacks = 0
17
18 async def call_extensions(name, _agent=None, **kwargs):
19 events.append(("extension", name))
20 if name == "response_stream_chunk":
21 stream_data = kwargs["stream_data"]
22 stream_data["full"] = "masked:" + stream_data["full"]
23
24 async def call_model(*, response_callback, **_kwargs):
25 nonlocal provider_callbacks
26 full = ""
27 for chunk in chunks:
28 provider_callbacks += 1
29 clock.now += advance
30 full += chunk
31 stopped = await response_callback(chunk, full)
32 if stopped:
33 return LLMResult(response=stopped)
34 return LLMResult(response=full)
35
36 async def handle_response_stream(full):
37 events.append(("handle", full))
38
39 async def done(_result):
40 return "done"
41
42 async def no_intervention(*_args):
43 return None
44
45 async def no_prompt(**_kwargs):
46 return []
47
48 monkeypatch.setattr(extension, "call_extensions_async", call_extensions)
49 monkeypatch.setattr(agent_module.time, "monotonic", lambda: clock.now)
50 agent.prepare_prompt = no_prompt
51 agent.handle_intervention = no_intervention
52 agent.call_chat_model_turn = call_model
53 agent.handle_response_stream = handle_response_stream
54 agent.hist_add_ai_response = lambda *_args, **_kwargs: SimpleNamespace(id="")
55 agent._remember_llm_result_state = lambda *_args: None
56 agent.process_llm_result_tools = done
57
58 result = await Agent.monologue.__wrapped__(agent)
59 return result, events, provider_callbacks
60
61
62 @pytest.mark.asyncio
63 async def test_response_stream_coalesces_fast_fragments_and_flushes_final(monkeypatch):
64 clock = SimpleNamespace(now=0.0)
65 result, events, provider_callbacks = await _run_monologue(
66 monkeypatch, list("x" * 130), clock
67 )
68
69 handled = [value for kind, value in events if kind == "handle"]
70 chunk_hooks = [
71 event
72 for event in events
73 if event == ("extension", "response_stream_chunk")
74 ]
75
76 assert result == "done"
77 assert provider_callbacks == len(chunk_hooks) == 130
78 assert handled == ["masked:" + "x" * size for size in (128, 130)]
79 assert events.index(("handle", handled[-1])) < events.index(
80 ("extension", "response_stream_end")
81 )
82 assert events.index(("extension", "response_stream_end")) < events.index(
83 ("extension", "message_loop_result")
84 )
85
86
87 @pytest.mark.asyncio
88 async def test_response_stream_time_bound_keeps_slow_fragments_live(monkeypatch):
89 clock = SimpleNamespace(now=0.0)
90 _, events, _ = await _run_monologue(
91 monkeypatch, ["a" * 10] * 3, clock, advance=0.03
92 )
93
94 handled = [value for kind, value in events if kind == "handle"]
95 assert handled == ["masked:" + "a" * size for size in (20, 30)]
96
97
98 @pytest.mark.asyncio
99 async def test_response_stream_still_stops_on_exact_canonical_root(monkeypatch):
100 clock = SimpleNamespace(now=0.0)
101 message = '{"tool_name":"response","tool_args":{"text":"ok"}}'
102 result, events, provider_callbacks = await _run_monologue(
103 monkeypatch, [message[:-1], message[-1], " unreachable"], clock
104 )
105
106 assert result == "done"
107 assert provider_callbacks == 2
108 assert [value for kind, value in events if kind == "handle"] == [message]
109 assert events.count(("extension", "response_stream_chunk")) == 1