Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
_tools.py269 linesDownload Raw Back to pregel
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 
codekingpro/portable-devtools · Team Ai