Team Ai
Datasetpublic

codekingpro/portable-devtools

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