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