diff --git a/libs/langgraph/langgraph/stream/transformers.py b/libs/langgraph/langgraph/stream/transformers.py index 0ce420af7..6ca225159 100644 --- a/libs/langgraph/langgraph/stream/transformers.py +++ b/libs/langgraph/langgraph/stream/transformers.py @@ -353,6 +353,23 @@ class LifecyclePayload(TypedDict, total=False): namespace: list[str] graph_name: NotRequired[str] trigger_call_id: NotRequired[str] + cause: NotRequired[dict[str, Any]] + """Optional generic descriptor of what spawned this subgraph. Forwarded + by protocol layers as the wire `lifecycle.started.cause` field. + + Shape: + + - `{"type": "tool_call", "subagent_type": "", "description": ""}` + — set when the subgraph was spawned by a tool invocation routed + through `langgraph.prebuilt.ToolNode`. Mined from the per-call + dispatched task's `input.tool_call.args`. Lets consumers attribute + `lifecycle.started` to the spawning intent without needing the + model's `tool_call_id` (consumers join on the existing + `trigger_call_id` field, which is the unique pregel task id of the + spawn). + + Absent for structurally-spawned subgraphs (parallel branches via + `Send` without ToolNode, nested `graph.invoke()`, etc.).""" error: NotRequired[str] @@ -371,9 +388,9 @@ class _TasksLifecycleBase(StreamTransformer): - `_should_track(ns)` — scope filter (e.g. multi-depth vs direct-children-only). - - `_on_started(ns, graph_name, trigger_call_id)` — first sighting - action (push payload / build handle / etc.). Called once per - discovered namespace. + - `_on_started(ns, graph_name, trigger_call_id, tool_call_id)` — + first sighting action (push payload / build handle / etc.). + Called once per discovered namespace. - `_on_terminal(ns, status, error)` — terminal action (push terminal payload / mark handle status). Called once per tracked namespace at result time, or via `finalize` / `fail` @@ -393,6 +410,13 @@ class _TasksLifecycleBase(StreamTransformer): # Maps tracked namespace -> task_id of the parent task whose # `TaskResultPayload` will close it. self._open: dict[tuple[str, ...], str] = {} + # Maps task_id -> spawn metadata for tasks whose `input` looked + # like a `{"tool_call": {...}, ...}` envelope (the shape + # `langgraph.prebuilt.ToolNode` Send-fans out via + # `ToolCallWithContext`). The lifecycle hook joins on this + # when a child subgraph fires its first task event so it can + # attribute the spawn to the model tool call's args. + self._spawn_metadata: dict[str, dict[str, str]] = {} # --- Template-method hooks (subclass overrides) --- @@ -405,8 +429,17 @@ class _TasksLifecycleBase(StreamTransformer): ns: tuple[str, ...], graph_name: str | None, trigger_call_id: str | None, + spawn_metadata: dict[str, str] | None = None, ) -> None: - """Fired once per discovered namespace (first observed task event).""" + """Fired once per discovered namespace (first observed task event). + + `spawn_metadata` carries spawn-intent fields mined from the + per-call dispatched task's `input.tool_call.args` envelope: + `{"subagent_type": str, "description": str}` (either or both + may be absent). `None` for structurally-spawned subgraphs. + Consumers join on `trigger_call_id` (the pregel task id) for + identity; this dict is purely descriptive metadata. + """ raise NotImplementedError def _on_terminal( @@ -430,21 +463,70 @@ class _TasksLifecycleBase(StreamTransformer): if "result" in data: self._handle_task_result(ns, data) else: - self._handle_task_start(ns) + self._handle_task_start(ns, data) # Tasks events are folded into the synthesized projections; # suppress from the main event log so iterators don't double-see # the same information in two shapes. return False - def _handle_task_start(self, ns: tuple[str, ...]) -> None: + def _handle_task_start(self, ns: tuple[str, ...], data: dict[str, Any]) -> None: + # Mine input shape on every tasks event (not just tracked ones) + # so we capture parent tasks that themselves live outside the + # tracked region but whose `id` will appear as `trigger_call_id` + # for a child subgraph. + self._record_spawn_metadata(data) if not self._should_track(ns) or ns in self._seen: return self._seen.add(ns) graph_name, trigger_call_id = _parse_ns_segment(ns[-1]) - self._on_started(ns, graph_name or None, trigger_call_id) + spawn_metadata = ( + self._spawn_metadata.get(trigger_call_id) + if trigger_call_id is not None + else None + ) + self._on_started( + ns, + graph_name or None, + trigger_call_id, + spawn_metadata, + ) if trigger_call_id is not None: self._open[ns] = trigger_call_id + def _record_spawn_metadata(self, data: dict[str, Any]) -> None: + """Remember `task_id -> spawn metadata` if the task input matches + the `ToolCallWithContext`-style envelope used by Send-fan-out + tool runners (`{"tool_call": {"id": ..., "args": {...}, ...}, ...}`). + + Captures `subagent_type` and `description` from the tool_call + args when present. Duck-typed on shape so 3rd-party tool runners + that mimic the layout participate without importing prebuilt + types. Tool_call_id is intentionally not extracted — consumers + join on `trigger_call_id` (the pregel task id), which is the + same `task.id` we cache by here. + """ + task_id = data.get("id") + if not isinstance(task_id, str): + return + payload = data.get("input") + if not isinstance(payload, dict): + return + tool_call = payload.get("tool_call") + if not isinstance(tool_call, dict): + return + args = tool_call.get("args") + if not isinstance(args, dict): + return + metadata: dict[str, str] = {} + subagent_type = args.get("subagent_type") + if isinstance(subagent_type, str): + metadata["subagent_type"] = subagent_type + description = args.get("description") + if isinstance(description, str): + metadata["description"] = description + if metadata: + self._spawn_metadata[task_id] = metadata + def _pop_terminal_transitions( self, ns: tuple[str, ...], data: dict[str, Any] ) -> list[tuple[tuple[str, ...], SubgraphStatus, str | None]]: @@ -540,6 +622,7 @@ class LifecycleTransformer(_TasksLifecycleBase): ns: tuple[str, ...], graph_name: str | None, trigger_call_id: str | None, + spawn_metadata: dict[str, str] | None = None, ) -> None: if trigger_call_id is None: # Without a task id we can't correlate a parent-result @@ -550,6 +633,10 @@ class LifecycleTransformer(_TasksLifecycleBase): if graph_name: payload["graph_name"] = graph_name payload["trigger_call_id"] = trigger_call_id + if spawn_metadata: + cause: dict[str, Any] = {"type": "tool_call"} + cause.update(spawn_metadata) + payload["cause"] = cause self._channel.push(payload) def _on_terminal( @@ -612,6 +699,7 @@ class SubgraphTransformer(_TasksLifecycleBase): ns: tuple[str, ...], graph_name: str | None, trigger_call_id: str | None, + spawn_metadata: dict[str, str] | None = None, # noqa: ARG002 ) -> None: if self._mux is None: return @@ -724,7 +812,7 @@ class SubgraphTransformer(_TasksLifecycleBase): for child_ns, status, error in self._pop_terminal_transitions(ns, data): await self._aon_terminal(child_ns, status, error) else: - self._handle_task_start(ns) + self._handle_task_start(ns, data) keep = False else: keep = True diff --git a/libs/langgraph/tests/test_stream_lifecycle_transformer.py b/libs/langgraph/tests/test_stream_lifecycle_transformer.py index 1972a245d..0dbaa4e6f 100644 --- a/libs/langgraph/tests/test_stream_lifecycle_transformer.py +++ b/libs/langgraph/tests/test_stream_lifecycle_transformer.py @@ -35,8 +35,16 @@ def _tasks_start( *, task_id: str, name: str, + input: Any = None, ) -> dict[str, Any]: - """Build a `tasks` ProtocolEvent carrying a TaskPayload (start).""" + """Build a `tasks` ProtocolEvent carrying a TaskPayload (start). + + Pass `input={"tool_call": {"args": {...}}}` (or any envelope with + that shape) to exercise the lifecycle transformer's input mining of + spawn-intent metadata (`subagent_type`, `description`) — this + mirrors the `ToolCallWithContext` payload `langgraph.prebuilt.ToolNode` + Send-fans out per tool call. + """ return { "type": "event", "method": "tasks", @@ -46,7 +54,7 @@ def _tasks_start( "data": { "id": task_id, "name": name, - "input": None, + "input": input, "triggers": [], }, }, @@ -126,6 +134,91 @@ def test_started_emitted_on_first_direct_child_task() -> None: assert payload["trigger_call_id"] == "abc123" +def test_started_carries_cause_when_parent_input_has_spawn_metadata() -> None: + """When a parent task's `input` is a `ToolCallWithContext`-shaped + envelope (`{"tool_call": {"args": {...}}, ...}`, the layout + `langgraph.prebuilt.ToolNode` Send-fans out per call), the + transformer mines `subagent_type` and `description` from + `tool_call.args` and remembers them keyed by `parent_task_id`. + When that parent task spawns a subgraph (the child's namespace + ends in `name:`), the `lifecycle.started` payload + carries `cause = {"type": "tool_call", "subagent_type": ..., "description": ...}`. + Consumers join on `trigger_call_id` (the pregel task id) for + identity; this dict is purely descriptive.""" + mux = _build_lifecycle_mux() + # Parent task at root ns whose input matches the Send envelope. + mux.push( + _tasks_start( + [], + task_id="abc123", + name="tools", + input={ + "tool_call": { + "id": "call_xyz", + "name": "task", + "args": { + "subagent_type": "researcher", + "description": "look up weather", + }, + } + }, + ) + ) + # Child subgraph's first task — trigger_call_id parsed from segment. + mux.push(_tasks_start(["agent:abc123"], task_id="t1", name="model")) + + [payload] = _drain_lifecycle(mux) + assert payload["event"] == "started" + assert payload["trigger_call_id"] == "abc123" + assert payload["cause"] == { + "type": "tool_call", + "subagent_type": "researcher", + "description": "look up weather", + } + # tool_call_id is intentionally not in cause — consumers join on + # trigger_call_id (the pregel task id) instead. + assert "tool_call_id" not in payload["cause"] + + +def test_started_cause_with_description_but_no_subagent_type() -> None: + """Partial spawn metadata (only `description`, or only `subagent_type`) + still produces a cause — both fields are optional within the dict.""" + mux = _build_lifecycle_mux() + mux.push( + _tasks_start( + [], + task_id="abc123", + name="tools", + input={ + "tool_call": { + "id": "call_xyz", + "name": "task", + "args": {"description": "do a thing"}, + } + }, + ) + ) + mux.push(_tasks_start(["agent:abc123"], task_id="t1", name="model")) + + [payload] = _drain_lifecycle(mux) + assert payload["cause"] == { + "type": "tool_call", + "description": "do a thing", + } + + +def test_started_omits_cause_for_structurally_spawned_subgraph() -> None: + """Subgraphs spawned without a recognizable tool-call envelope on + the parent's input (Send with custom payloads, plain nested + `graph.invoke`, etc.) don't get a `cause` field on + `lifecycle.started`.""" + mux = _build_lifecycle_mux() + mux.push(_tasks_start(["agent:abc123"], task_id="t1", name="tool")) + + [payload] = _drain_lifecycle(mux) + assert "cause" not in payload + + def test_started_dedup_on_repeat_namespace() -> None: mux = _build_lifecycle_mux() mux.push(_tasks_start(["agent:abc"], task_id="t1", name="a"))