|
| 1 | +"""Tests for the sync LangGraph -> Agentex stream event converter. |
| 2 | +
|
| 3 | +Covers: |
| 4 | +- Basic text, tool call, and tool response emission |
| 5 | +- on_final_ai_message callback for usage capture |
| 6 | +- Deprecation warning emitted by create_langgraph_tracing_handler |
| 7 | +
|
| 8 | +NOTE: langchain_core imports must be deferred to test-function scope because |
| 9 | +conftest.py stubs out ``langchain_core.messages`` with MagicMock for ADK |
| 10 | +package-level tests. The real classes are imported lazily inside each test. |
| 11 | +""" |
| 12 | + |
| 13 | +from __future__ import annotations |
| 14 | + |
| 15 | +import sys |
| 16 | +import warnings |
| 17 | +from typing import Any, AsyncIterator |
| 18 | + |
| 19 | +import pytest |
| 20 | + |
| 21 | +from agentex.types.task_message_update import ( |
| 22 | + StreamTaskMessageFull, |
| 23 | +) |
| 24 | +from agentex.types.tool_request_content import ToolRequestContent |
| 25 | +from agentex.types.tool_response_content import ToolResponseContent |
| 26 | +from agentex.lib.adk._modules._langgraph_sync import convert_langgraph_to_agentex_events |
| 27 | + |
| 28 | +# --------------------------------------------------------------------------- |
| 29 | +# Helpers |
| 30 | +# --------------------------------------------------------------------------- |
| 31 | + |
| 32 | + |
| 33 | +async def _collect(stream: AsyncIterator[Any]) -> list[Any]: |
| 34 | + return [e async for e in stream] |
| 35 | + |
| 36 | + |
| 37 | +def _make_stream(events: list[tuple[str, Any]]) -> AsyncIterator[tuple[str, Any]]: |
| 38 | + async def _gen(): |
| 39 | + for e in events: |
| 40 | + yield e |
| 41 | + |
| 42 | + return _gen() |
| 43 | + |
| 44 | + |
| 45 | +# --------------------------------------------------------------------------- |
| 46 | +# Remove the conftest stubs for langchain_core so real classes are used |
| 47 | +# --------------------------------------------------------------------------- |
| 48 | + |
| 49 | + |
| 50 | +@pytest.fixture(autouse=True) |
| 51 | +def _real_langchain_core(): |
| 52 | + """Remove conftest MagicMock stubs so real langchain_core types are used.""" |
| 53 | + stub_keys = [k for k in sys.modules if k.startswith("langchain_core") or k.startswith("langgraph")] |
| 54 | + saved = {k: sys.modules.pop(k) for k in stub_keys} |
| 55 | + # Re-import the real modules |
| 56 | + import importlib |
| 57 | + |
| 58 | + importlib.import_module("langchain_core.messages") |
| 59 | + yield |
| 60 | + # Restore stubs after the test |
| 61 | + sys.modules.update(saved) |
| 62 | + |
| 63 | + |
| 64 | +class TestTextStreaming: |
| 65 | + async def test_plain_text_emits_start_delta_done(self): |
| 66 | + from langchain_core.messages import AIMessage, AIMessageChunk |
| 67 | + |
| 68 | + chunk = AIMessageChunk(content="Hello, world!") |
| 69 | + events = [ |
| 70 | + ("messages", (chunk, {})), |
| 71 | + ("updates", {"agent": {"messages": [AIMessage(content="Hello, world!")]}}), |
| 72 | + ] |
| 73 | + out = await _collect(convert_langgraph_to_agentex_events(_make_stream(events))) |
| 74 | + types = [type(e).__name__ for e in out] |
| 75 | + assert "StreamTaskMessageStart" in types |
| 76 | + assert "StreamTaskMessageDelta" in types |
| 77 | + assert "StreamTaskMessageDone" in types |
| 78 | + |
| 79 | + async def test_empty_chunk_content_is_skipped(self): |
| 80 | + from langchain_core.messages import AIMessageChunk |
| 81 | + |
| 82 | + chunk = AIMessageChunk(content="") |
| 83 | + events = [("messages", (chunk, {}))] |
| 84 | + out = await _collect(convert_langgraph_to_agentex_events(_make_stream(events))) |
| 85 | + assert out == [] |
| 86 | + |
| 87 | + |
| 88 | +class TestToolCallEmission: |
| 89 | + async def test_tool_call_emits_full_message(self): |
| 90 | + from langchain_core.messages import AIMessage |
| 91 | + |
| 92 | + tc = {"id": "call_1", "name": "get_weather", "args": {"city": "Paris"}} |
| 93 | + ai_msg = AIMessage(content="", tool_calls=[tc]) |
| 94 | + events = [("updates", {"agent": {"messages": [ai_msg]}})] |
| 95 | + out = await _collect(convert_langgraph_to_agentex_events(_make_stream(events))) |
| 96 | + assert len(out) == 1 |
| 97 | + assert isinstance(out[0], StreamTaskMessageFull) |
| 98 | + content = out[0].content |
| 99 | + assert isinstance(content, ToolRequestContent) |
| 100 | + assert content.tool_call_id == "call_1" |
| 101 | + assert content.name == "get_weather" |
| 102 | + assert content.arguments == {"city": "Paris"} |
| 103 | + assert content.author == "agent" |
| 104 | + |
| 105 | + async def test_tool_response_emits_full_message(self): |
| 106 | + from langchain_core.messages import ToolMessage |
| 107 | + |
| 108 | + tool_msg = ToolMessage(content="Sunny, 72F", tool_call_id="call_1", name="get_weather") |
| 109 | + events = [("updates", {"tools": {"messages": [tool_msg]}})] |
| 110 | + out = await _collect(convert_langgraph_to_agentex_events(_make_stream(events))) |
| 111 | + assert len(out) == 1 |
| 112 | + assert isinstance(out[0], StreamTaskMessageFull) |
| 113 | + content = out[0].content |
| 114 | + assert isinstance(content, ToolResponseContent) |
| 115 | + assert content.tool_call_id == "call_1" |
| 116 | + assert content.name == "get_weather" |
| 117 | + assert content.content == "Sunny, 72F" |
| 118 | + assert content.author == "agent" |
| 119 | + |
| 120 | + |
| 121 | +class TestOnFinalAiMessageCallback: |
| 122 | + async def test_callback_called_for_ai_message_in_agent_node(self): |
| 123 | + from langchain_core.messages import AIMessage |
| 124 | + |
| 125 | + captured: list[Any] = [] |
| 126 | + ai_msg = AIMessage(content="Hello!") |
| 127 | + |
| 128 | + events = [("updates", {"agent": {"messages": [ai_msg]}})] |
| 129 | + await _collect(convert_langgraph_to_agentex_events(_make_stream(events), on_final_ai_message=captured.append)) |
| 130 | + assert len(captured) == 1 |
| 131 | + assert captured[0] is ai_msg |
| 132 | + |
| 133 | + async def test_callback_not_called_for_tool_messages(self): |
| 134 | + from langchain_core.messages import ToolMessage |
| 135 | + |
| 136 | + captured: list[Any] = [] |
| 137 | + tool_msg = ToolMessage(content="result", tool_call_id="c1", name="t") |
| 138 | + |
| 139 | + events = [("updates", {"tools": {"messages": [tool_msg]}})] |
| 140 | + await _collect(convert_langgraph_to_agentex_events(_make_stream(events), on_final_ai_message=captured.append)) |
| 141 | + assert captured == [] |
| 142 | + |
| 143 | + async def test_callback_receives_usage_metadata(self): |
| 144 | + from langchain_core.messages import AIMessage |
| 145 | + |
| 146 | + captured: list[Any] = [] |
| 147 | + usage = {"input_tokens": 10, "output_tokens": 5, "total_tokens": 15} |
| 148 | + ai_msg = AIMessage(content="Answer.", usage_metadata=usage) |
| 149 | + |
| 150 | + events = [("updates", {"agent": {"messages": [ai_msg]}})] |
| 151 | + await _collect(convert_langgraph_to_agentex_events(_make_stream(events), on_final_ai_message=captured.append)) |
| 152 | + assert len(captured) == 1 |
| 153 | + assert captured[0].usage_metadata == usage |
| 154 | + |
| 155 | + async def test_no_callback_is_noop(self): |
| 156 | + from langchain_core.messages import AIMessage |
| 157 | + |
| 158 | + ai_msg = AIMessage(content="Hello!") |
| 159 | + events = [("updates", {"agent": {"messages": [ai_msg]}})] |
| 160 | + out = await _collect(convert_langgraph_to_agentex_events(_make_stream(events))) |
| 161 | + assert isinstance(out, list) |
| 162 | + |
| 163 | + async def test_callback_called_multiple_times_for_multi_step(self): |
| 164 | + from langchain_core.messages import AIMessage |
| 165 | + |
| 166 | + captured: list[Any] = [] |
| 167 | + ai_msg_1 = AIMessage(content="Step 1") |
| 168 | + ai_msg_2 = AIMessage(content="Step 2") |
| 169 | + |
| 170 | + events = [ |
| 171 | + ("updates", {"agent": {"messages": [ai_msg_1]}}), |
| 172 | + ("updates", {"agent": {"messages": [ai_msg_2]}}), |
| 173 | + ] |
| 174 | + await _collect(convert_langgraph_to_agentex_events(_make_stream(events), on_final_ai_message=captured.append)) |
| 175 | + assert len(captured) == 2 |
| 176 | + assert captured[0] is ai_msg_1 |
| 177 | + assert captured[1] is ai_msg_2 |
| 178 | + |
| 179 | + async def test_callback_called_after_tool_call_events_yielded(self): |
| 180 | + """The callback fires after all events for that AIMessage are yielded.""" |
| 181 | + from langchain_core.messages import AIMessage |
| 182 | + |
| 183 | + yield_order: list[str] = [] |
| 184 | + |
| 185 | + async def _gen(): |
| 186 | + tc = {"id": "c1", "name": "t", "args": {}} |
| 187 | + ai_msg = AIMessage(content="", tool_calls=[tc]) |
| 188 | + yield ("updates", {"agent": {"messages": [ai_msg]}}) |
| 189 | + |
| 190 | + def _cb(msg): |
| 191 | + yield_order.append("callback") |
| 192 | + |
| 193 | + async for _ in convert_langgraph_to_agentex_events(_gen(), on_final_ai_message=_cb): |
| 194 | + yield_order.append("event") |
| 195 | + |
| 196 | + # The tool call Full event is emitted before the callback fires |
| 197 | + assert yield_order.index("event") < yield_order.index("callback") |
| 198 | + |
| 199 | + |
| 200 | +class TestDeprecationWarning: |
| 201 | + def test_create_langgraph_tracing_handler_emits_deprecation_warning(self): |
| 202 | + from agentex.lib.adk._modules._langgraph_tracing import create_langgraph_tracing_handler |
| 203 | + |
| 204 | + with warnings.catch_warnings(record=True) as w: |
| 205 | + warnings.simplefilter("always") |
| 206 | + create_langgraph_tracing_handler(trace_id="t1") |
| 207 | + assert any(issubclass(warning.category, DeprecationWarning) for warning in w), ( |
| 208 | + "create_langgraph_tracing_handler must emit a DeprecationWarning" |
| 209 | + ) |
0 commit comments