Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
_messages.py379 linesDownload Raw Back to pregel
1from __future__ import annotations2 3from collections.abc import AsyncIterator, Callable, Iterator, Sequence4from dataclasses import fields, is_dataclass5from typing import (6    Any,7    TypeVar,8    cast,9)10from uuid import UUID, uuid411 12from langchain_core.callbacks import BaseCallbackHandler13from langchain_core.messages import BaseMessage14from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, LLMResult15from pydantic import BaseModel16 17from langgraph._internal._constants import NS_SEP18from langgraph.constants import TAG_HIDDEN, TAG_NOSTREAM19from langgraph.pregel.protocol import StreamChunk20from langgraph.types import Command21 22try:23    from langchain_core.tracers._streaming import _StreamingCallbackHandler24except ImportError:25    _StreamingCallbackHandler = object  # type: ignore26 27try:28    from langchain_core.tracers._streaming import _V2StreamingCallbackHandler29except ImportError:30    _V2StreamingCallbackHandler = object  # type: ignore31 32T = TypeVar("T")33Meta = tuple[tuple[str, ...], dict[str, Any]]34 35 36def _state_values(obj: Any) -> Sequence[Any]:37    """Extract top-level field values from a state object (dict, BaseModel, or dataclass)."""38    if isinstance(obj, dict):39        return list(obj.values())40    elif isinstance(obj, BaseModel):41        return [getattr(obj, k) for k in type(obj).model_fields]42    elif is_dataclass(obj) and not isinstance(obj, type):43        return [getattr(obj, f.name) for f in fields(obj)]44    return ()45 46 47class StreamMessagesHandler(BaseCallbackHandler, _StreamingCallbackHandler):48    """A callback handler that implements stream_mode=messages.49 50    Collects messages from:51    (1) chat model stream events; and52    (2) node outputs.53    """54 55    run_inline = True56    """We want this callback to run in the main thread to avoid order/locking issues."""57 58    def __init__(59        self,60        stream: Callable[[StreamChunk], None],61        subgraphs: bool,62        *,63        parent_ns: tuple[str, ...] | None = None,64    ) -> None:65        """Configure the handler to stream messages from LLMs and nodes.66 67        Args:68            stream: A callable that takes a StreamChunk and emits it.69            subgraphs: Whether to emit messages from subgraphs.70            parent_ns: The namespace where the handler was created.71                We keep track of this namespace to allow calls to subgraphs that72                were explicitly requested as a stream with `messages` mode73                configured.74 75        Example:76            parent_ns is used to handle scenarios where the subgraph is explicitly77            streamed with `stream_mode="messages"`.78 79            ```python80            def parent_graph_node():81                # This node is in the parent graph.82                async for event in some_subgraph(..., stream_mode="messages"):83                    do something with event # <-- these events will be emitted84                return ...85 86            parent_graph.invoke(subgraphs=False)87            ```88        """89        self.stream = stream90        self.subgraphs = subgraphs91        self.metadata: dict[UUID, Meta] = {}92        self.seen: set[int | str] = set()93        self.parent_ns = parent_ns94 95    def _emit(self, meta: Meta, message: BaseMessage, *, dedupe: bool = False) -> None:96        if dedupe and message.id in self.seen:97            return98        else:99            if message.id is None:100                message.id = str(uuid4())101            self.seen.add(message.id)102            self.stream((meta[0], "messages", (message, meta[1])))103 104    def _find_and_emit_messages(self, meta: Meta, response: Any) -> None:105        if isinstance(response, BaseMessage):106            self._emit(meta, response, dedupe=True)107        elif isinstance(response, Sequence):108            for value in response:109                if isinstance(value, BaseMessage):110                    self._emit(meta, value, dedupe=True)111        else:112            for value in _state_values(response):113                if isinstance(value, BaseMessage):114                    self._emit(meta, value, dedupe=True)115                elif isinstance(value, Sequence):116                    for item in value:117                        if isinstance(item, BaseMessage):118                            self._emit(meta, item, dedupe=True)119 120    def tap_output_aiter(121        self, run_id: UUID, output: AsyncIterator[T]122    ) -> AsyncIterator[T]:123        return output124 125    def tap_output_iter(self, run_id: UUID, output: Iterator[T]) -> Iterator[T]:126        return output127 128    def on_chat_model_start(129        self,130        serialized: dict[str, Any],131        messages: list[list[BaseMessage]],132        *,133        run_id: UUID,134        parent_run_id: UUID | None = None,135        tags: list[str] | None = None,136        metadata: dict[str, Any] | None = None,137        **kwargs: Any,138    ) -> Any:139        if metadata and (not tags or (TAG_NOSTREAM not in tags)):140            ns = tuple(cast(str, metadata["langgraph_checkpoint_ns"]).split(NS_SEP))[141                :-1142            ]143            if not self.subgraphs and len(ns) > 0 and ns != self.parent_ns:144                return145            if tags:146                if filtered_tags := [t for t in tags if not t.startswith("seq:step")]:147                    metadata["tags"] = filtered_tags148            self.metadata[run_id] = (ns, metadata)149 150    def on_llm_new_token(151        self,152        token: str,153        *,154        chunk: ChatGenerationChunk | None = None,155        run_id: UUID,156        parent_run_id: UUID | None = None,157        tags: list[str] | None = None,158        **kwargs: Any,159    ) -> Any:160        if not isinstance(chunk, ChatGenerationChunk):161            return162        if meta := self.metadata.get(run_id):163            self._emit(meta, chunk.message)164 165    def on_llm_end(166        self,167        response: LLMResult,168        *,169        run_id: UUID,170        parent_run_id: UUID | None = None,171        **kwargs: Any,172    ) -> Any:173        if meta := self.metadata.get(run_id):174            if response.generations and response.generations[0]:175                gen = response.generations[0][0]176                if isinstance(gen, ChatGeneration):177                    self._emit(meta, gen.message, dedupe=True)178        self.metadata.pop(run_id, None)179 180    def on_llm_error(181        self,182        error: BaseException,183        *,184        run_id: UUID,185        parent_run_id: UUID | None = None,186        **kwargs: Any,187    ) -> Any:188        self.metadata.pop(run_id, None)189 190    def on_chain_start(191        self,192        serialized: dict[str, Any],193        inputs: dict[str, Any],194        *,195        run_id: UUID,196        parent_run_id: UUID | None = None,197        tags: list[str] | None = None,198        metadata: dict[str, Any] | None = None,199        **kwargs: Any,200    ) -> Any:201        if (202            metadata203            and kwargs.get("name") == metadata.get("langgraph_node")204            and (not tags or TAG_HIDDEN not in tags)205        ):206            ns = tuple(cast(str, metadata["langgraph_checkpoint_ns"]).split(NS_SEP))[207                :-1208            ]209            if not self.subgraphs and len(ns) > 0:210                return211            self.metadata[run_id] = (ns, metadata)212            for value in _state_values(inputs):213                if isinstance(value, BaseMessage):214                    if value.id is not None:215                        self.seen.add(value.id)216                elif isinstance(value, Sequence) and not isinstance(value, str):217                    for item in value:218                        if isinstance(item, BaseMessage):219                            if item.id is not None:220                                self.seen.add(item.id)221 222    def on_chain_end(223        self,224        response: Any,225        *,226        run_id: UUID,227        parent_run_id: UUID | None = None,228        **kwargs: Any,229    ) -> Any:230        if meta := self.metadata.pop(run_id, None):231            # Handle Command node updates232            if isinstance(response, Command):233                self._find_and_emit_messages(meta, response.update)234            # Handle list of Command updates235            elif isinstance(response, Sequence) and any(236                isinstance(value, Command) for value in response237            ):238                for value in response:239                    if isinstance(value, Command):240                        self._find_and_emit_messages(meta, value.update)241                    else:242                        self._find_and_emit_messages(meta, value)243            # Handle basic updates / streaming244            else:245                self._find_and_emit_messages(meta, response)246 247    def on_chain_error(248        self,249        error: BaseException,250        *,251        run_id: UUID,252        parent_run_id: UUID | None = None,253        **kwargs: Any,254    ) -> Any:255        self.metadata.pop(run_id, None)256 257 258class StreamMessagesHandlerV2(StreamMessagesHandler, _V2StreamingCallbackHandler):259    """v2 variant of `StreamMessagesHandler`.260 261    Declaring `_V2StreamingCallbackHandler` as a base flips262    `BaseChatModel.invoke` to route through `_stream_chat_model_events`263    (firing `on_stream_event`) instead of `_stream` (firing264    `on_llm_new_token`). Inherits `on_stream_event` from the parent,265    which forwards protocol events onto the messages stream channel.266 267    Pregel attaches this class instead of the v1 handler only when268    `StreamingHandler` opts in via the internal269    `CONFIG_KEY_STREAM_MESSAGES_V2` config key; direct270    `graph.stream(stream_mode="messages")` callers keep the v1271    AIMessageChunk shape.272    """273 274    def on_llm_new_token(275        self,276        token: str,277        *,278        chunk: ChatGenerationChunk | None = None,279        run_id: UUID,280        parent_run_id: UUID | None = None,281        tags: list[str] | None = None,282        **kwargs: Any,283    ) -> Any:284        """Intentional no-op — v1 chunks are not used on v2-flagged runs.285 286        The v2 marker already steers `invoke` to the event generator, so287        `on_llm_new_token` should not fire under normal routing. This288        override stays a pass-through (no call to `super()`) to make289        the intent explicit and to guard against any caller (e.g. a290        node that calls `model.stream()` directly, which still fires291        the v1 callback) leaking AIMessageChunks onto a v2-flagged292        messages stream.293        """294        # Intentionally empty: v2 handler does not forward v1 chunks.295 296    def __init__(297        self,298        stream: Callable[[StreamChunk], None],299        subgraphs: bool,300        *,301        parent_ns: tuple[str, ...] | None = None,302    ) -> None:303        super().__init__(stream, subgraphs, parent_ns=parent_ns)304        self._streamed_run_ids: set[UUID] = set()305 306    def on_llm_end(307        self,308        response: LLMResult,309        *,310        run_id: UUID,311        parent_run_id: UUID | None = None,312        **kwargs: Any,313    ) -> Any:314        if meta := self.metadata.get(run_id):315            if response.generations and response.generations[0]:316                gen = response.generations[0][0]317                if isinstance(gen, ChatGeneration):318                    if run_id in self._streamed_run_ids:319                        if gen.message.id is None:320                            gen.message.id = str(uuid4())321                        self.seen.add(gen.message.id)322                    else:323                        self._emit(meta, gen.message, dedupe=True)324        self._streamed_run_ids.discard(run_id)325        self.metadata.pop(run_id, None)326 327    def on_llm_error(328        self,329        error: BaseException,330        *,331        run_id: UUID,332        parent_run_id: UUID | None = None,333        **kwargs: Any,334    ) -> Any:335        self._streamed_run_ids.discard(run_id)336        super().on_llm_error(337            error,338            run_id=run_id,339            parent_run_id=parent_run_id,340            **kwargs,341        )342 343    def on_stream_event(344        self,345        event: dict[str, Any],346        *,347        run_id: UUID,348        parent_run_id: UUID | None = None,349        tags: list[str] | None = None,350        **kwargs: Any,351    ) -> Any:352        """Forward a protocol event from `stream_events(version="v3")` as a messages stream part.353 354        Fires once per `MessagesData` event (`message-start`, per-block355        `content-block-*`, `message-finish`). The transformer layer356        correlates events back to a single `ChatModelStream` via357        `metadata["run_id"]` — attached here so the v1358        `stream_mode="messages"` output (which emits359        `(AIMessageChunk, metadata)` via `on_llm_new_token`) keeps its360        original metadata shape.361 362        Lives on the v2 handler rather than the v1 base: content-block363        events are a v2-only concept, and forwarding them only when the364        v2 handler is attached keeps the message channel's shape365        predictable for v1 callers.366        """367        if meta := self.metadata.get(run_id):368            # Record message_id on message-start so on_chain_end's369            # dedupe skips the finalized AIMessage the node returns370            # (otherwise the messages projection double-counts: once371            # from streaming, once from the chain output).372            if event.get("event") == "message-start":373                self._streamed_run_ids.add(run_id)374                msg_id = event.get("message_id")375                if msg_id:376                    self.seen.add(msg_id)377            v2_meta = {**meta[1], "run_id": str(run_id)}378            self.stream((meta[0], "messages", (event, v2_meta)))379 
codekingpro/portable-devtools · Team Ai