Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
config.py673 linesDownload Raw Back to runnables
1"""Configuration utilities for `Runnable` objects."""2 3from __future__ import annotations4 5import asyncio6 7# Cannot move uuid to TYPE_CHECKING as RunnableConfig is used in Pydantic models8import uuid  # noqa: TC0039import warnings10from collections.abc import Awaitable, Callable, Generator, Iterable, Iterator, Sequence11from concurrent.futures import Executor, Future, ThreadPoolExecutor12from contextlib import contextmanager13from contextvars import Context, ContextVar, Token, copy_context14from functools import partial15from typing import (16    TYPE_CHECKING,17    Any,18    ParamSpec,19    TypeVar,20    cast,21)22 23from typing_extensions import TypedDict24 25from langchain_core.callbacks.manager import AsyncCallbackManager, CallbackManager26from langchain_core.runnables.utils import (27    Input,28    Output,29    accepts_config,30    accepts_run_manager,31)32 33if TYPE_CHECKING:34    from langchain_core.callbacks.base import BaseCallbackManager, Callbacks35    from langchain_core.callbacks.manager import (36        AsyncCallbackManagerForChainRun,37        CallbackManagerForChainRun,38    )39else:40    # Pydantic validates through typed dicts, but41    # the callbacks need forward refs updated42    Callbacks = list | Any | None43 44 45class EmptyDict(TypedDict, total=False):46    """Empty dict type."""47 48 49class RunnableConfig(TypedDict, total=False):50    """Configuration for a `Runnable`.51 52    !!! note Custom values53 54        The `TypedDict` has `total=False` set intentionally to:55 56        - Allow partial configs to be created and merged together via `merge_configs`57        - Support config propagation from parent to child runnables via58            `var_child_runnable_config` (a `ContextVar` that automatically passes59            config down the call stack without explicit parameter passing), where60            configs are merged rather than replaced61 62        !!! example63 64            ```python65            # Parent sets tags66            chain.invoke(input, config={"tags": ["parent"]})67            # Child automatically inherits and can add:68            # ensure_config({"tags": ["child"]}) -> {"tags": ["parent", "child"]}69            ```70    """71 72    tags: list[str]73    """Tags for this call and any sub-calls (e.g. a Chain calling an LLM).74 75    You can use these to filter calls.76    """77 78    metadata: dict[str, Any]79    """Metadata for this call and any sub-calls (e.g. a Chain calling an LLM).80 81    Keys should be strings, values should be JSON-serializable.82    """83 84    callbacks: Callbacks85    """Callbacks for this call and any sub-calls (e.g. a Chain calling an LLM).86 87    Tags are passed to all callbacks, metadata is passed to handle*Start callbacks.88    """89 90    run_name: str91    """Name for the tracer run for this call.92 93    Defaults to the name of the class."""94 95    max_concurrency: int | None96    """Maximum number of parallel calls to make.97 98    If not provided, defaults to `ThreadPoolExecutor`'s default.99    """100 101    recursion_limit: int102    """Maximum number of times a call can recurse.103 104    If not provided, defaults to `25`.105    """106 107    configurable: dict[str, Any]108    """Runtime values for attributes previously made configurable on this `Runnable`,109    or sub-`Runnable` objects, through `configurable_fields` or110    `configurable_alternatives`.111 112    Check `output_schema` for a description of the attributes that have been made113    configurable.114    """115 116    run_id: uuid.UUID | None117    """Unique identifier for the tracer run for this call.118 119    If not provided, a new UUID will be generated.120    """121 122 123CONFIG_KEYS = [124    "tags",125    "metadata",126    "callbacks",127    "run_name",128    "max_concurrency",129    "recursion_limit",130    "configurable",131    "run_id",132]133 134COPIABLE_KEYS = [135    "tags",136    "metadata",137    "callbacks",138    "configurable",139]140 141 142# Users are expected to use the `context` API with a context object143# (which does not get traced)144CONFIGURABLE_TO_TRACING_METADATA_EXCLUDED_KEYS = frozenset(("api_key",))145 146 147def _get_langsmith_inheritable_metadata_from_config(148    config: RunnableConfig,149) -> dict[str, Any] | None:150    """Get LangSmith-only inheritable metadata defaults derived from config."""151    configurable = config.get("configurable") or {}152    metadata = {153        key: value154        for key, value in configurable.items()155        if not key.startswith("__")156        and isinstance(value, (str, int, float, bool))157        and key not in config.get("metadata", {})158        and key not in CONFIGURABLE_TO_TRACING_METADATA_EXCLUDED_KEYS159    }160    return metadata or None161 162 163DEFAULT_RECURSION_LIMIT = 25164 165 166var_child_runnable_config: ContextVar[RunnableConfig | None] = ContextVar(167    "child_runnable_config", default=None168)169 170 171# This is imported and used in langgraph, so don't break.172def _set_config_context(173    config: RunnableConfig,174) -> tuple[Token[RunnableConfig | None], dict[str, Any] | None]:175    """Set the child Runnable config + tracing context.176 177    Args:178        config: The config to set.179 180    Returns:181        The token to reset the config and the previous tracing context.182    """183    # Deferred to avoid importing langsmith at module level (~132ms).184    from langsmith.run_helpers import (  # noqa: PLC0415185        _set_tracing_context,186        get_tracing_context,187    )188 189    from langchain_core.tracers.langchain import LangChainTracer  # noqa: PLC0415190 191    config_token = var_child_runnable_config.set(config)192    current_context = None193    if (194        (callbacks := config.get("callbacks"))195        and (196            parent_run_id := getattr(callbacks, "parent_run_id", None)197        )  # Is callback manager198        and (199            tracer := next(200                (201                    handler202                    for handler in getattr(callbacks, "handlers", [])203                    if isinstance(handler, LangChainTracer)204                ),205                None,206            )207        )208        and (run := tracer.run_map.get(str(parent_run_id)))209    ):210        current_context = get_tracing_context()211        _set_tracing_context({"parent": run})212    return config_token, current_context213 214 215@contextmanager216def set_config_context(config: RunnableConfig) -> Generator[Context, None, None]:217    """Set the child Runnable config + tracing context.218 219    Args:220        config: The config to set.221 222    Yields:223        The config context.224    """225    # Deferred to avoid importing langsmith at module level (~132ms).226    from langsmith.run_helpers import _set_tracing_context  # noqa: PLC0415227 228    ctx = copy_context()229    config_token, _ = ctx.run(_set_config_context, config)230    try:231        yield ctx232    finally:233        ctx.run(var_child_runnable_config.reset, config_token)234        ctx.run(235            _set_tracing_context,236            {237                "parent": None,238                "project_name": None,239                "tags": None,240                "metadata": None,241                "enabled": None,242                "client": None,243            },244        )245 246 247def ensure_config(config: RunnableConfig | None = None) -> RunnableConfig:248    """Ensure that a config is a dict with all keys present.249 250    Args:251        config: The config to ensure.252 253    Returns:254        The ensured config.255    """256    empty = RunnableConfig(257        tags=[],258        metadata={},259        callbacks=None,260        recursion_limit=DEFAULT_RECURSION_LIMIT,261        configurable={},262    )263    if var_config := var_child_runnable_config.get():264        empty.update(265            cast(266                "RunnableConfig",267                {268                    k: v.copy() if k in COPIABLE_KEYS else v  # type: ignore[attr-defined]269                    for k, v in var_config.items()270                    if v is not None271                },272            )273        )274    if config is not None:275        empty.update(276            cast(277                "RunnableConfig",278                {279                    k: v.copy() if k in COPIABLE_KEYS else v  # type: ignore[attr-defined]280                    for k, v in config.items()281                    if v is not None and k in CONFIG_KEYS282                },283            )284        )285    if config is not None:286        for k, v in config.items():287            if k not in CONFIG_KEYS and v is not None:288                empty["configurable"][k] = v289    for configurable_key in ("model", "checkpoint_ns"):290        if (291            isinstance(292                configurable_value := empty.get("configurable", {}).get(293                    configurable_key294                ),295                str,296            )297            and configurable_key not in empty["metadata"]298        ):299            empty["metadata"][configurable_key] = configurable_value300    return empty301 302 303def get_config_list(304    config: RunnableConfig | Sequence[RunnableConfig] | None, length: int305) -> list[RunnableConfig]:306    """Get a list of configs from a single config or a list of configs.307 308     It is useful for subclasses overriding batch() or abatch().309 310    Args:311        config: The config or list of configs.312        length: The length of the list.313 314    Returns:315        The list of configs.316 317    Raises:318        ValueError: If the length of the list is not equal to the length of the inputs.319 320    """321    if length < 0:322        msg = f"length must be >= 0, but got {length}"323        raise ValueError(msg)324    if isinstance(config, Sequence) and len(config) != length:325        msg = (326            f"config must be a list of the same length as inputs, "327            f"but got {len(config)} configs for {length} inputs"328        )329        raise ValueError(msg)330 331    if isinstance(config, Sequence):332        return list(map(ensure_config, config))333    if length > 1 and isinstance(config, dict) and config.get("run_id") is not None:334        warnings.warn(335            "Provided run_id be used only for the first element of the batch.",336            category=RuntimeWarning,337            stacklevel=3,338        )339        subsequent = cast(340            "RunnableConfig", {k: v for k, v in config.items() if k != "run_id"}341        )342        return [343            ensure_config(subsequent) if i else ensure_config(config)344            for i in range(length)345        ]346    return [ensure_config(config) for i in range(length)]347 348 349def patch_config(350    config: RunnableConfig | None,351    *,352    callbacks: BaseCallbackManager | None = None,353    recursion_limit: int | None = None,354    max_concurrency: int | None = None,355    run_name: str | None = None,356    configurable: dict[str, Any] | None = None,357) -> RunnableConfig:358    """Patch a config with new values.359 360    Args:361        config: The config to patch.362        callbacks: The callbacks to set.363        recursion_limit: The recursion limit to set.364        max_concurrency: The max concurrency to set.365        run_name: The run name to set.366        configurable: The configurable to set.367 368    Returns:369        The patched config.370    """371    config = ensure_config(config)372    if callbacks is not None:373        # If we're replacing callbacks, we need to unset run_name374        # As that should apply only to the same run as the original callbacks375        config["callbacks"] = callbacks376        if "run_name" in config:377            del config["run_name"]378        if "run_id" in config:379            del config["run_id"]380    if recursion_limit is not None:381        config["recursion_limit"] = recursion_limit382    if max_concurrency is not None:383        config["max_concurrency"] = max_concurrency384    if run_name is not None:385        config["run_name"] = run_name386    if configurable is not None:387        config["configurable"] = {**config.get("configurable", {}), **configurable}388    return config389 390 391def merge_configs(*configs: RunnableConfig | None) -> RunnableConfig:392    """Merge multiple configs into one.393 394    Args:395        *configs: The configs to merge.396 397    Returns:398        The merged config.399    """400    base: RunnableConfig = {}401    # Even though the keys aren't literals, this is correct402    # because both dicts are the same type403    for config in (ensure_config(c) for c in configs if c is not None):404        for key in config:405            if key == "metadata":406                base["metadata"] = {407                    **base.get("metadata", {}),408                    **(config.get("metadata") or {}),409                }410            elif key == "tags":411                base["tags"] = sorted(412                    set(base.get("tags", []) + (config.get("tags") or [])),413                )414            elif key == "configurable":415                base["configurable"] = {416                    **base.get("configurable", {}),417                    **(config.get("configurable") or {}),418                }419            elif key == "callbacks":420                base_callbacks = base.get("callbacks")421                these_callbacks = config["callbacks"]422                # callbacks can be either None, list[handler] or manager423                # so merging two callbacks values has 6 cases424                if isinstance(these_callbacks, list):425                    if base_callbacks is None:426                        base["callbacks"] = these_callbacks.copy()427                    elif isinstance(base_callbacks, list):428                        base["callbacks"] = base_callbacks + these_callbacks429                    else:430                        # base_callbacks is a manager431                        mngr = base_callbacks.copy()432                        for callback in these_callbacks:433                            mngr.add_handler(callback, inherit=True)434                        base["callbacks"] = mngr435                elif these_callbacks is not None:436                    # these_callbacks is a manager437                    if base_callbacks is None:438                        base["callbacks"] = these_callbacks.copy()439                    elif isinstance(base_callbacks, list):440                        mngr = these_callbacks.copy()441                        for callback in base_callbacks:442                            mngr.add_handler(callback, inherit=True)443                        base["callbacks"] = mngr444                    else:445                        # base_callbacks is also a manager446                        base["callbacks"] = base_callbacks.merge(these_callbacks)447            elif key == "recursion_limit":448                if config["recursion_limit"] != DEFAULT_RECURSION_LIMIT:449                    base["recursion_limit"] = config["recursion_limit"]450            elif key in COPIABLE_KEYS and config[key] is not None:  # type: ignore[literal-required]451                base[key] = config[key].copy()  # type: ignore[literal-required]452            else:453                base[key] = config[key] or base.get(key)  # type: ignore[literal-required]454    return base455 456 457def call_func_with_variable_args(458    func: Callable[[Input], Output]459    | Callable[[Input, RunnableConfig], Output]460    | Callable[[Input, CallbackManagerForChainRun], Output]461    | Callable[[Input, CallbackManagerForChainRun, RunnableConfig], Output],462    input: Input,463    config: RunnableConfig,464    run_manager: CallbackManagerForChainRun | None = None,465    **kwargs: Any,466) -> Output:467    """Call function that may optionally accept a run_manager and/or config.468 469    Args:470        func: The function to call.471        input: The input to the function.472        config: The config to pass to the function.473        run_manager: The run manager to pass to the function.474        **kwargs: The keyword arguments to pass to the function.475 476    Returns:477        The output of the function.478    """479    if accepts_config(func):480        if run_manager is not None:481            kwargs["config"] = patch_config(config, callbacks=run_manager.get_child())482        else:483            kwargs["config"] = config484    if run_manager is not None and accepts_run_manager(func):485        kwargs["run_manager"] = run_manager486    return func(input, **kwargs)  # type: ignore[call-arg]487 488 489def acall_func_with_variable_args(490    func: Callable[[Input], Awaitable[Output]]491    | Callable[[Input, RunnableConfig], Awaitable[Output]]492    | Callable[[Input, AsyncCallbackManagerForChainRun], Awaitable[Output]]493    | Callable[494        [Input, AsyncCallbackManagerForChainRun, RunnableConfig], Awaitable[Output]495    ],496    input: Input,497    config: RunnableConfig,498    run_manager: AsyncCallbackManagerForChainRun | None = None,499    **kwargs: Any,500) -> Awaitable[Output]:501    """Async call function that may optionally accept a run_manager and/or config.502 503    Args:504        func: The function to call.505        input: The input to the function.506        config: The config to pass to the function.507        run_manager: The run manager to pass to the function.508        **kwargs: The keyword arguments to pass to the function.509 510    Returns:511        The output of the function.512    """513    if accepts_config(func):514        if run_manager is not None:515            kwargs["config"] = patch_config(config, callbacks=run_manager.get_child())516        else:517            kwargs["config"] = config518    if run_manager is not None and accepts_run_manager(func):519        kwargs["run_manager"] = run_manager520    return func(input, **kwargs)  # type: ignore[call-arg]521 522 523def get_callback_manager_for_config(config: RunnableConfig) -> CallbackManager:524    """Get a callback manager for a config.525 526    Args:527        config: The config.528 529    Returns:530        The callback manager.531    """532    return CallbackManager.configure(533        inheritable_callbacks=config.get("callbacks"),534        inheritable_tags=config.get("tags"),535        inheritable_metadata=config.get("metadata"),536        langsmith_inheritable_metadata=_get_langsmith_inheritable_metadata_from_config(537            config538        ),539    )540 541 542def get_async_callback_manager_for_config(543    config: RunnableConfig,544) -> AsyncCallbackManager:545    """Get an async callback manager for a config.546 547    Args:548        config: The config.549 550    Returns:551        The async callback manager.552    """553    return AsyncCallbackManager.configure(554        inheritable_callbacks=config.get("callbacks"),555        inheritable_tags=config.get("tags"),556        inheritable_metadata=config.get("metadata"),557        langsmith_inheritable_metadata=_get_langsmith_inheritable_metadata_from_config(558            config559        ),560    )561 562 563P = ParamSpec("P")564T = TypeVar("T")565 566 567class ContextThreadPoolExecutor(ThreadPoolExecutor):568    """ThreadPoolExecutor that copies the context to the child thread."""569 570    def submit(  # type: ignore[override]571        self,572        func: Callable[P, T],573        *args: P.args,574        **kwargs: P.kwargs,575    ) -> Future[T]:576        """Submit a function to the executor.577 578        Args:579            func: The function to submit.580            *args: The positional arguments to the function.581            **kwargs: The keyword arguments to the function.582 583        Returns:584            The future for the function.585        """586        return super().submit(587            cast("Callable[..., T]", partial(copy_context().run, func, *args, **kwargs))588        )589 590    def map(591        self,592        fn: Callable[..., T],593        *iterables: Iterable[Any],594        **kwargs: Any,595    ) -> Iterator[T]:596        """Map a function to multiple iterables.597 598        Args:599            fn: The function to map.600            *iterables: The iterables to map over.601            timeout: The timeout for the map.602            chunksize: The chunksize for the map.603 604        Returns:605            The iterator for the mapped function.606        """607        contexts = [copy_context() for _ in range(len(iterables[0]))]  # type: ignore[arg-type]608 609        def _wrapped_fn(*args: Any) -> T:610            return contexts.pop().run(fn, *args)611 612        return super().map(613            _wrapped_fn,614            *iterables,615            **kwargs,616        )617 618 619@contextmanager620def get_executor_for_config(621    config: RunnableConfig | None,622) -> Generator[Executor, None, None]:623    """Get an executor for a config.624 625    Args:626        config: The config.627 628    Yields:629        The executor.630    """631    config = config or {}632    with ContextThreadPoolExecutor(633        max_workers=config.get("max_concurrency")634    ) as executor:635        yield executor636 637 638async def run_in_executor(639    executor_or_config: Executor | RunnableConfig | None,640    func: Callable[P, T],641    *args: P.args,642    **kwargs: P.kwargs,643) -> T:644    """Run a function in an executor.645 646    Args:647        executor_or_config: The executor or config to run in.648        func: The function.649        *args: The positional arguments to the function.650        **kwargs: The keyword arguments to the function.651 652    Returns:653        The output of the function.654    """655 656    def wrapper() -> T:657        try:658            return func(*args, **kwargs)659        except StopIteration as exc:660            # StopIteration can't be set on an asyncio.Future661            # it raises a TypeError and leaves the Future pending forever662            # so we need to convert it to a RuntimeError663            raise RuntimeError from exc664 665    if executor_or_config is None or isinstance(executor_or_config, dict):666        # Use default executor with context copied from current context667        return await asyncio.get_running_loop().run_in_executor(668            None,669            cast("Callable[..., T]", partial(copy_context().run, wrapper)),670        )671 672    return await asyncio.get_running_loop().run_in_executor(executor_or_config, wrapper)673 
codekingpro/portable-devtools · Team Ai