codekingpro/portable-devtools
114k
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 