Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
log_stream.py770 linesDownload Raw Back to tracers
1"""Tracer that streams run logs to a stream."""2 3from __future__ import annotations4 5import asyncio6import contextlib7import copy8import threading9from collections import defaultdict10from pprint import pformat11from typing import (12    TYPE_CHECKING,13    Any,14    Literal,15    TypeVar,16    overload,17)18 19import jsonpatch  # type: ignore[import-untyped]20from typing_extensions import NotRequired, TypedDict, override21 22from langchain_core.callbacks.base import BaseCallbackManager23from langchain_core.load import dumps24from langchain_core.load.load import load25from langchain_core.outputs import ChatGenerationChunk, GenerationChunk26from langchain_core.runnables import RunnableConfig, ensure_config27from langchain_core.tracers._streaming import _StreamingCallbackHandler28from langchain_core.tracers.base import BaseTracer29from langchain_core.tracers.memory_stream import _MemoryStream30 31if TYPE_CHECKING:32    from collections.abc import AsyncIterator, Iterator, Sequence33    from uuid import UUID34 35    from langchain_core.runnables import Runnable36    from langchain_core.runnables.utils import Input, Output37    from langchain_core.tracers.schemas import Run38 39 40class LogEntry(TypedDict):41    """A single entry in the run log."""42 43    id: str44    """ID of the sub-run."""45 46    name: str47    """Name of the object being run."""48 49    type: str50    """Type of the object being run, eg. prompt, chain, llm, etc."""51 52    tags: list[str]53    """List of tags for the run."""54 55    metadata: dict[str, Any]56    """Key-value pairs of metadata for the run."""57 58    start_time: str59    """ISO-8601 timestamp of when the run started."""60 61    streamed_output_str: list[str]62    """List of LLM tokens streamed by this run, if applicable."""63 64    streamed_output: list[Any]65    """List of output chunks streamed by this run, if available."""66 67    inputs: NotRequired[Any | None]68    """Inputs to this run. Not available currently via `astream_log`."""69 70    final_output: Any | None71    """Final output of this run.72 73    Only available after the run has finished successfully.74    """75 76    end_time: str | None77    """ISO-8601 timestamp of when the run ended.78 79    Only available after the run has finished.80    """81 82 83class RunState(TypedDict):84    """State of the run."""85 86    id: str87    """ID of the run."""88 89    streamed_output: list[Any]90    """List of output chunks streamed by `Runnable.stream()`"""91 92    final_output: Any | None93    """Final output of the run, usually the result of aggregating (`+`) streamed_output.94 95    Updated throughout the run when supported by the `Runnable`.96    """97 98    name: str99    """Name of the object being run."""100 101    type: str102    """Type of the object being run, e.g. prompt, chain, llm, etc."""103 104    # Do we want tags/metadata on the root run? Client kinda knows it in most situations105    # tags: list[str]106 107    logs: dict[str, LogEntry]108    """Map of run names to sub-runs.109 110    If filters were supplied, this list will contain only the runs that matched the111    filters.112    """113 114 115class RunLogPatch:116    """Patch to the run log."""117 118    ops: list[dict[str, Any]]119    """List of `JSONPatch` operations, which describe how to create the run state120    from an empty dict.121 122    This is the minimal representation of the log, designed to be serialized as JSON and123    sent over the wire to reconstruct the log on the other side. Reconstruction of the124    state can be done with any JSONPatch-compliant library, see https://jsonpatch.com125    for more information.126    """127 128    def __init__(self, *ops: dict[str, Any]) -> None:129        """Create a RunLogPatch.130 131        Args:132            *ops: The operations to apply to the state.133        """134        self.ops = list(ops)135 136    def __add__(self, other: RunLogPatch | Any) -> RunLog:137        """Combine two `RunLogPatch` instances.138 139        Args:140            other: The other `RunLogPatch` to combine with.141 142        Raises:143            TypeError: If the other object is not a `RunLogPatch`.144 145        Returns:146            A new `RunLog` representing the combination of the two.147        """148        if type(other) is RunLogPatch:149            ops = self.ops + other.ops150            state = jsonpatch.apply_patch(None, copy.deepcopy(ops))151            return RunLog(*ops, state=state)152 153        msg = f"unsupported operand type(s) for +: '{type(self)}' and '{type(other)}'"154        raise TypeError(msg)155 156    @override157    def __repr__(self) -> str:158        # 1:-1 to get rid of the [] around the list159        return f"RunLogPatch({pformat(self.ops)[1:-1]})"160 161    @override162    def __eq__(self, other: object) -> bool:163        return isinstance(other, RunLogPatch) and self.ops == other.ops164 165    __hash__ = None  # type: ignore[assignment]166 167 168class RunLog(RunLogPatch):169    """Run log."""170 171    state: RunState172    """Current state of the log, obtained from applying all ops in sequence."""173 174    def __init__(self, *ops: dict[str, Any], state: RunState) -> None:175        """Create a RunLog.176 177        Args:178            *ops: The operations to apply to the state.179            state: The initial state of the run log.180        """181        super().__init__(*ops)182        self.state = state183 184    def __add__(self, other: RunLogPatch | Any) -> RunLog:185        """Combine two `RunLog` objects.186 187        Args:188            other: The other `RunLog` or `RunLogPatch` to combine with.189 190        Raises:191            TypeError: If the other object is not a `RunLog` or `RunLogPatch`.192 193        Returns:194            A new `RunLog` representing the combination of the two.195        """196        if type(other) is RunLogPatch:197            ops = self.ops + other.ops198            state = jsonpatch.apply_patch(self.state, other.ops)199            return RunLog(*ops, state=state)200 201        msg = f"unsupported operand type(s) for +: '{type(self)}' and '{type(other)}'"202        raise TypeError(msg)203 204    @override205    def __repr__(self) -> str:206        return f"RunLog({pformat(self.state)})"207 208    @override209    def __eq__(self, other: object) -> bool:210        """Check if two `RunLog`s are equal.211 212        Args:213            other: The other `RunLog` to compare to.214 215        Returns:216            `True` if the `RunLog`s are equal, `False` otherwise.217        """218        # First compare that the state is the same219        if not isinstance(other, RunLog):220            return False221        if self.state != other.state:222            return False223        # Then compare that the ops are the same224        return super().__eq__(other)225 226    __hash__ = None227 228 229T = TypeVar("T")230 231 232class LogStreamCallbackHandler(BaseTracer, _StreamingCallbackHandler):233    """Tracer that streams run logs to a stream."""234 235    def __init__(236        self,237        *,238        auto_close: bool = True,239        include_names: Sequence[str] | None = None,240        include_types: Sequence[str] | None = None,241        include_tags: Sequence[str] | None = None,242        exclude_names: Sequence[str] | None = None,243        exclude_types: Sequence[str] | None = None,244        exclude_tags: Sequence[str] | None = None,245        # Schema format is for internal use only.246        _schema_format: Literal["original", "streaming_events"] = "streaming_events",247    ) -> None:248        """A tracer that streams run logs to a stream.249 250        Args:251            auto_close: Whether to close the stream when the root run finishes.252            include_names: Only include runs from `Runnable` objects with matching253                names.254            include_types: Only include runs from `Runnable` objects with matching255                types.256            include_tags: Only include runs from `Runnable` objects with matching tags.257            exclude_names: Exclude runs from `Runnable` objects with matching names.258            exclude_types: Exclude runs from `Runnable` objects with matching types.259            exclude_tags: Exclude runs from `Runnable` objects with matching tags.260            _schema_format: Primarily changes how the inputs and outputs are handled.261 262                **For internal use only. This API will change.**263 264                - `'original'` is the format used by all current tracers. This format is265                    slightly inconsistent with respect to inputs and outputs.266                - 'streaming_events' is used for supporting streaming events, for267                    internal usage. It will likely change in the future,268                    or be deprecated entirely in favor of a dedicated async269                    tracer for streaming events.270 271        Raises:272            ValueError: If an invalid schema format is provided (internal use only).273        """274        if _schema_format not in {"original", "streaming_events"}:275            msg = (276                f"Invalid schema format: {_schema_format}. "277                f"Expected one of 'original', 'streaming_events'."278            )279            raise ValueError(msg)280        super().__init__(_schema_format=_schema_format)281 282        self.auto_close = auto_close283        self.include_names = include_names284        self.include_types = include_types285        self.include_tags = include_tags286        self.exclude_names = exclude_names287        self.exclude_types = exclude_types288        self.exclude_tags = exclude_tags289 290        try:291            loop = asyncio.get_event_loop()292        except RuntimeError:293            loop = asyncio.new_event_loop()294        memory_stream = _MemoryStream[RunLogPatch](loop)295        self.lock = threading.Lock()296        self.send_stream = memory_stream.get_send_stream()297        self.receive_stream = memory_stream.get_receive_stream()298        self._key_map_by_run_id: dict[UUID, str] = {}299        self._counter_map_by_name: dict[str, int] = defaultdict(int)300        self.root_id: UUID | None = None301 302    def __aiter__(self) -> AsyncIterator[RunLogPatch]:303        """Iterate over the stream of run logs.304 305        Returns:306            An async iterator over the run log patches.307        """308        return self.receive_stream.__aiter__()309 310    def send(self, *ops: dict[str, Any]) -> bool:311        """Send a patch to the stream, return `False` if the stream is closed.312 313        Args:314            *ops: The operations to send to the stream.315 316        Returns:317            `True` if the patch was sent successfully, `False` if the stream is closed.318        """319        # We will likely want to wrap this in try / except at some point320        # to handle exceptions that might arise at run time.321        # For now we'll let the exception bubble up, and always return322        # True on the happy path.323        self.send_stream.send_nowait(RunLogPatch(*ops))324        return True325 326    async def tap_output_aiter(327        self, run_id: UUID, output: AsyncIterator[T]328    ) -> AsyncIterator[T]:329        """Tap an output async iterator to stream its values to the log.330 331        Args:332            run_id: The ID of the run.333            output: The output async iterator.334 335        Yields:336            The output value.337        """338        async for chunk in output:339            # root run is handled in .astream_log()340            # if we can't find the run silently ignore341            # eg. because this run wasn't included in the log342            if (343                run_id != self.root_id344                and (key := self._key_map_by_run_id.get(run_id))345                and (346                    not self.send(347                        {348                            "op": "add",349                            "path": f"/logs/{key}/streamed_output/-",350                            "value": chunk,351                        }352                    )353                )354            ):355                break356 357            yield chunk358 359    def tap_output_iter(self, run_id: UUID, output: Iterator[T]) -> Iterator[T]:360        """Tap an output iterator to stream its values to the log.361 362        Args:363            run_id: The ID of the run.364            output: The output iterator.365 366        Yields:367            The output value.368        """369        for chunk in output:370            # root run is handled in .astream_log()371            # if we can't find the run silently ignore372            # eg. because this run wasn't included in the log373            if (374                run_id != self.root_id375                and (key := self._key_map_by_run_id.get(run_id))376                and (377                    not self.send(378                        {379                            "op": "add",380                            "path": f"/logs/{key}/streamed_output/-",381                            "value": chunk,382                        }383                    )384                )385            ):386                break387 388            yield chunk389 390    def include_run(self, run: Run) -> bool:391        """Check if a `Run` should be included in the log.392 393        Args:394            run: The `Run` to check.395 396        Returns:397            `True` if the `Run` should be included, `False` otherwise.398        """399        if run.id == self.root_id:400            return False401 402        run_tags = run.tags or []403 404        if (405            self.include_names is None406            and self.include_types is None407            and self.include_tags is None408        ):409            include = True410        else:411            include = False412 413        if self.include_names is not None:414            include = include or run.name in self.include_names415        if self.include_types is not None:416            include = include or run.run_type in self.include_types417        if self.include_tags is not None:418            include = include or any(tag in self.include_tags for tag in run_tags)419 420        if self.exclude_names is not None:421            include = include and run.name not in self.exclude_names422        if self.exclude_types is not None:423            include = include and run.run_type not in self.exclude_types424        if self.exclude_tags is not None:425            include = include and all(tag not in self.exclude_tags for tag in run_tags)426 427        return include428 429    def _persist_run(self, run: Run) -> None:430        # This is a legacy method only called once for an entire run tree431        # therefore not useful here432        pass433 434    def _on_run_create(self, run: Run) -> None:435        """Start a run."""436        if self.root_id is None:437            self.root_id = run.id438            if not self.send(439                {440                    "op": "replace",441                    "path": "",442                    "value": RunState(443                        id=str(run.id),444                        streamed_output=[],445                        final_output=None,446                        logs={},447                        name=run.name,448                        type=run.run_type,449                    ),450                }451            ):452                return453 454        if not self.include_run(run):455            return456 457        # Determine previous index, increment by 1458        with self.lock:459            self._counter_map_by_name[run.name] += 1460            count = self._counter_map_by_name[run.name]461            self._key_map_by_run_id[run.id] = (462                run.name if count == 1 else f"{run.name}:{count}"463            )464 465        entry = LogEntry(466            id=str(run.id),467            name=run.name,468            type=run.run_type,469            tags=run.tags or [],470            metadata=(run.extra or {}).get("metadata", {}),471            start_time=run.start_time.isoformat(timespec="milliseconds"),472            streamed_output=[],473            streamed_output_str=[],474            final_output=None,475            end_time=None,476        )477 478        if self._schema_format == "streaming_events":479            # If using streaming events let's add inputs as well480            entry["inputs"] = _get_standardized_inputs(run, self._schema_format)481 482        # Add the run to the stream483        self.send(484            {485                "op": "add",486                "path": f"/logs/{self._key_map_by_run_id[run.id]}",487                "value": entry,488            }489        )490 491    def _on_run_update(self, run: Run) -> None:492        """Finish a `Run`."""493        try:494            index = self._key_map_by_run_id.get(run.id)495 496            if index is None:497                return498 499            ops = []500 501            if self._schema_format == "streaming_events":502                ops.append(503                    {504                        "op": "replace",505                        "path": f"/logs/{index}/inputs",506                        "value": _get_standardized_inputs(run, self._schema_format),507                    }508                )509 510            ops.extend(511                [512                    # Replace 'inputs' with final inputs513                    # This is needed because in many cases the inputs are not514                    # known until after the run is finished and the entire515                    # input stream has been processed by the runnable.516                    {517                        "op": "add",518                        "path": f"/logs/{index}/final_output",519                        # to undo the dumpd done by some runnables / tracer / etc520                        "value": _get_standardized_outputs(run, self._schema_format),521                    },522                    {523                        "op": "add",524                        "path": f"/logs/{index}/end_time",525                        "value": run.end_time.isoformat(timespec="milliseconds")526                        if run.end_time is not None527                        else None,528                    },529                ]530            )531 532            self.send(*ops)533        finally:534            if run.id == self.root_id and self.auto_close:535                self.send_stream.close()536 537    def _on_llm_new_token(538        self,539        run: Run,540        token: str,541        chunk: GenerationChunk | ChatGenerationChunk | None,542    ) -> None:543        """Process new LLM token."""544        index = self._key_map_by_run_id.get(run.id)545 546        if index is None:547            return548 549        self.send(550            {551                "op": "add",552                "path": f"/logs/{index}/streamed_output_str/-",553                "value": token,554            },555            {556                "op": "add",557                "path": f"/logs/{index}/streamed_output/-",558                "value": chunk.message559                if isinstance(chunk, ChatGenerationChunk)560                else token,561            },562        )563 564 565def _get_standardized_inputs(566    run: Run, schema_format: Literal["original", "streaming_events"]567) -> Any:568    """Extract standardized inputs from a `Run`.569 570    Standardizes the inputs based on the type of the runnable used.571 572    Args:573        run: `Run` object574        schema_format: The schema format to use.575 576    Returns:577        Valid inputs are only dict. By conventions, inputs always represented invocation578            using named arguments. `None` means that the input is not yet known!579    """580    if schema_format == "original":581        msg = (582            "Do not assign inputs with original schema drop the key for now."583            "When inputs are added to astream_log they should be added with "584            "standardized schema for streaming events."585        )586        raise NotImplementedError(msg)587 588    inputs = load(run.inputs, allowed_objects="messages")589 590    if run.run_type in {"retriever", "llm", "chat_model"}:591        return inputs592 593    # new style chains594    # These nest an additional 'input' key inside the 'inputs' to make sure595    # the input is always a dict. We need to unpack and use the inner value.596    inputs = inputs["input"]597    # We should try to fix this in Runnables and callbacks/tracers598    # Runnables should be using a None type here not a placeholder599    # dict.600    if inputs == {"input": ""}:  # Workaround for Runnables not using None601        # The input is not known, so we don't assign data['input']602        return None603    return inputs604 605 606def _get_standardized_outputs(607    run: Run, schema_format: Literal["original", "streaming_events", "original+chat"]608) -> Any | None:609    """Extract standardized output from a run.610 611    Standardizes the outputs based on the type of the runnable used.612 613    Args:614        run: the run object.615        schema_format: The schema format to use.616 617    Returns:618        An output if returned, otherwise `None`.619    """620    outputs = load(run.outputs, allowed_objects="messages")621    if schema_format == "original":622        if run.run_type == "prompt" and "output" in outputs:623            # These were previously dumped before the tracer.624            # Now we needn't do anything to them.625            return outputs["output"]626        # Return the old schema, without standardizing anything627        return outputs628 629    if run.run_type in {"retriever", "llm", "chat_model"}:630        return outputs631 632    if isinstance(outputs, dict):633        return outputs.get("output", None)634 635    return None636 637 638@overload639def _astream_log_implementation(640    runnable: Runnable[Input, Output],641    value: Any,642    config: RunnableConfig | None = None,643    *,644    stream: LogStreamCallbackHandler,645    diff: Literal[True] = True,646    with_streamed_output_list: bool = True,647    **kwargs: Any,648) -> AsyncIterator[RunLogPatch]: ...649 650 651@overload652def _astream_log_implementation(653    runnable: Runnable[Input, Output],654    value: Any,655    config: RunnableConfig | None = None,656    *,657    stream: LogStreamCallbackHandler,658    diff: Literal[False],659    with_streamed_output_list: bool = True,660    **kwargs: Any,661) -> AsyncIterator[RunLog]: ...662 663 664async def _astream_log_implementation(665    runnable: Runnable[Input, Output],666    value: Any,667    config: RunnableConfig | None = None,668    *,669    stream: LogStreamCallbackHandler,670    diff: bool = True,671    with_streamed_output_list: bool = True,672    **kwargs: Any,673) -> AsyncIterator[RunLogPatch] | AsyncIterator[RunLog]:674    """Implementation of astream_log for a given runnable.675 676    The implementation has been factored out (at least temporarily) as both677    `astream_log` and `astream_events` rely on it.678 679    Args:680        runnable: The runnable to run in streaming mode.681        value: The input to the runnable.682        config: The config to pass to the runnable.683        stream: The stream to send the run logs to.684        diff: Whether to yield run log patches (`True`) or full run logs (`False`).685        with_streamed_output_list: Whether to include a list of all streamed outputs in686            each patch. If `False`, only the final output will be included in the687            patches.688        **kwargs: Additional keyword arguments to pass to the `Runnable`.689 690    Raises:691        ValueError: If the callbacks in the config are of an unexpected type.692 693    Yields:694        The run log patches or states, depending on the value of `diff`.695    """696    # Assign the stream handler to the config697    config = ensure_config(config)698    callbacks = config.get("callbacks")699    if callbacks is None:700        config["callbacks"] = [stream]701    elif isinstance(callbacks, list):702        config["callbacks"] = [*callbacks, stream]703    elif isinstance(callbacks, BaseCallbackManager):704        callbacks = callbacks.copy()705        callbacks.add_handler(stream, inherit=True)706        config["callbacks"] = callbacks707    else:708        msg = (709            f"Unexpected type for callbacks: {callbacks}."710            "Expected None, list or AsyncCallbackManager."711        )712        raise ValueError(msg)713 714    # Call the runnable in streaming mode,715    # add each chunk to the output stream716    async def consume_astream() -> None:717        try:718            prev_final_output: Output | None = None719            final_output: Output | None = None720 721            async for chunk in runnable.astream(value, config, **kwargs):722                prev_final_output = final_output723                if final_output is None:724                    final_output = chunk725                else:726                    try:727                        final_output = final_output + chunk  # type: ignore[operator]728                    except TypeError:729                        prev_final_output = None730                        final_output = chunk731                patches: list[dict[str, Any]] = []732                if with_streamed_output_list:733                    patches.append(734                        {735                            "op": "add",736                            "path": "/streamed_output/-",737                            # chunk cannot be shared between738                            # streamed_output and final_output739                            # otherwise jsonpatch.apply will740                            # modify both741                            "value": copy.deepcopy(chunk),742                        }743                    )744                patches.extend(745                    {**op, "path": f"/final_output{op['path']}"}746                    for op in jsonpatch.JsonPatch.from_diff(747                        prev_final_output, final_output, dumps=dumps748                    )749                )750                await stream.send_stream.send(RunLogPatch(*patches))751        finally:752            await stream.send_stream.aclose()753 754    # Start the runnable in a task, so we can start consuming output755    task = asyncio.create_task(consume_astream())756    try:757        # Yield each chunk from the output stream758        if diff:759            async for log in stream:760                yield log761        else:762            state = RunLog(state=None)  # type: ignore[arg-type]763            async for log in stream:764                state += log765                yield state766    finally:767        # Wait for the runnable to finish, if not cancelled (eg. by break)768        with contextlib.suppress(asyncio.CancelledError):769            await task770 
codekingpro/portable-devtools · Team Ai