Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
core.py725 linesDownload Raw Back to tracers
1"""Utilities for the root listener."""2 3from __future__ import annotations4 5import logging6import traceback7from abc import ABC, abstractmethod8from datetime import datetime, timezone9from typing import (10    TYPE_CHECKING,11    Any,12    Literal,13    cast,14)15 16from langchain_core.exceptions import TracerException17from langchain_core.load import dumpd18from langchain_core.tracers.schemas import Run19 20if TYPE_CHECKING:21    from collections.abc import Coroutine, Sequence22    from uuid import UUID23 24    from tenacity import RetryCallState25 26    from langchain_core.documents import Document27    from langchain_core.messages import BaseMessage28    from langchain_core.outputs import (29        ChatGeneration,30        ChatGenerationChunk,31        GenerationChunk,32        LLMResult,33    )34 35logger = logging.getLogger(__name__)36 37SCHEMA_FORMAT_TYPE = Literal["original", "streaming_events"]38 39 40class _TracerCore(ABC):41    """Abstract base class for tracers.42 43    This class provides common methods, and reusable methods for tracers.44    """45 46    log_missing_parent: bool = True47 48    def __init__(49        self,50        *,51        _schema_format: Literal[52            "original", "streaming_events", "original+chat"53        ] = "original",54        run_map: dict[str, Run] | None = None,55        order_map: dict[UUID, tuple[UUID, str]] | None = None,56        _external_run_ids: dict[str, int] | None = None,57        **kwargs: Any,58    ) -> None:59        """Initialize the tracer.60 61        Args:62            _schema_format: Primarily changes how the inputs and outputs are handled.63 64                For internal use only. This API will change.65 66                - `'original'` is the format used by all current tracers.67 68                    This format is slightly inconsistent with respect to inputs and69                    outputs.70                - `'streaming_events'` is used for supporting streaming events, for71                    internal usage. It will likely change in the future, or be72                    deprecated entirely in favor of a dedicated async tracer for73                    streaming events.74                - `'original+chat'` is a format that is the same as `'original'` except75                    it does NOT raise an attribute error `on_chat_model_start`76            run_map: Optional shared map of run ID to run.77            order_map: Optional shared map of run ID to trace ordering data.78            _external_run_ids: Optional shared set of externally injected run IDs.79            **kwargs: Additional keyword arguments that will be passed to the80                superclass.81        """82        super().__init__(**kwargs)83 84        self._schema_format = _schema_format  # For internal use only API will change.85 86        self.run_map = run_map if run_map is not None else {}87        """Map of run ID to run. Cleared on run end."""88 89        self.order_map = order_map if order_map is not None else {}90        """Map of run ID to (trace_id, dotted_order). Cleared when tracer GCed."""91 92        self._external_run_ids: dict[str, int] = (93            _external_run_ids if _external_run_ids is not None else {}94        )95        """Refcount of active children per externally-injected run ID.96 97        These runs are added to `run_map` so child runs can find their parent,98        but they are not managed by the tracer's callback lifecycle.  When99        the last child finishes the entry is evicted to avoid memory leaks.100        """101 102    @abstractmethod103    def _persist_run(self, run: Run) -> Coroutine[Any, Any, None] | None:104        """Persist a run."""105 106    @staticmethod107    def _add_child_run(108        parent_run: Run,109        child_run: Run,110    ) -> None:111        """Add child run to a chain run or tool run."""112        parent_run.child_runs.append(child_run)113 114    @staticmethod115    def _get_stacktrace(error: BaseException) -> str:116        """Get the stacktrace of the parent error."""117        msg = repr(error)118        try:119            tb = traceback.format_exception(error)120            return (msg + "\n\n".join(tb)).strip()121        except Exception:122            return msg123 124    def _start_trace(self, run: Run) -> Coroutine[Any, Any, None] | None:  # type: ignore[return]125        current_dotted_order = run.start_time.strftime("%Y%m%dT%H%M%S%fZ") + str(run.id)126        if run.parent_run_id:127            if parent := self.order_map.get(run.parent_run_id):128                run.trace_id, run.dotted_order = parent129                run.dotted_order += "." + current_dotted_order130                if parent_run := self.run_map.get(str(run.parent_run_id)):131                    self._add_child_run(parent_run, run)132                    parent_key = str(run.parent_run_id)133                    if parent_key in self._external_run_ids:134                        self._external_run_ids[parent_key] += 1135            else:136                if self.log_missing_parent:137                    logger.debug(138                        "Parent run %s not found for run %s. Treating as a root run.",139                        run.parent_run_id,140                        run.id,141                    )142                run.parent_run_id = None143                run.trace_id = run.id144                run.dotted_order = current_dotted_order145        else:146            run.trace_id = run.id147            run.dotted_order = current_dotted_order148        self.order_map[run.id] = (run.trace_id, run.dotted_order)149        self.run_map[str(run.id)] = run150 151    def _get_run(self, run_id: UUID, run_type: str | set[str] | None = None) -> Run:152        try:153            run = self.run_map[str(run_id)]154        except KeyError as exc:155            msg = f"No indexed run ID {run_id}."156            raise TracerException(msg) from exc157 158        if isinstance(run_type, str):159            run_types: set[str] | None = {run_type}160        else:161            run_types = run_type162        if run_types is not None and run.run_type not in run_types:163            msg = (164                f"Found {run.run_type} run at ID {run_id}, "165                f"but expected {run_types} run."166            )167            raise TracerException(msg)168        return run169 170    def _create_chat_model_run(171        self,172        serialized: dict[str, Any],173        messages: list[list[BaseMessage]],174        run_id: UUID,175        tags: list[str] | None = None,176        parent_run_id: UUID | None = None,177        metadata: dict[str, Any] | None = None,178        name: str | None = None,179        **kwargs: Any,180    ) -> Run:181        """Create a chat model run."""182        if self._schema_format not in {"streaming_events", "original+chat"}:183            # Please keep this un-implemented for backwards compatibility.184            # When it's unimplemented old tracers that use the "original" format185            # fallback on the on_llm_start method implementation if they186            # find that the on_chat_model_start method is not implemented.187            # This can eventually be cleaned up by writing a "modern" tracer188            # that has all the updated schema changes corresponding to189            # the "streaming_events" format.190            msg = (191                f"Chat model tracing is not supported in "192                f"for {self._schema_format} format."193            )194            raise NotImplementedError(msg)195        start_time = datetime.now(timezone.utc)196        if metadata:197            kwargs.update({"metadata": metadata})198        return Run(199            id=run_id,200            parent_run_id=parent_run_id,201            serialized=serialized,202            inputs={"messages": [[dumpd(msg) for msg in batch] for batch in messages]},203            extra=kwargs,204            events=[{"name": "start", "time": start_time}],205            start_time=start_time,206            # WARNING: This is valid ONLY for streaming_events.207            # run_type="llm" is what's used by virtually all tracers.208            # Changing this to "chat_model" may break triggering on_llm_start209            run_type="chat_model",210            tags=tags,211            name=name,212        )213 214    def _create_llm_run(215        self,216        serialized: dict[str, Any],217        prompts: list[str],218        run_id: UUID,219        tags: list[str] | None = None,220        parent_run_id: UUID | None = None,221        metadata: dict[str, Any] | None = None,222        name: str | None = None,223        **kwargs: Any,224    ) -> Run:225        """Create a llm run."""226        start_time = datetime.now(timezone.utc)227        if metadata:228            kwargs.update({"metadata": metadata})229        return Run(230            id=run_id,231            parent_run_id=parent_run_id,232            serialized=serialized,233            # TODO: Figure out how to expose kwargs here234            inputs={"prompts": prompts},235            extra=kwargs,236            events=[{"name": "start", "time": start_time}],237            start_time=start_time,238            run_type="llm",239            tags=tags or [],240            name=name,241        )242 243    def _llm_run_with_token_event(244        self,245        token: str,246        run_id: UUID,247        chunk: GenerationChunk | ChatGenerationChunk | None = None,248        parent_run_id: UUID | None = None,249    ) -> Run:250        """Append token event to LLM run and return the run."""251        _ = parent_run_id252        llm_run = self._get_run(run_id, run_type={"llm", "chat_model"})253        event_kwargs: dict[str, Any] = {"token": token}254        if chunk:255            event_kwargs["chunk"] = chunk256        llm_run.events.append(257            {258                "name": "new_token",259                "time": datetime.now(timezone.utc),260                "kwargs": event_kwargs,261            },262        )263        return llm_run264 265    def _llm_run_with_retry_event(266        self,267        retry_state: RetryCallState,268        run_id: UUID,269    ) -> Run:270        llm_run = self._get_run(run_id)271        retry_d: dict[str, Any] = {272            "slept": retry_state.idle_for,273            "attempt": retry_state.attempt_number,274        }275        if retry_state.outcome is None:276            retry_d["outcome"] = "N/A"277        elif retry_state.outcome.failed:278            retry_d["outcome"] = "failed"279            exception = retry_state.outcome.exception()280            retry_d["exception"] = str(exception)281            retry_d["exception_type"] = exception.__class__.__name__282        else:283            retry_d["outcome"] = "success"284            retry_d["result"] = str(retry_state.outcome.result())285        llm_run.events.append(286            {287                "name": "retry",288                "time": datetime.now(timezone.utc),289                "kwargs": retry_d,290            },291        )292        return llm_run293 294    def _complete_llm_run(self, response: LLMResult, run_id: UUID) -> Run:295        llm_run = self._get_run(run_id, run_type={"llm", "chat_model"})296        if getattr(llm_run, "outputs", None) is None:297            llm_run.outputs = {}298        else:299            llm_run.outputs = cast("dict[str, Any]", llm_run.outputs)300        if not llm_run.extra.get("__omit_auto_outputs", False):301            llm_run.outputs.update(response.model_dump())302        for i, generations in enumerate(response.generations):303            for j, generation in enumerate(generations):304                output_generation = llm_run.outputs["generations"][i][j]305                if "message" in output_generation:306                    output_generation["message"] = dumpd(307                        cast("ChatGeneration", generation).message308                    )309        llm_run.end_time = datetime.now(timezone.utc)310        llm_run.events.append({"name": "end", "time": llm_run.end_time})311 312        tool_call_count = 0313        for generations in response.generations:314            for generation in generations:315                if hasattr(generation, "message"):316                    msg = generation.message317                    if hasattr(msg, "tool_calls") and msg.tool_calls:318                        tool_call_count += len(msg.tool_calls)319        if tool_call_count > 0:320            llm_run.extra["tool_call_count"] = tool_call_count321 322        return llm_run323 324    def _errored_llm_run(325        self, error: BaseException, run_id: UUID, response: LLMResult | None = None326    ) -> Run:327        llm_run = self._get_run(run_id, run_type={"llm", "chat_model"})328        llm_run.error = self._get_stacktrace(error)329        if response:330            if getattr(llm_run, "outputs", None) is None:331                llm_run.outputs = {}332            else:333                llm_run.outputs = cast("dict[str, Any]", llm_run.outputs)334            if not llm_run.extra.get("__omit_auto_outputs", False):335                llm_run.outputs.update(response.model_dump())336            for i, generations in enumerate(response.generations):337                for j, generation in enumerate(generations):338                    output_generation = llm_run.outputs["generations"][i][j]339                    if "message" in output_generation:340                        output_generation["message"] = dumpd(341                            cast("ChatGeneration", generation).message342                        )343        llm_run.end_time = datetime.now(timezone.utc)344        llm_run.events.append({"name": "error", "time": llm_run.end_time})345 346        return llm_run347 348    def _create_chain_run(349        self,350        serialized: dict[str, Any],351        inputs: dict[str, Any],352        run_id: UUID,353        tags: list[str] | None = None,354        parent_run_id: UUID | None = None,355        metadata: dict[str, Any] | None = None,356        run_type: str | None = None,357        name: str | None = None,358        **kwargs: Any,359    ) -> Run:360        """Create a chain Run."""361        start_time = datetime.now(timezone.utc)362        if metadata:363            kwargs.update({"metadata": metadata})364        return Run(365            id=run_id,366            parent_run_id=parent_run_id,367            serialized=serialized,368            inputs=self._get_chain_inputs(inputs),369            extra=kwargs,370            events=[{"name": "start", "time": start_time}],371            start_time=start_time,372            child_runs=[],373            run_type=run_type or "chain",374            name=name,375            tags=tags or [],376        )377 378    def _get_chain_inputs(self, inputs: Any) -> Any:379        """Get the inputs for a chain run."""380        if self._schema_format in {"original", "original+chat"}:381            return inputs if isinstance(inputs, dict) else {"input": inputs}382        if self._schema_format == "streaming_events":383            return {384                "input": inputs,385            }386        msg = f"Invalid format: {self._schema_format}"387        raise ValueError(msg)388 389    def _get_chain_outputs(self, outputs: Any) -> Any:390        """Get the outputs for a chain run."""391        if self._schema_format in {"original", "original+chat"}:392            return outputs if isinstance(outputs, dict) else {"output": outputs}393        if self._schema_format == "streaming_events":394            return {395                "output": outputs,396            }397        msg = f"Invalid format: {self._schema_format}"398        raise ValueError(msg)399 400    def _complete_chain_run(401        self,402        outputs: dict[str, Any],403        run_id: UUID,404        inputs: dict[str, Any] | None = None,405    ) -> Run:406        """Update a chain run with outputs and end time."""407        chain_run = self._get_run(run_id)408        if getattr(chain_run, "outputs", None) is None:409            chain_run.outputs = {}410        if not chain_run.extra.get("__omit_auto_outputs", False):411            cast("dict[str, Any]", chain_run.outputs).update(412                self._get_chain_outputs(outputs)413            )414        chain_run.end_time = datetime.now(timezone.utc)415        chain_run.events.append({"name": "end", "time": chain_run.end_time})416        if inputs is not None:417            chain_run.inputs = self._get_chain_inputs(inputs)418        return chain_run419 420    def _errored_chain_run(421        self,422        error: BaseException,423        inputs: dict[str, Any] | None,424        run_id: UUID,425    ) -> Run:426        chain_run = self._get_run(run_id)427        chain_run.error = self._get_stacktrace(error)428        chain_run.end_time = datetime.now(timezone.utc)429        chain_run.events.append({"name": "error", "time": chain_run.end_time})430        if inputs is not None:431            chain_run.inputs = self._get_chain_inputs(inputs)432        return chain_run433 434    def _create_tool_run(435        self,436        serialized: dict[str, Any],437        input_str: str,438        run_id: UUID,439        tags: list[str] | None = None,440        parent_run_id: UUID | None = None,441        metadata: dict[str, Any] | None = None,442        name: str | None = None,443        inputs: dict[str, Any] | None = None,444        **kwargs: Any,445    ) -> Run:446        """Create a tool run."""447        start_time = datetime.now(timezone.utc)448        if metadata:449            kwargs.update({"metadata": metadata})450 451        if self._schema_format in {"original", "original+chat"}:452            inputs = inputs if isinstance(inputs, dict) else {"input": input_str}453        elif self._schema_format == "streaming_events":454            inputs = {"input": inputs}455        else:456            msg = f"Invalid format: {self._schema_format}"457            raise AssertionError(msg)458 459        return Run(460            id=run_id,461            parent_run_id=parent_run_id,462            serialized=serialized,463            # Wrapping in dict since Run requires a dict object.464            inputs=inputs,465            extra=kwargs,466            events=[{"name": "start", "time": start_time}],467            start_time=start_time,468            child_runs=[],469            run_type="tool",470            tags=tags or [],471            name=name,472        )473 474    def _complete_tool_run(475        self,476        output: dict[str, Any],477        run_id: UUID,478    ) -> Run:479        """Update a tool run with outputs and end time."""480        tool_run = self._get_run(run_id, run_type="tool")481        if getattr(tool_run, "outputs", None) is None:482            tool_run.outputs = {}483        if not tool_run.extra.get("__omit_auto_outputs", False):484            cast("dict[str, Any]", tool_run.outputs).update({"output": output})485        tool_run.end_time = datetime.now(timezone.utc)486        tool_run.events.append({"name": "end", "time": tool_run.end_time})487        return tool_run488 489    def _errored_tool_run(490        self,491        error: BaseException,492        run_id: UUID,493    ) -> Run:494        """Update a tool run with error and end time."""495        tool_run = self._get_run(run_id, run_type="tool")496        tool_run.error = self._get_stacktrace(error)497        tool_run.end_time = datetime.now(timezone.utc)498        tool_run.events.append({"name": "error", "time": tool_run.end_time})499        return tool_run500 501    def _create_retrieval_run(502        self,503        serialized: dict[str, Any],504        query: str,505        run_id: UUID,506        parent_run_id: UUID | None = None,507        tags: list[str] | None = None,508        metadata: dict[str, Any] | None = None,509        name: str | None = None,510        **kwargs: Any,511    ) -> Run:512        """Create a retrieval run."""513        start_time = datetime.now(timezone.utc)514        if metadata:515            kwargs.update({"metadata": metadata})516        return Run(517            id=run_id,518            name=name or "Retriever",519            parent_run_id=parent_run_id,520            serialized=serialized,521            inputs={"query": query},522            extra=kwargs,523            events=[{"name": "start", "time": start_time}],524            start_time=start_time,525            tags=tags,526            child_runs=[],527            run_type="retriever",528        )529 530    def _complete_retrieval_run(531        self,532        documents: Sequence[Document],533        run_id: UUID,534    ) -> Run:535        """Update a retrieval run with outputs and end time."""536        retrieval_run = self._get_run(run_id, run_type="retriever")537        if getattr(retrieval_run, "outputs", None) is None:538            retrieval_run.outputs = {}539        if not retrieval_run.extra.get("__omit_auto_outputs", False):540            cast("dict[str, Any]", retrieval_run.outputs).update(541                {"documents": documents}542            )543        retrieval_run.end_time = datetime.now(timezone.utc)544        retrieval_run.events.append({"name": "end", "time": retrieval_run.end_time})545        return retrieval_run546 547    def _errored_retrieval_run(548        self,549        error: BaseException,550        run_id: UUID,551    ) -> Run:552        retrieval_run = self._get_run(run_id, run_type="retriever")553        retrieval_run.error = self._get_stacktrace(error)554        retrieval_run.end_time = datetime.now(timezone.utc)555        retrieval_run.events.append({"name": "error", "time": retrieval_run.end_time})556        return retrieval_run557 558    def __deepcopy__(self, memo: dict) -> _TracerCore:559        """Return self deepcopied."""560        return self561 562    def __copy__(self) -> _TracerCore:563        """Return self copied."""564        return self565 566    def _end_trace(self, run: Run) -> Coroutine[Any, Any, None] | None:567        """End a trace for a run.568 569        Args:570            run: The run.571        """572        _ = run573        return None574 575    def _on_run_create(self, run: Run) -> Coroutine[Any, Any, None] | None:576        """Process a run upon creation.577 578        Args:579            run: The created run.580        """581        _ = run582        return None583 584    def _on_run_update(self, run: Run) -> Coroutine[Any, Any, None] | None:585        """Process a run upon update.586 587        Args:588            run: The updated run.589        """590        _ = run591        return None592 593    def _on_llm_start(self, run: Run) -> Coroutine[Any, Any, None] | None:594        """Process the LLM Run upon start.595 596        Args:597            run: The LLM run.598        """599        _ = run600        return None601 602    def _on_llm_new_token(603        self,604        run: Run,605        token: str,606        chunk: GenerationChunk | ChatGenerationChunk | None,607    ) -> Coroutine[Any, Any, None] | None:608        """Process new LLM token.609 610        Args:611            run: The LLM run.612            token: The new token.613            chunk: Optional chunk.614        """615        _ = (run, token, chunk)616        return None617 618    def _on_llm_end(self, run: Run) -> Coroutine[Any, Any, None] | None:619        """Process the LLM Run.620 621        Args:622            run: The LLM run.623        """624        _ = run625        return None626 627    def _on_llm_error(self, run: Run) -> Coroutine[Any, Any, None] | None:628        """Process the LLM Run upon error.629 630        Args:631            run: The LLM run.632        """633        _ = run634        return None635 636    def _on_chain_start(self, run: Run) -> Coroutine[Any, Any, None] | None:637        """Process the Chain Run upon start.638 639        Args:640            run: The chain run.641        """642        _ = run643        return None644 645    def _on_chain_end(self, run: Run) -> Coroutine[Any, Any, None] | None:646        """Process the Chain Run.647 648        Args:649            run: The chain run.650        """651        _ = run652        return None653 654    def _on_chain_error(self, run: Run) -> Coroutine[Any, Any, None] | None:655        """Process the Chain Run upon error.656 657        Args:658            run: The chain run.659        """660        _ = run661        return None662 663    def _on_tool_start(self, run: Run) -> Coroutine[Any, Any, None] | None:664        """Process the Tool Run upon start.665 666        Args:667            run: The tool run.668        """669        _ = run670        return None671 672    def _on_tool_end(self, run: Run) -> Coroutine[Any, Any, None] | None:673        """Process the Tool Run.674 675        Args:676            run: The tool run.677        """678        _ = run679        return None680 681    def _on_tool_error(self, run: Run) -> Coroutine[Any, Any, None] | None:682        """Process the Tool Run upon error.683 684        Args:685            run: The tool run.686        """687        _ = run688        return None689 690    def _on_chat_model_start(self, run: Run) -> Coroutine[Any, Any, None] | None:691        """Process the Chat Model Run upon start.692 693        Args:694            run: The chat model run.695        """696        _ = run697        return None698 699    def _on_retriever_start(self, run: Run) -> Coroutine[Any, Any, None] | None:700        """Process the Retriever Run upon start.701 702        Args:703            run: The retriever run.704        """705        _ = run706        return None707 708    def _on_retriever_end(self, run: Run) -> Coroutine[Any, Any, None] | None:709        """Process the Retriever Run.710 711        Args:712            run: The retriever run.713        """714        _ = run715        return None716 717    def _on_retriever_error(self, run: Run) -> Coroutine[Any, Any, None] | None:718        """Process the Retriever Run upon error.719 720        Args:721            run: The retriever run.722        """723        _ = run724        return None725 
codekingpro/portable-devtools · Team Ai