|
| 1 | +"""Tests for LangGraph node-to-node Messages edge emission.""" |
| 2 | + |
| 3 | +from __future__ import annotations |
| 4 | + |
| 5 | +from typing import Any |
| 6 | + |
| 7 | +from agent_assembly.adapters.langgraph import patch as lg_patch |
| 8 | + |
| 9 | + |
| 10 | +class RecordingEdgeEmitter: |
| 11 | + """Synchronous test double that records emitted edges.""" |
| 12 | + |
| 13 | + def __init__(self) -> None: |
| 14 | + self.edges: list[tuple[str, str, str, dict | None]] = [] |
| 15 | + |
| 16 | + def emit(self, source: str, target: str, edge_type: str, metadata: dict | None = None) -> None: |
| 17 | + self.edges.append((source, target, edge_type, metadata)) |
| 18 | + |
| 19 | + |
| 20 | +class _NullCallbackHandler: |
| 21 | + """Minimal callback handler that satisfies the patch interface.""" |
| 22 | + |
| 23 | + def on_graph_node_start(self, **_kwargs: Any) -> None: |
| 24 | + pass |
| 25 | + |
| 26 | + def on_graph_node_end(self, **_kwargs: Any) -> None: |
| 27 | + pass |
| 28 | + |
| 29 | + |
| 30 | +def _make_node_map(*node_names: str, handler: Any) -> dict[str, Any]: |
| 31 | + """Build a fake compiled-graph node map with simple callables.""" |
| 32 | + node_map: dict[str, Any] = {} |
| 33 | + for name in node_names: |
| 34 | + def make_func(n: str) -> Any: |
| 35 | + def node_fn(state: Any) -> dict[str, Any]: |
| 36 | + return {"node": n, **state} |
| 37 | + |
| 38 | + return node_fn |
| 39 | + |
| 40 | + node_map[name] = make_func(name) |
| 41 | + return node_map |
| 42 | + |
| 43 | + |
| 44 | +def _run_two_node_graph( |
| 45 | + node_a: str, |
| 46 | + node_b: str, |
| 47 | + emitter: RecordingEdgeEmitter, |
| 48 | + initial_state: dict[str, Any] | None = None, |
| 49 | +) -> None: |
| 50 | + """Simulate a 2-node sequential graph execution with patched wrappers.""" |
| 51 | + handler = _NullCallbackHandler() |
| 52 | + node_map = _make_node_map(node_a, node_b, handler=handler) |
| 53 | + |
| 54 | + lg_patch.set_edge_emitter(emitter) |
| 55 | + try: |
| 56 | + lg_patch._wrap_node_map(node_map, handler) |
| 57 | + |
| 58 | + state: dict[str, Any] = initial_state or {} |
| 59 | + state = node_map[node_a](state) |
| 60 | + state = node_map[node_b](state) |
| 61 | + finally: |
| 62 | + lg_patch.set_edge_emitter(None) |
| 63 | + lg_patch._NODE_TRANSITION.name = None |
| 64 | + |
| 65 | + |
| 66 | +def test_messages_edge_emitted_between_two_sequential_nodes() -> None: |
| 67 | + emitter = RecordingEdgeEmitter() |
| 68 | + _run_two_node_graph("node_a", "node_b", emitter) |
| 69 | + |
| 70 | + assert len(emitter.edges) == 1 |
| 71 | + src, tgt, etype, meta = emitter.edges[0] |
| 72 | + assert src == "node_a" |
| 73 | + assert tgt == "node_b" |
| 74 | + assert etype == "messages" |
| 75 | + assert isinstance(meta, dict) |
| 76 | + assert "transition_input_keys" in meta |
| 77 | + |
| 78 | + |
| 79 | +def test_transition_input_keys_contains_state_keys() -> None: |
| 80 | + emitter = RecordingEdgeEmitter() |
| 81 | + _run_two_node_graph("node_a", "node_b", emitter, initial_state={"msg": "hi", "count": 0}) |
| 82 | + |
| 83 | + assert len(emitter.edges) == 1 |
| 84 | + meta = emitter.edges[0][3] |
| 85 | + assert meta is not None |
| 86 | + # transition_input_keys reflects the state keys seen when node_a completed |
| 87 | + assert set(meta["transition_input_keys"]) >= {"msg", "count"} |
| 88 | + |
| 89 | + |
| 90 | +def test_no_edge_emitted_for_first_node_in_graph() -> None: |
| 91 | + """The very first node has no predecessor so no edge should be emitted.""" |
| 92 | + emitter = RecordingEdgeEmitter() |
| 93 | + lg_patch.set_edge_emitter(emitter) |
| 94 | + try: |
| 95 | + handler = _NullCallbackHandler() |
| 96 | + node_map = _make_node_map("only_node", handler=handler) |
| 97 | + lg_patch._wrap_node_map(node_map, handler) |
| 98 | + node_map["only_node"]({}) |
| 99 | + finally: |
| 100 | + lg_patch.set_edge_emitter(None) |
| 101 | + lg_patch._NODE_TRANSITION.name = None |
| 102 | + |
| 103 | + assert emitter.edges == [] |
| 104 | + |
| 105 | + |
| 106 | +def test_three_node_graph_emits_two_edges() -> None: |
| 107 | + emitter = RecordingEdgeEmitter() |
| 108 | + lg_patch.set_edge_emitter(emitter) |
| 109 | + try: |
| 110 | + handler = _NullCallbackHandler() |
| 111 | + node_map = _make_node_map("a", "b", "c", handler=handler) |
| 112 | + lg_patch._wrap_node_map(node_map, handler) |
| 113 | + state: dict[str, Any] = {} |
| 114 | + state = node_map["a"](state) |
| 115 | + state = node_map["b"](state) |
| 116 | + node_map["c"](state) |
| 117 | + finally: |
| 118 | + lg_patch.set_edge_emitter(None) |
| 119 | + lg_patch._NODE_TRANSITION.name = None |
| 120 | + |
| 121 | + assert len(emitter.edges) == 2 |
| 122 | + assert emitter.edges[0][:3] == ("a", "b", "messages") |
| 123 | + assert emitter.edges[1][:3] == ("b", "c", "messages") |
| 124 | + |
| 125 | + |
| 126 | +def test_no_edge_emitted_when_emitter_is_none() -> None: |
| 127 | + """When no emitter is registered, transitions must not raise.""" |
| 128 | + lg_patch.set_edge_emitter(None) |
| 129 | + try: |
| 130 | + handler = _NullCallbackHandler() |
| 131 | + node_map = _make_node_map("x", "y", handler=handler) |
| 132 | + lg_patch._wrap_node_map(node_map, handler) |
| 133 | + state: dict[str, Any] = {} |
| 134 | + state = node_map["x"](state) |
| 135 | + node_map["y"](state) # Should not raise |
| 136 | + finally: |
| 137 | + lg_patch._NODE_TRANSITION.name = None |
0 commit comments