codekingpro/portable-devtools
114k
1from __future__ import annotations2 3from collections.abc import AsyncIterator, Callable, Iterator4from contextvars import ContextVar, Token5from typing import Any, TypeVar, cast6from uuid import UUID7 8from langchain_core.callbacks import BaseCallbackHandler9 10from langgraph._internal._constants import NS_SEP11from langgraph.constants import TAG_NOSTREAM12from langgraph.pregel.protocol import StreamChunk13 14try:15 from langchain_core.tracers._streaming import _StreamingCallbackHandler16except ImportError:17 _StreamingCallbackHandler = object # type: ignore[assignment,misc]18 19 20T = TypeVar("T")21 22ToolCallWriter = Callable[[Any], None]23"""A closure bound to a single tool call that emits `tool-output-delta` events."""24 25_tool_call_writer: ContextVar[ToolCallWriter | None] = ContextVar(26 "langgraph_tool_call_writer", default=None27)28"""ContextVar holding the writer for the currently-executing tool call.29 30Set by `StreamToolCallHandler.on_tool_start` and reset on end/error.31Read by `ToolRuntime.emit_output_delta` (in `langgraph.prebuilt`).32"""33 34 35class StreamToolCallHandler(BaseCallbackHandler, _StreamingCallbackHandler):36 """Callback handler that emits tool-call lifecycle events on the stream.37 38 Fires on LangChain's `on_tool_*` callbacks and pushes to the `tools`39 stream mode. Emits `tool-started` / `tool-output-delta` /40 `tool-finished` / `tool-error` payloads keyed by `tool_call_id`.41 42 While a tool is executing, this handler sets `_tool_call_writer` to a43 closure bound to that call's namespace and `tool_call_id`.44 `ToolRuntime.emit_output_delta` reads that ContextVar so tool bodies45 can stream partial output without threading the writer through their46 own signature.47 48 Attached by `Pregel.stream` / `astream` when `"tools"` is in49 `stream_modes`. `run_inline = True` keeps event ordering50 deterministic.51 """52 53 run_inline = True54 55 def __init__(56 self,57 stream: Callable[[StreamChunk], None],58 subgraphs: bool,59 *,60 parent_ns: tuple[str, ...] | None = None,61 ) -> None:62 """Configure the handler to stream tool-call events.63 64 Args:65 stream: Callable that accepts a `StreamChunk` tuple66 `(namespace, mode, payload)` and enqueues it.67 subgraphs: Whether to emit events from tools called inside68 nested subgraphs. When False, only tools at the69 handler's own scope (`parent_ns`) emit.70 parent_ns: Namespace where the handler was attached.71 Mirrors the `StreamMessagesHandler` escape hatch:72 tools whose containing namespace equals `parent_ns`73 still emit even with `subgraphs=False`, so a node that74 explicitly streams a subgraph with `stream_mode="tools"`75 sees its own tools.76 """77 self.stream = stream78 self.subgraphs = subgraphs79 self.parent_ns = parent_ns80 # run_id → (namespace, tool_call_id, ContextVar token)81 # `on_tool_end` does not receive `tool_call_id` in kwargs, so82 # we correlate by `run_id` which is present on every callback.83 self._run_to_call: dict[84 UUID, tuple[tuple[str, ...], str, Token[ToolCallWriter | None]]85 ] = {}86 87 def _ns_for_emit(88 self,89 metadata: dict[str, Any] | None,90 tags: list[str] | None,91 ) -> tuple[str, ...] | None:92 """Resolve the namespace this tool call should emit at, or `None` to skip.93 94 Mirrors `StreamMessagesHandler.on_chat_model_start`'s namespace95 derivation: parses `langgraph_checkpoint_ns` (which ends with96 the `node_name:task_id` of the calling node), drops that97 trailing segment, and returns the containing subgraph's own98 namespace. Returns `None` when the call should be silently99 suppressed:100 101 - `metadata` is missing — handler is attached to a context102 without Pregel routing info.103 - `TAG_NOSTREAM` is in `tags` — caller explicitly opted out.104 - Tool runs in a subgraph (`len(ns) > 0`) and the handler was105 attached with `subgraphs=False` and a different `parent_ns`106 than the call's containing subgraph.107 """108 if not metadata:109 return None110 if tags and TAG_NOSTREAM in tags:111 return None112 nskey = metadata.get("langgraph_checkpoint_ns")113 if not nskey:114 ns: tuple[str, ...] = ()115 else:116 ns = tuple(cast(str, nskey).split(NS_SEP))[:-1]117 if not self.subgraphs and len(ns) > 0 and ns != self.parent_ns:118 return None119 return ns120 121 def _start(122 self,123 serialized: dict[str, Any] | None,124 input_str: str,125 *,126 run_id: UUID,127 metadata: dict[str, Any] | None,128 tags: list[str] | None,129 inputs: dict[str, Any] | None,130 kwargs: dict[str, Any],131 ) -> None:132 ns = self._ns_for_emit(metadata, tags)133 if ns is None:134 return135 tool_call_id = cast("str | None", kwargs.get("tool_call_id")) or str(run_id)136 tool_name = (137 (serialized or {}).get("name")138 or cast("str | None", kwargs.get("name"))139 or ""140 )141 142 def writer(delta: Any) -> None:143 self.stream(144 (145 ns,146 "tools",147 {148 "event": "tool-output-delta",149 "tool_call_id": tool_call_id,150 "delta": delta,151 },152 )153 )154 155 token = _tool_call_writer.set(writer)156 self._run_to_call[run_id] = (ns, tool_call_id, token)157 158 payload: dict[str, Any] = {159 "event": "tool-started",160 "tool_call_id": tool_call_id,161 "tool_name": tool_name,162 }163 if inputs is not None:164 payload["input"] = inputs165 self.stream((ns, "tools", payload))166 167 def _end(self, output: Any, *, run_id: UUID) -> None:168 info = self._run_to_call.pop(run_id, None)169 if info is None:170 return171 ns, tool_call_id, token = info172 self._reset_writer(token)173 self.stream(174 (175 ns,176 "tools",177 {178 "event": "tool-finished",179 "tool_call_id": tool_call_id,180 "output": output,181 },182 )183 )184 185 def _error(self, error: BaseException, *, run_id: UUID) -> None:186 info = self._run_to_call.pop(run_id, None)187 if info is None:188 return189 ns, tool_call_id, token = info190 self._reset_writer(token)191 self.stream(192 (193 ns,194 "tools",195 {196 "event": "tool-error",197 "tool_call_id": tool_call_id,198 "message": str(error),199 },200 )201 )202 203 def tap_output_aiter(204 self, run_id: UUID, output: AsyncIterator[T]205 ) -> AsyncIterator[T]:206 """Pass-through — required by the `_StreamingCallbackHandler` protocol."""207 return output208 209 def tap_output_iter(self, run_id: UUID, output: Iterator[T]) -> Iterator[T]:210 """Pass-through — sync counterpart to `tap_output_aiter`."""211 return output212 213 @staticmethod214 def _reset_writer(token: Token[ToolCallWriter | None]) -> None:215 # Token is invalid if `on_tool_end` runs in a different context216 # than `on_tool_start` (e.g. langchain may hand off to a thread217 # worker without copying the context). Swallow that case; the218 # ContextVar lifetime is bounded by the enclosing task anyway.219 try:220 _tool_call_writer.reset(token)221 except ValueError:222 pass223 224 # ------------------------------------------------------------------225 # Sync callbacks226 # ------------------------------------------------------------------227 228 def on_tool_start(229 self,230 serialized: dict[str, Any],231 input_str: str,232 *,233 run_id: UUID,234 parent_run_id: UUID | None = None,235 tags: list[str] | None = None,236 metadata: dict[str, Any] | None = None,237 inputs: dict[str, Any] | None = None,238 **kwargs: Any,239 ) -> Any:240 self._start(241 serialized,242 input_str,243 run_id=run_id,244 metadata=metadata,245 tags=tags,246 inputs=inputs,247 kwargs=kwargs,248 )249 250 def on_tool_end(251 self,252 output: Any,253 *,254 run_id: UUID,255 parent_run_id: UUID | None = None,256 **kwargs: Any,257 ) -> Any:258 self._end(output, run_id=run_id)259 260 def on_tool_error(261 self,262 error: BaseException,263 *,264 run_id: UUID,265 parent_run_id: UUID | None = None,266 **kwargs: Any,267 ) -> Any:268 self._error(error, run_id=run_id)269 