codekingpro/portable-devtools
114k
1"""Internal tracer to power the event stream API."""2 3from __future__ import annotations4 5import asyncio6import contextlib7import logging8from typing import (9 TYPE_CHECKING,10 Any,11 TypedDict,12 TypeVar,13 cast,14)15 16from typing_extensions import NotRequired, override17 18from langchain_core.callbacks.base import AsyncCallbackHandler, BaseCallbackManager19from langchain_core.messages import AIMessageChunk, BaseMessage, BaseMessageChunk20from langchain_core.outputs import (21 ChatGenerationChunk,22 GenerationChunk,23 LLMResult,24)25from langchain_core.runnables import ensure_config26from langchain_core.runnables.schema import (27 CustomStreamEvent,28 EventData,29 StandardStreamEvent,30 StreamEvent,31)32from langchain_core.runnables.utils import (33 Input,34 Output,35 _RootEventFilter,36)37from langchain_core.tracers._streaming import _StreamingCallbackHandler38from langchain_core.tracers.log_stream import (39 LogStreamCallbackHandler,40 RunLog,41 _astream_log_implementation,42)43from langchain_core.tracers.memory_stream import _MemoryStream44from langchain_core.utils.aiter import aclosing45from langchain_core.utils.uuid import uuid746 47if TYPE_CHECKING:48 from collections.abc import AsyncIterator, Iterator, Sequence49 from uuid import UUID50 51 from langchain_core.documents import Document52 from langchain_core.runnables import Runnable, RunnableConfig53 from langchain_core.tracers.log_stream import LogEntry54 55logger = logging.getLogger(__name__)56 57 58class RunInfo(TypedDict):59 """Information about a run.60 61 This is used to keep track of the metadata associated with a run.62 """63 64 name: str65 """The name of the run."""66 67 tags: list[str]68 """The tags associated with the run."""69 70 metadata: dict[str, Any]71 """The metadata associated with the run."""72 73 run_type: str74 """The type of the run."""75 76 inputs: NotRequired[Any]77 """The inputs to the run."""78 79 parent_run_id: UUID | None80 """The ID of the parent run."""81 82 tool_call_id: NotRequired[str | None]83 """The tool call ID associated with the run."""84 85 86def _assign_name(name: str | None, serialized: dict[str, Any] | None) -> str:87 """Assign a name to a run."""88 if name is not None:89 return name90 if serialized is not None:91 if "name" in serialized:92 return cast("str", serialized["name"])93 if "id" in serialized:94 return cast("str", serialized["id"][-1])95 return "Unnamed"96 97 98T = TypeVar("T")99 100 101class _AstreamEventsCallbackHandler(AsyncCallbackHandler, _StreamingCallbackHandler):102 """An implementation of an async callback handler for astream events."""103 104 def __init__(105 self,106 *args: Any,107 include_names: Sequence[str] | None = None,108 include_types: Sequence[str] | None = None,109 include_tags: Sequence[str] | None = None,110 exclude_names: Sequence[str] | None = None,111 exclude_types: Sequence[str] | None = None,112 exclude_tags: Sequence[str] | None = None,113 **kwargs: Any,114 ) -> None:115 """Initialize the tracer."""116 super().__init__(*args, **kwargs)117 # Map of run ID to run info.118 # the entry corresponding to a given run id is cleaned119 # up when each corresponding run ends.120 self.run_map: dict[UUID, RunInfo] = {}121 # The callback event that corresponds to the end of a parent run122 # may be invoked BEFORE the callback event that corresponds to the end123 # of a child run, which results in clean up of run_map.124 # So we keep track of the mapping between children and parent run IDs125 # in a separate container. This container is GCed when the tracer is GCed.126 self.parent_map: dict[UUID, UUID | None] = {}127 128 self.is_tapped: dict[UUID, Any] = {}129 130 # Filter which events will be sent over the queue.131 self.root_event_filter = _RootEventFilter(132 include_names=include_names,133 include_types=include_types,134 include_tags=include_tags,135 exclude_names=exclude_names,136 exclude_types=exclude_types,137 exclude_tags=exclude_tags,138 )139 140 try:141 loop = asyncio.get_event_loop()142 except RuntimeError:143 loop = asyncio.new_event_loop()144 memory_stream = _MemoryStream[StreamEvent](loop)145 self.send_stream = memory_stream.get_send_stream()146 self.receive_stream = memory_stream.get_receive_stream()147 148 def _get_parent_ids(self, run_id: UUID) -> list[str]:149 """Get the parent IDs of a run (non-recursively) cast to strings."""150 parent_ids = []151 152 while parent_id := self.parent_map.get(run_id):153 str_parent_id = str(parent_id)154 if str_parent_id in parent_ids:155 msg = (156 f"Parent ID {parent_id} is already in the parent_ids list. "157 f"This should never happen."158 )159 raise AssertionError(msg)160 parent_ids.append(str_parent_id)161 run_id = parent_id162 163 # Return the parent IDs in reverse order, so that the first164 # parent ID is the root and the last ID is the immediate parent.165 return parent_ids[::-1]166 167 def _send(self, event: StreamEvent, event_type: str) -> None:168 """Send an event to the stream."""169 if self.root_event_filter.include_event(event, event_type):170 self.send_stream.send_nowait(event)171 172 def __aiter__(self) -> AsyncIterator[Any]:173 """Iterate over the receive stream.174 175 Returns:176 An async iterator over the receive stream.177 """178 return self.receive_stream.__aiter__()179 180 async def tap_output_aiter(181 self, run_id: UUID, output: AsyncIterator[T]182 ) -> AsyncIterator[T]:183 """Tap the output aiter.184 185 This method is used to tap the output of a `Runnable` that produces an async186 iterator. It is used to generate stream events for the output of the `Runnable`.187 188 Args:189 run_id: The ID of the run.190 output: The output of the `Runnable`.191 192 Yields:193 The output of the `Runnable`.194 """195 sentinel = object()196 # atomic check and set197 tap = self.is_tapped.setdefault(run_id, sentinel)198 # wait for first chunk199 first = await anext(output, sentinel)200 if first is sentinel:201 return202 # get run info203 run_info = self.run_map.get(run_id)204 if run_info is None:205 # run has finished, don't issue any stream events206 yield cast("T", first)207 return208 if tap is sentinel:209 # if we are the first to tap, issue stream events210 event: StandardStreamEvent = {211 "event": f"on_{run_info['run_type']}_stream",212 "run_id": str(run_id),213 "name": run_info["name"],214 "tags": run_info["tags"],215 "metadata": run_info["metadata"],216 "data": {},217 "parent_ids": self._get_parent_ids(run_id),218 }219 self._send({**event, "data": {"chunk": first}}, run_info["run_type"])220 yield cast("T", first)221 # consume the rest of the output222 async for chunk in output:223 self._send(224 {**event, "data": {"chunk": chunk}},225 run_info["run_type"],226 )227 yield chunk228 else:229 # otherwise just pass through230 yield cast("T", first)231 # consume the rest of the output232 async for chunk in output:233 yield chunk234 235 def tap_output_iter(self, run_id: UUID, output: Iterator[T]) -> Iterator[T]:236 """Tap the output iter.237 238 Args:239 run_id: The ID of the run.240 output: The output of the `Runnable`.241 242 Yields:243 The output of the `Runnable`.244 """245 sentinel = object()246 # atomic check and set247 tap = self.is_tapped.setdefault(run_id, sentinel)248 # wait for first chunk249 first = next(output, sentinel)250 if first is sentinel:251 return252 # get run info253 run_info = self.run_map.get(run_id)254 if run_info is None:255 # run has finished, don't issue any stream events256 yield cast("T", first)257 return258 if tap is sentinel:259 # if we are the first to tap, issue stream events260 event: StandardStreamEvent = {261 "event": f"on_{run_info['run_type']}_stream",262 "run_id": str(run_id),263 "name": run_info["name"],264 "tags": run_info["tags"],265 "metadata": run_info["metadata"],266 "data": {},267 "parent_ids": self._get_parent_ids(run_id),268 }269 self._send({**event, "data": {"chunk": first}}, run_info["run_type"])270 yield cast("T", first)271 # consume the rest of the output272 for chunk in output:273 self._send(274 {**event, "data": {"chunk": chunk}},275 run_info["run_type"],276 )277 yield chunk278 else:279 # otherwise just pass through280 yield cast("T", first)281 # consume the rest of the output282 for chunk in output:283 yield chunk284 285 def _write_run_start_info(286 self,287 run_id: UUID,288 *,289 tags: list[str] | None,290 metadata: dict[str, Any] | None,291 parent_run_id: UUID | None,292 name_: str,293 run_type: str,294 **kwargs: Any,295 ) -> None:296 """Update the run info."""297 info: RunInfo = {298 "tags": tags or [],299 "metadata": metadata or {},300 "name": name_,301 "run_type": run_type,302 "parent_run_id": parent_run_id,303 }304 305 if "inputs" in kwargs:306 # Handle inputs in a special case to allow inputs to be an307 # optionally provided and distinguish between missing value308 # vs. None value.309 info["inputs"] = kwargs["inputs"]310 311 if "tool_call_id" in kwargs:312 # Store tool_call_id in run info for linking errors to tool calls313 info["tool_call_id"] = kwargs["tool_call_id"]314 315 self.run_map[run_id] = info316 self.parent_map[run_id] = parent_run_id317 318 @override319 async def on_chat_model_start(320 self,321 serialized: dict[str, Any],322 messages: list[list[BaseMessage]],323 *,324 run_id: UUID,325 tags: list[str] | None = None,326 parent_run_id: UUID | None = None,327 metadata: dict[str, Any] | None = None,328 name: str | None = None,329 **kwargs: Any,330 ) -> None:331 """Start a trace for a chat model run."""332 name_ = _assign_name(name, serialized)333 run_type = "chat_model"334 335 self._write_run_start_info(336 run_id,337 tags=tags,338 metadata=metadata,339 parent_run_id=parent_run_id,340 name_=name_,341 run_type=run_type,342 inputs={"messages": messages},343 )344 345 self._send(346 {347 "event": "on_chat_model_start",348 "data": {349 "input": {"messages": messages},350 },351 "name": name_,352 "tags": tags or [],353 "run_id": str(run_id),354 "metadata": metadata or {},355 "parent_ids": self._get_parent_ids(run_id),356 },357 run_type,358 )359 360 @override361 async def on_llm_start(362 self,363 serialized: dict[str, Any],364 prompts: list[str],365 *,366 run_id: UUID,367 tags: list[str] | None = None,368 parent_run_id: UUID | None = None,369 metadata: dict[str, Any] | None = None,370 name: str | None = None,371 **kwargs: Any,372 ) -> None:373 """Start a trace for a (non-chat model) LLM run."""374 name_ = _assign_name(name, serialized)375 run_type = "llm"376 377 self._write_run_start_info(378 run_id,379 tags=tags,380 metadata=metadata,381 parent_run_id=parent_run_id,382 name_=name_,383 run_type=run_type,384 inputs={"prompts": prompts},385 )386 387 self._send(388 {389 "event": "on_llm_start",390 "data": {391 "input": {392 "prompts": prompts,393 }394 },395 "name": name_,396 "tags": tags or [],397 "run_id": str(run_id),398 "metadata": metadata or {},399 "parent_ids": self._get_parent_ids(run_id),400 },401 run_type,402 )403 404 @override405 async def on_custom_event(406 self,407 name: str,408 data: Any,409 *,410 run_id: UUID,411 tags: list[str] | None = None,412 metadata: dict[str, Any] | None = None,413 **kwargs: Any,414 ) -> None:415 """Generate a custom astream event."""416 event = CustomStreamEvent(417 event="on_custom_event",418 run_id=str(run_id),419 name=name,420 tags=tags or [],421 metadata=metadata or {},422 data=data,423 parent_ids=self._get_parent_ids(run_id),424 )425 self._send(event, name)426 427 @override428 async def on_llm_new_token(429 self,430 token: str,431 *,432 chunk: GenerationChunk | ChatGenerationChunk | None = None,433 run_id: UUID,434 parent_run_id: UUID | None = None,435 **kwargs: Any,436 ) -> None:437 """Run on new output token.438 439 Only available when streaming is enabled.440 441 For both chat models and non-chat models (legacy text-completion LLMs).442 443 Raises:444 ValueError: If the run type is not `llm` or `chat_model`.445 AssertionError: If the run ID is not found in the run map.446 """447 run_info = self.run_map.get(run_id)448 chunk_: GenerationChunk | BaseMessageChunk449 450 if run_info is None:451 msg = f"Run ID {run_id} not found in run map."452 raise AssertionError(msg)453 if self.is_tapped.get(run_id):454 return455 if run_info["run_type"] == "chat_model":456 event = "on_chat_model_stream"457 458 if chunk is None:459 chunk_ = AIMessageChunk(content=token)460 else:461 chunk_ = cast("ChatGenerationChunk", chunk).message462 463 elif run_info["run_type"] == "llm":464 event = "on_llm_stream"465 if chunk is None:466 chunk_ = GenerationChunk(text=token)467 else:468 chunk_ = cast("GenerationChunk", chunk)469 else:470 msg = f"Unexpected run type: {run_info['run_type']}"471 raise ValueError(msg)472 473 self._send(474 {475 "event": event,476 "data": {477 "chunk": chunk_,478 },479 "run_id": str(run_id),480 "name": run_info["name"],481 "tags": run_info["tags"],482 "metadata": run_info["metadata"],483 "parent_ids": self._get_parent_ids(run_id),484 },485 run_info["run_type"],486 )487 488 @override489 async def on_llm_end(490 self, response: LLMResult, *, run_id: UUID, **kwargs: Any491 ) -> None:492 """End a trace for a model run.493 494 For both chat models and non-chat models (legacy text-completion LLMs).495 496 Raises:497 ValueError: If the run type is not `'llm'` or `'chat_model'`.498 """499 run_info = self.run_map.pop(run_id)500 inputs_ = run_info.get("inputs")501 502 generations: list[list[GenerationChunk]] | list[list[ChatGenerationChunk]]503 output: dict | BaseMessage = {}504 505 if run_info["run_type"] == "chat_model":506 generations = cast("list[list[ChatGenerationChunk]]", response.generations)507 for gen in generations:508 if output != {}:509 break510 for chunk in gen:511 output = chunk.message512 break513 514 event = "on_chat_model_end"515 elif run_info["run_type"] == "llm":516 generations = cast("list[list[GenerationChunk]]", response.generations)517 output = {518 "generations": [519 [520 {521 "text": chunk.text,522 "generation_info": chunk.generation_info,523 "type": chunk.type,524 }525 for chunk in gen526 ]527 for gen in generations528 ],529 "llm_output": response.llm_output,530 }531 event = "on_llm_end"532 else:533 msg = f"Unexpected run type: {run_info['run_type']}"534 raise ValueError(msg)535 536 self._send(537 {538 "event": event,539 "data": {"output": output, "input": inputs_},540 "run_id": str(run_id),541 "name": run_info["name"],542 "tags": run_info["tags"],543 "metadata": run_info["metadata"],544 "parent_ids": self._get_parent_ids(run_id),545 },546 run_info["run_type"],547 )548 549 async def on_chain_start(550 self,551 serialized: dict[str, Any],552 inputs: dict[str, Any],553 *,554 run_id: UUID,555 tags: list[str] | None = None,556 parent_run_id: UUID | None = None,557 metadata: dict[str, Any] | None = None,558 run_type: str | None = None,559 name: str | None = None,560 **kwargs: Any,561 ) -> None:562 """Start a trace for a chain run."""563 name_ = _assign_name(name, serialized)564 run_type_ = run_type or "chain"565 566 data: EventData = {}567 568 # Work-around Runnable core code not sending input in some569 # cases.570 if inputs != {"input": ""}:571 data["input"] = inputs572 kwargs["inputs"] = inputs573 574 self._write_run_start_info(575 run_id,576 tags=tags,577 metadata=metadata,578 parent_run_id=parent_run_id,579 name_=name_,580 run_type=run_type_,581 **kwargs,582 )583 584 self._send(585 {586 "event": f"on_{run_type_}_start",587 "data": data,588 "name": name_,589 "tags": tags or [],590 "run_id": str(run_id),591 "metadata": metadata or {},592 "parent_ids": self._get_parent_ids(run_id),593 },594 run_type_,595 )596 597 @override598 async def on_chain_end(599 self,600 outputs: dict[str, Any],601 *,602 run_id: UUID,603 inputs: dict[str, Any] | None = None,604 **kwargs: Any,605 ) -> None:606 """End a trace for a chain run."""607 run_info = self.run_map.pop(run_id)608 run_type = run_info["run_type"]609 610 event = f"on_{run_type}_end"611 612 inputs = inputs or run_info.get("inputs") or {}613 614 data: EventData = {615 "output": outputs,616 "input": inputs,617 }618 619 self._send(620 {621 "event": event,622 "data": data,623 "run_id": str(run_id),624 "name": run_info["name"],625 "tags": run_info["tags"],626 "metadata": run_info["metadata"],627 "parent_ids": self._get_parent_ids(run_id),628 },629 run_type,630 )631 632 def _get_tool_run_info_with_inputs(self, run_id: UUID) -> tuple[RunInfo, Any]:633 """Get run info for a tool and extract inputs, with validation.634 635 Args:636 run_id: The run ID of the tool.637 638 Returns:639 A tuple of `(run_info, inputs)`.640 641 Raises:642 AssertionError: If the run ID is a tool call and does not have inputs.643 """644 run_info = self.run_map.pop(run_id)645 if "inputs" not in run_info:646 msg = (647 f"Run ID {run_id} is a tool call and is expected to have "648 f"inputs associated with it."649 )650 raise AssertionError(msg)651 inputs = run_info["inputs"]652 return run_info, inputs653 654 @override655 async def on_tool_start(656 self,657 serialized: dict[str, Any],658 input_str: str,659 *,660 run_id: UUID,661 tags: list[str] | None = None,662 parent_run_id: UUID | None = None,663 metadata: dict[str, Any] | None = None,664 name: str | None = None,665 inputs: dict[str, Any] | None = None,666 **kwargs: Any,667 ) -> None:668 """Start a trace for a tool run."""669 name_ = _assign_name(name, serialized)670 671 self._write_run_start_info(672 run_id,673 tags=tags,674 metadata=metadata,675 parent_run_id=parent_run_id,676 name_=name_,677 run_type="tool",678 inputs=inputs,679 tool_call_id=kwargs.get("tool_call_id"),680 )681 682 self._send(683 {684 "event": "on_tool_start",685 "data": {686 "input": inputs or {},687 },688 "name": name_,689 "tags": tags or [],690 "run_id": str(run_id),691 "metadata": metadata or {},692 "parent_ids": self._get_parent_ids(run_id),693 },694 "tool",695 )696 697 @override698 async def on_tool_error(699 self,700 error: BaseException,701 *,702 run_id: UUID,703 parent_run_id: UUID | None = None,704 tags: list[str] | None = None,705 **kwargs: Any,706 ) -> None:707 """Run when tool errors."""708 # Extract tool_call_id from kwargs if passed directly, or from run_info709 # (which was stored during on_tool_start) as a fallback710 tool_call_id = kwargs.get("tool_call_id")711 run_info, inputs = self._get_tool_run_info_with_inputs(run_id)712 if tool_call_id is None:713 tool_call_id = run_info.get("tool_call_id")714 715 event: StandardStreamEvent = {716 "event": "on_tool_error",717 "data": {718 "error": error,719 "input": inputs,720 "tool_call_id": tool_call_id,721 },722 "run_id": str(run_id),723 "name": run_info["name"],724 "tags": run_info["tags"],725 "metadata": run_info["metadata"],726 "parent_ids": self._get_parent_ids(run_id),727 }728 self._send(event, "tool")729 730 @override731 async def on_tool_end(self, output: Any, *, run_id: UUID, **kwargs: Any) -> None:732 """End a trace for a tool run."""733 run_info, inputs = self._get_tool_run_info_with_inputs(run_id)734 735 self._send(736 {737 "event": "on_tool_end",738 "data": {739 "output": output,740 "input": inputs,741 },742 "run_id": str(run_id),743 "name": run_info["name"],744 "tags": run_info["tags"],745 "metadata": run_info["metadata"],746 "parent_ids": self._get_parent_ids(run_id),747 },748 "tool",749 )750 751 @override752 async def on_retriever_start(753 self,754 serialized: dict[str, Any],755 query: str,756 *,757 run_id: UUID,758 parent_run_id: UUID | None = None,759 tags: list[str] | None = None,760 metadata: dict[str, Any] | None = None,761 name: str | None = None,762 **kwargs: Any,763 ) -> None:764 """Run when `Retriever` starts running."""765 name_ = _assign_name(name, serialized)766 run_type = "retriever"767 768 self._write_run_start_info(769 run_id,770 tags=tags,771 metadata=metadata,772 parent_run_id=parent_run_id,773 name_=name_,774 run_type=run_type,775 inputs={"query": query},776 )777 778 self._send(779 {780 "event": "on_retriever_start",781 "data": {782 "input": {783 "query": query,784 }785 },786 "name": name_,787 "tags": tags or [],788 "run_id": str(run_id),789 "metadata": metadata or {},790 "parent_ids": self._get_parent_ids(run_id),791 },792 run_type,793 )794 795 @override796 async def on_retriever_end(797 self, documents: Sequence[Document], *, run_id: UUID, **kwargs: Any798 ) -> None:799 """Run when `Retriever` ends running."""800 run_info = self.run_map.pop(run_id)801 802 self._send(803 {804 "event": "on_retriever_end",805 "data": {806 "output": documents,807 "input": run_info.get("inputs"),808 },809 "run_id": str(run_id),810 "name": run_info["name"],811 "tags": run_info["tags"],812 "metadata": run_info["metadata"],813 "parent_ids": self._get_parent_ids(run_id),814 },815 run_info["run_type"],816 )817 818 def __deepcopy__(self, memo: dict) -> _AstreamEventsCallbackHandler:819 """Return self."""820 return self821 822 def __copy__(self) -> _AstreamEventsCallbackHandler:823 """Return self."""824 return self825 826 827async def _astream_events_implementation_v1(828 runnable: Runnable[Input, Output],829 value: Any,830 config: RunnableConfig | None = None,831 *,832 include_names: Sequence[str] | None = None,833 include_types: Sequence[str] | None = None,834 include_tags: Sequence[str] | None = None,835 exclude_names: Sequence[str] | None = None,836 exclude_types: Sequence[str] | None = None,837 exclude_tags: Sequence[str] | None = None,838 **kwargs: Any,839) -> AsyncIterator[StandardStreamEvent]:840 stream = LogStreamCallbackHandler(841 auto_close=False,842 include_names=include_names,843 include_types=include_types,844 include_tags=include_tags,845 exclude_names=exclude_names,846 exclude_types=exclude_types,847 exclude_tags=exclude_tags,848 _schema_format="streaming_events",849 )850 851 run_log = RunLog(state=None) # type: ignore[arg-type]852 encountered_start_event = False853 854 root_event_filter = _RootEventFilter(855 include_names=include_names,856 include_types=include_types,857 include_tags=include_tags,858 exclude_names=exclude_names,859 exclude_types=exclude_types,860 exclude_tags=exclude_tags,861 )862 863 config = ensure_config(config)864 root_tags = config.get("tags", [])865 root_metadata = config.get("metadata", {})866 root_name = config.get("run_name", runnable.get_name())867 868 async for log in _astream_log_implementation(869 runnable,870 value,871 config=config,872 stream=stream,873 diff=True,874 with_streamed_output_list=True,875 **kwargs,876 ):877 run_log += log878 879 if not encountered_start_event:880 # Yield the start event for the root runnable.881 encountered_start_event = True882 state = run_log.state.copy()883 884 event = StandardStreamEvent(885 event=f"on_{state['type']}_start",886 run_id=state["id"],887 name=root_name,888 tags=root_tags,889 metadata=root_metadata,890 data={891 "input": value,892 },893 parent_ids=[], # Not supported in v1894 )895 896 if root_event_filter.include_event(event, state["type"]):897 yield event898 899 paths = {900 op["path"].split("/")[2]901 for op in log.ops902 if op["path"].startswith("/logs/")903 }904 # Elements in a set should be iterated in the same order905 # as they were inserted in modern python versions.906 for path in paths:907 data: EventData = {}908 log_entry: LogEntry = run_log.state["logs"][path]909 if log_entry["end_time"] is None:910 event_type = "stream" if log_entry["streamed_output"] else "start"911 else:912 event_type = "end"913 914 if event_type == "start":915 # Include the inputs with the start event if they are available.916 # Usually they will NOT be available for components that operate917 # on streams, since those components stream the input and918 # don't know its final value until the end of the stream.919 inputs = log_entry.get("inputs")920 if inputs is not None:921 data["input"] = inputs922 923 if event_type == "end":924 inputs = log_entry.get("inputs")925 if inputs is not None:926 data["input"] = inputs927 928 # None is a VALID output for an end event929 data["output"] = log_entry["final_output"]930 931 if event_type == "stream":932 num_chunks = len(log_entry["streamed_output"])933 if num_chunks != 1:934 msg = (935 f"Expected exactly one chunk of streamed output, "936 f"got {num_chunks} instead. This is impossible. "937 f"Encountered in: {log_entry['name']}"938 )939 raise AssertionError(msg)940 941 data = {"chunk": log_entry["streamed_output"][0]}942 # Clean up the stream, we don't need it anymore.943 # And this avoids duplicates as well!944 log_entry["streamed_output"] = []945 946 yield StandardStreamEvent(947 event=f"on_{log_entry['type']}_{event_type}",948 name=log_entry["name"],949 run_id=log_entry["id"],950 tags=log_entry["tags"],951 metadata=log_entry["metadata"],952 data=data,953 parent_ids=[], # Not supported in v1954 )955 956 # Finally, we take care of the streaming output from the root chain957 # if there is any.958 state = run_log.state959 if state["streamed_output"]:960 num_chunks = len(state["streamed_output"])961 if num_chunks != 1:962 msg = (963 f"Expected exactly one chunk of streamed output, "964 f"got {num_chunks} instead. This is impossible. "965 f"Encountered in: {state['name']}"966 )967 raise AssertionError(msg)968 969 data = {"chunk": state["streamed_output"][0]}970 # Clean up the stream, we don't need it anymore.971 state["streamed_output"] = []972 973 event = StandardStreamEvent(974 event=f"on_{state['type']}_stream",975 run_id=state["id"],976 tags=root_tags,977 metadata=root_metadata,978 name=root_name,979 data=data,980 parent_ids=[], # Not supported in v1981 )982 if root_event_filter.include_event(event, state["type"]):983 yield event984 985 state = run_log.state986 987 # Finally yield the end event for the root runnable.988 event = StandardStreamEvent(989 event=f"on_{state['type']}_end",990 name=root_name,991 run_id=state["id"],992 tags=root_tags,993 metadata=root_metadata,994 data={995 "output": state["final_output"],996 },997 parent_ids=[], # Not supported in v1998 )999 if root_event_filter.include_event(event, state["type"]):1000 yield event1001 1002 1003async def _astream_events_implementation_v2(1004 runnable: Runnable[Input, Output],1005 value: Any,1006 config: RunnableConfig | None = None,1007 *,1008 include_names: Sequence[str] | None = None,1009 include_types: Sequence[str] | None = None,1010 include_tags: Sequence[str] | None = None,1011 exclude_names: Sequence[str] | None = None,1012 exclude_types: Sequence[str] | None = None,1013 exclude_tags: Sequence[str] | None = None,1014 **kwargs: Any,1015) -> AsyncIterator[StandardStreamEvent]:1016 """Implementation of the astream events API for v2 runnables."""1017 event_streamer = _AstreamEventsCallbackHandler(1018 include_names=include_names,1019 include_types=include_types,1020 include_tags=include_tags,1021 exclude_names=exclude_names,1022 exclude_types=exclude_types,1023 exclude_tags=exclude_tags,1024 )1025 1026 # Assign the stream handler to the config1027 config = ensure_config(config)1028 if "run_id" in config:1029 run_id = cast("UUID", config["run_id"])1030 else:1031 run_id = uuid7()1032 config["run_id"] = run_id1033 callbacks = config.get("callbacks")1034 if callbacks is None:1035 config["callbacks"] = [event_streamer]1036 elif isinstance(callbacks, list):1037 config["callbacks"] = [*callbacks, event_streamer]1038 elif isinstance(callbacks, BaseCallbackManager):1039 callbacks = callbacks.copy()1040 callbacks.add_handler(event_streamer, inherit=True)1041 config["callbacks"] = callbacks1042 else:1043 msg = (1044 f"Unexpected type for callbacks: {callbacks}."1045 "Expected None, list or AsyncCallbackManager."1046 )1047 raise ValueError(msg)1048 1049 # Call the runnable in streaming mode,1050 # add each chunk to the output stream1051 async def consume_astream() -> None:1052 try:1053 # if astream also calls tap_output_aiter this will be a no-op1054 async with aclosing(runnable.astream(value, config, **kwargs)) as stream:1055 async for _ in event_streamer.tap_output_aiter(run_id, stream):1056 # All the content will be picked up1057 pass1058 finally:1059 await event_streamer.send_stream.aclose()1060 1061 # Start the runnable in a task, so we can start consuming output1062 task = asyncio.create_task(consume_astream())1063 1064 first_event_sent = False1065 first_event_run_id = None1066 1067 try:1068 async for event in event_streamer:1069 if not first_event_sent:1070 first_event_sent = True1071 # This is a work-around an issue where the inputs into the1072 # chain are not available until the entire input is consumed.1073 # As a temporary solution, we'll modify the input to be the input1074 # that was passed into the chain.1075 event["data"]["input"] = value1076 first_event_run_id = event["run_id"]1077 yield event1078 continue1079 1080 # If it's the end event corresponding to the root runnable1081 # we don't include the input in the event since it's guaranteed1082 # to be included in the first event.1083 if (1084 event["run_id"] == first_event_run_id1085 and event["event"].endswith("_end")1086 and "input" in event["data"]1087 ):1088 del event["data"]["input"]1089 1090 yield event1091 except asyncio.CancelledError as exc:1092 # Cancel the task if it's still running1093 task.cancel(exc.args[0] if exc.args else None)1094 raise1095 finally:1096 # Cancel the task if it's still running1097 task.cancel()1098 # Await it anyway, to run any cleanup code, and propagate any exceptions1099 with contextlib.suppress(asyncio.CancelledError):1100 await task1101 