codekingpro/portable-devtools
114k
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 