|
| 1 | +"""Characterization tests for stream_langgraph_events. |
| 2 | +
|
| 3 | +These tests record the current behavior of the bespoke ``stream_langgraph_events`` |
| 4 | +implementation BEFORE the unified-surface refactor (Task 4). They act as a |
| 5 | +contract test: after Task 4 rewrites the internals, these tests must still pass, |
| 6 | +proving behavioral parity. |
| 7 | +
|
| 8 | +NOTE: langchain_core imports are deferred to test scope because conftest.py |
| 9 | +stubs ``langchain_core.messages`` with MagicMock. |
| 10 | +""" |
| 11 | + |
| 12 | +from __future__ import annotations |
| 13 | + |
| 14 | +import sys |
| 15 | +from typing import Any |
| 16 | +from dataclasses import field, dataclass |
| 17 | + |
| 18 | +import pytest |
| 19 | + |
| 20 | +from agentex.types.task_message import TaskMessage |
| 21 | +from agentex.types.text_content import TextContent |
| 22 | +from agentex.types.task_message_delta import TextDelta |
| 23 | +from agentex.types.task_message_update import StreamTaskMessageDelta |
| 24 | +from agentex.lib.adk._modules._langgraph_async import stream_langgraph_events |
| 25 | + |
| 26 | +TASK_ID = "task-test" |
| 27 | + |
| 28 | + |
| 29 | +# --------------------------------------------------------------------------- |
| 30 | +# Remove conftest stubs so real langchain_core types are used |
| 31 | +# --------------------------------------------------------------------------- |
| 32 | + |
| 33 | + |
| 34 | +@pytest.fixture(autouse=True) |
| 35 | +def _real_langchain_core(): |
| 36 | + stub_keys = [k for k in sys.modules if k.startswith("langchain_core") or k.startswith("langgraph")] |
| 37 | + saved = {k: sys.modules.pop(k) for k in stub_keys} |
| 38 | + import importlib |
| 39 | + |
| 40 | + importlib.import_module("langchain_core.messages") |
| 41 | + yield |
| 42 | + sys.modules.update(saved) |
| 43 | + |
| 44 | + |
| 45 | +# --------------------------------------------------------------------------- |
| 46 | +# Fake streaming infrastructure (mirrors test_pydantic_ai_async.py pattern) |
| 47 | +# --------------------------------------------------------------------------- |
| 48 | + |
| 49 | + |
| 50 | +@dataclass |
| 51 | +class FakeContext: |
| 52 | + initial_content: Any |
| 53 | + task_message: TaskMessage |
| 54 | + closed: bool = False |
| 55 | + updates: list[StreamTaskMessageDelta] = field(default_factory=list) |
| 56 | + |
| 57 | + async def __aenter__(self) -> "FakeContext": |
| 58 | + return self |
| 59 | + |
| 60 | + async def __aexit__(self, exc_type, exc_val, exc_tb) -> bool: |
| 61 | + await self.close() |
| 62 | + return False |
| 63 | + |
| 64 | + async def stream_update(self, update: StreamTaskMessageDelta) -> None: |
| 65 | + if self.closed: |
| 66 | + raise AssertionError("stream_update called after close") |
| 67 | + self.updates.append(update) |
| 68 | + |
| 69 | + async def close(self) -> None: |
| 70 | + self.closed = True |
| 71 | + |
| 72 | + |
| 73 | +class FakeStreamingModule: |
| 74 | + def __init__(self) -> None: |
| 75 | + self.contexts: list[FakeContext] = [] |
| 76 | + |
| 77 | + def streaming_task_message_context(self, *, task_id: str, initial_content: Any) -> FakeContext: |
| 78 | + tm = TaskMessage( |
| 79 | + id=f"m{len(self.contexts) + 1}", |
| 80 | + task_id=task_id, |
| 81 | + content=initial_content, |
| 82 | + streaming_status="IN_PROGRESS", |
| 83 | + ) |
| 84 | + ctx = FakeContext(initial_content=initial_content, task_message=tm) |
| 85 | + self.contexts.append(ctx) |
| 86 | + return ctx |
| 87 | + |
| 88 | + |
| 89 | +class FakeMessagesModule: |
| 90 | + def __init__(self) -> None: |
| 91 | + self.created: list[dict[str, Any]] = [] |
| 92 | + |
| 93 | + async def create(self, *, task_id: str, content: Any) -> TaskMessage: |
| 94 | + self.created.append({"task_id": task_id, "content": content}) |
| 95 | + return TaskMessage( |
| 96 | + id=f"created-{len(self.created)}", |
| 97 | + task_id=task_id, |
| 98 | + content=content, |
| 99 | + streaming_status="DONE", |
| 100 | + ) |
| 101 | + |
| 102 | + |
| 103 | +@pytest.fixture |
| 104 | +def fake_adk(monkeypatch): |
| 105 | + from agentex.lib import adk as adk_module |
| 106 | + |
| 107 | + streaming = FakeStreamingModule() |
| 108 | + messages = FakeMessagesModule() |
| 109 | + monkeypatch.setattr(adk_module, "streaming", streaming) |
| 110 | + monkeypatch.setattr(adk_module, "messages", messages) |
| 111 | + return streaming, messages |
| 112 | + |
| 113 | + |
| 114 | +def _make_stream(events: list[tuple[str, Any]]): |
| 115 | + async def _gen(): |
| 116 | + for e in events: |
| 117 | + yield e |
| 118 | + |
| 119 | + return _gen() |
| 120 | + |
| 121 | + |
| 122 | +def _text_deltas(ctx: FakeContext) -> list[str]: |
| 123 | + out: list[str] = [] |
| 124 | + for u in ctx.updates: |
| 125 | + if isinstance(u.delta, TextDelta): |
| 126 | + out.append(u.delta.text_delta or "") |
| 127 | + return out |
| 128 | + |
| 129 | + |
| 130 | +# --------------------------------------------------------------------------- |
| 131 | +# Characterization tests |
| 132 | +# --------------------------------------------------------------------------- |
| 133 | + |
| 134 | + |
| 135 | +class TestCharacterization: |
| 136 | + async def test_plain_text_streams_and_returns_final_text( |
| 137 | + self, fake_adk: tuple[FakeStreamingModule, FakeMessagesModule] |
| 138 | + ) -> None: |
| 139 | + from langchain_core.messages import AIMessage, AIMessageChunk |
| 140 | + |
| 141 | + streaming, messages = fake_adk |
| 142 | + chunk = AIMessageChunk(content="Hello, world!") |
| 143 | + ai_msg = AIMessage(content="Hello, world!") |
| 144 | + stream = _make_stream( |
| 145 | + [ |
| 146 | + ("messages", (chunk, {})), |
| 147 | + ("updates", {"agent": {"messages": [ai_msg]}}), |
| 148 | + ] |
| 149 | + ) |
| 150 | + |
| 151 | + final = await stream_langgraph_events(stream, TASK_ID) |
| 152 | + |
| 153 | + assert final == "Hello, world!" |
| 154 | + assert len(streaming.contexts) == 1 |
| 155 | + ctx = streaming.contexts[0] |
| 156 | + assert isinstance(ctx.initial_content, TextContent) |
| 157 | + assert _text_deltas(ctx) == ["Hello, world!"] |
| 158 | + assert ctx.closed is True |
| 159 | + |
| 160 | + async def test_empty_stream_returns_empty_string( |
| 161 | + self, fake_adk: tuple[FakeStreamingModule, FakeMessagesModule] |
| 162 | + ) -> None: |
| 163 | + streaming, _ = fake_adk |
| 164 | + final = await stream_langgraph_events(_make_stream([]), TASK_ID) |
| 165 | + assert final == "" |
| 166 | + assert streaming.contexts == [] |
| 167 | + |
| 168 | + async def test_tool_call_creates_tool_request_message( |
| 169 | + self, fake_adk: tuple[FakeStreamingModule, FakeMessagesModule] |
| 170 | + ) -> None: |
| 171 | + from langchain_core.messages import AIMessage |
| 172 | + |
| 173 | + _, messages = fake_adk |
| 174 | + tc = {"id": "call_1", "name": "get_weather", "args": {"city": "Paris"}} |
| 175 | + ai_msg = AIMessage(content="", tool_calls=[tc]) |
| 176 | + stream = _make_stream([("updates", {"agent": {"messages": [ai_msg]}})]) |
| 177 | + |
| 178 | + await stream_langgraph_events(stream, TASK_ID) |
| 179 | + |
| 180 | + assert len(messages.created) == 1 |
| 181 | + content = messages.created[0]["content"] |
| 182 | + from agentex.types.tool_request_content import ToolRequestContent |
| 183 | + |
| 184 | + assert isinstance(content, ToolRequestContent) |
| 185 | + assert content.tool_call_id == "call_1" |
| 186 | + assert content.name == "get_weather" |
| 187 | + assert content.arguments == {"city": "Paris"} |
| 188 | + |
| 189 | + async def test_tool_response_creates_tool_response_message( |
| 190 | + self, fake_adk: tuple[FakeStreamingModule, FakeMessagesModule] |
| 191 | + ) -> None: |
| 192 | + from langchain_core.messages import ToolMessage |
| 193 | + |
| 194 | + _, messages = fake_adk |
| 195 | + tool_msg = ToolMessage(content="Sunny, 72F", tool_call_id="call_1", name="get_weather") |
| 196 | + stream = _make_stream([("updates", {"tools": {"messages": [tool_msg]}})]) |
| 197 | + |
| 198 | + await stream_langgraph_events(stream, TASK_ID) |
| 199 | + |
| 200 | + assert len(messages.created) == 1 |
| 201 | + content = messages.created[0]["content"] |
| 202 | + from agentex.types.tool_response_content import ToolResponseContent |
| 203 | + |
| 204 | + assert isinstance(content, ToolResponseContent) |
| 205 | + assert content.tool_call_id == "call_1" |
| 206 | + assert content.name == "get_weather" |
| 207 | + assert content.content == "Sunny, 72F" |
| 208 | + |
| 209 | + async def test_multi_step_text_then_tool_then_text( |
| 210 | + self, fake_adk: tuple[FakeStreamingModule, FakeMessagesModule] |
| 211 | + ) -> None: |
| 212 | + from langchain_core.messages import AIMessage, ToolMessage, AIMessageChunk |
| 213 | + |
| 214 | + streaming, messages = fake_adk |
| 215 | + chunk1 = AIMessageChunk(content="Looking up...") |
| 216 | + ai_msg1 = AIMessage(content="Looking up...", tool_calls=[{"id": "c1", "name": "search", "args": {}}]) |
| 217 | + tool_msg = ToolMessage(content="result", tool_call_id="c1", name="search") |
| 218 | + chunk2 = AIMessageChunk(content="Found it!") |
| 219 | + ai_msg2 = AIMessage(content="Found it!") |
| 220 | + |
| 221 | + stream = _make_stream( |
| 222 | + [ |
| 223 | + ("messages", (chunk1, {})), |
| 224 | + ("updates", {"agent": {"messages": [ai_msg1]}}), |
| 225 | + ("updates", {"tools": {"messages": [tool_msg]}}), |
| 226 | + ("messages", (chunk2, {})), |
| 227 | + ("updates", {"agent": {"messages": [ai_msg2]}}), |
| 228 | + ] |
| 229 | + ) |
| 230 | + |
| 231 | + final = await stream_langgraph_events(stream, TASK_ID) |
| 232 | + |
| 233 | + assert final == "Found it!" |
| 234 | + # Tool request + tool response messages |
| 235 | + assert len(messages.created) == 2 |
| 236 | + # Two text streaming contexts |
| 237 | + assert len(streaming.contexts) == 2 |
| 238 | + assert all(ctx.closed for ctx in streaming.contexts) |
| 239 | + |
| 240 | + async def test_context_closed_on_exception(self, fake_adk: tuple[FakeStreamingModule, FakeMessagesModule]) -> None: |
| 241 | + from langchain_core.messages import AIMessageChunk |
| 242 | + |
| 243 | + streaming, _ = fake_adk |
| 244 | + |
| 245 | + async def _boom(): |
| 246 | + chunk = AIMessageChunk(content="partial") |
| 247 | + yield ("messages", (chunk, {})) |
| 248 | + raise RuntimeError("upstream exploded") |
| 249 | + |
| 250 | + with pytest.raises(RuntimeError, match="upstream exploded"): |
| 251 | + await stream_langgraph_events(_boom(), TASK_ID) |
| 252 | + |
| 253 | + assert streaming.contexts[0].closed is True |
0 commit comments