codekingpro/portable-devtools
114k
1"""Implementation of the `RunnablePassthrough`."""2 3from __future__ import annotations4 5import asyncio6import inspect7import threading8from collections.abc import Awaitable, Callable9from typing import (10 TYPE_CHECKING,11 Any,12 cast,13)14 15from pydantic import BaseModel, RootModel16from typing_extensions import override17 18from langchain_core.runnables.base import (19 Other,20 Runnable,21 RunnableParallel,22 RunnableSerializable,23)24from langchain_core.runnables.config import (25 RunnableConfig,26 acall_func_with_variable_args,27 call_func_with_variable_args,28 ensure_config,29 get_executor_for_config,30 patch_config,31)32from langchain_core.runnables.utils import (33 AddableDict,34 ConfigurableFieldSpec,35)36from langchain_core.utils.aiter import atee37from langchain_core.utils.iter import safetee38from langchain_core.utils.pydantic import create_model_v239 40if TYPE_CHECKING:41 from collections.abc import AsyncIterator, Iterator, Mapping42 43 from langchain_core.callbacks.manager import (44 AsyncCallbackManagerForChainRun,45 CallbackManagerForChainRun,46 )47 from langchain_core.runnables.graph import Graph48 49 50def identity(x: Other) -> Other:51 """Identity function.52 53 Args:54 x: Input.55 56 Returns:57 Output.58 """59 return x60 61 62async def aidentity(x: Other) -> Other:63 """Async identity function.64 65 Args:66 x: Input.67 68 Returns:69 Output.70 """71 return x72 73 74class RunnablePassthrough(RunnableSerializable[Other, Other]):75 """Runnable to passthrough inputs unchanged or with additional keys.76 77 This `Runnable` behaves almost like the identity function, except that it78 can be configured to add additional keys to the output, if the input is a79 dict.80 81 The examples below demonstrate this `Runnable` works using a few simple82 chains. The chains rely on simple lambdas to make the examples easy to execute83 and experiment with.84 85 Examples:86 ```python87 from langchain_core.runnables import (88 RunnableLambda,89 RunnableParallel,90 RunnablePassthrough,91 )92 93 runnable = RunnableParallel(94 origin=RunnablePassthrough(), modified=lambda x: x + 195 )96 97 runnable.invoke(1) # {'origin': 1, 'modified': 2}98 99 100 def fake_llm(prompt: str) -> str: # Fake LLM for the example101 return "completion"102 103 104 chain = RunnableLambda(fake_llm) | {105 "original": RunnablePassthrough(), # Original LLM output106 "parsed": lambda text: text[::-1], # Parsing logic107 }108 109 chain.invoke("hello") # {'original': 'completion', 'parsed': 'noitelpmoc'}110 ```111 112 In some cases, it may be useful to pass the input through while adding some113 keys to the output. In this case, you can use the `assign` method:114 115 ```python116 from langchain_core.runnables import RunnablePassthrough117 118 119 def fake_llm(prompt: str) -> str: # Fake LLM for the example120 return "completion"121 122 123 runnable = {124 "llm1": fake_llm,125 "llm2": fake_llm,126 } | RunnablePassthrough.assign(127 total_chars=lambda inputs: len(inputs["llm1"] + inputs["llm2"])128 )129 130 runnable.invoke("hello")131 # {'llm1': 'completion', 'llm2': 'completion', 'total_chars': 20}132 ```133 """134 135 input_type: type[Other] | None = None136 137 func: Callable[[Other], None] | Callable[[Other, RunnableConfig], None] | None = (138 None139 )140 141 afunc: (142 Callable[[Other], Awaitable[None]]143 | Callable[[Other, RunnableConfig], Awaitable[None]]144 | None145 ) = None146 147 @override148 def __repr_args__(self) -> Any:149 # Without this repr(self) raises a RecursionError150 # See https://github.com/pydantic/pydantic/issues/7327151 return []152 153 def __init__(154 self,155 func: Callable[[Other], None]156 | Callable[[Other, RunnableConfig], None]157 | Callable[[Other], Awaitable[None]]158 | Callable[[Other, RunnableConfig], Awaitable[None]]159 | None = None,160 afunc: Callable[[Other], Awaitable[None]]161 | Callable[[Other, RunnableConfig], Awaitable[None]]162 | None = None,163 *,164 input_type: type[Other] | None = None,165 **kwargs: Any,166 ) -> None:167 """Create a `RunnablePassthrough`.168 169 Args:170 func: Function to be called with the input.171 afunc: Async function to be called with the input.172 input_type: Type of the input.173 """174 if inspect.iscoroutinefunction(func):175 afunc = func176 func = None177 178 super().__init__(func=func, afunc=afunc, input_type=input_type, **kwargs)179 180 @classmethod181 @override182 def is_lc_serializable(cls) -> bool:183 """Return `True` as this class is serializable."""184 return True185 186 @classmethod187 def get_lc_namespace(cls) -> list[str]:188 """Get the namespace of the LangChain object.189 190 Returns:191 `["langchain", "schema", "runnable"]`192 """193 return ["langchain", "schema", "runnable"]194 195 @property196 @override197 def InputType(self) -> Any:198 return self.input_type or Any199 200 @property201 @override202 def OutputType(self) -> Any:203 return self.input_type or Any204 205 @classmethod206 @override207 def assign(208 cls,209 **kwargs: Runnable[dict[str, Any], Any]210 | Callable[[dict[str, Any]], Any]211 | Mapping[str, Runnable[dict[str, Any], Any] | Callable[[dict[str, Any]], Any]],212 ) -> RunnableAssign:213 """Merge the Dict input with the output produced by the mapping argument.214 215 Args:216 **kwargs: `Runnable`, `Callable` or a `Mapping` from keys to `Runnable`217 objects or `Callable`s.218 219 Returns:220 A `Runnable` that merges the `dict` input with the output produced by the221 mapping argument.222 """223 return RunnableAssign(RunnableParallel[dict[str, Any]](kwargs))224 225 @override226 def invoke(227 self, input: Other, config: RunnableConfig | None = None, **kwargs: Any228 ) -> Other:229 if self.func is not None:230 call_func_with_variable_args(231 self.func, input, ensure_config(config), **kwargs232 )233 return self._call_with_config(identity, input, config)234 235 @override236 async def ainvoke(237 self,238 input: Other,239 config: RunnableConfig | None = None,240 **kwargs: Any | None,241 ) -> Other:242 if self.afunc is not None:243 await acall_func_with_variable_args(244 self.afunc, input, ensure_config(config), **kwargs245 )246 elif self.func is not None:247 call_func_with_variable_args(248 self.func, input, ensure_config(config), **kwargs249 )250 return await self._acall_with_config(aidentity, input, config)251 252 @override253 def transform(254 self,255 input: Iterator[Other],256 config: RunnableConfig | None = None,257 **kwargs: Any,258 ) -> Iterator[Other]:259 if self.func is None:260 for chunk in self._transform_stream_with_config(input, identity, config):261 yield chunk262 else:263 final: Other264 got_first_chunk = False265 266 for chunk in self._transform_stream_with_config(input, identity, config):267 yield chunk268 269 if not got_first_chunk:270 final = chunk271 got_first_chunk = True272 else:273 try:274 final = final + chunk # type: ignore[operator]275 except TypeError:276 final = chunk277 278 if got_first_chunk:279 call_func_with_variable_args(280 self.func, final, ensure_config(config), **kwargs281 )282 283 @override284 async def atransform(285 self,286 input: AsyncIterator[Other],287 config: RunnableConfig | None = None,288 **kwargs: Any,289 ) -> AsyncIterator[Other]:290 if self.afunc is None and self.func is None:291 async for chunk in self._atransform_stream_with_config(292 input, identity, config293 ):294 yield chunk295 else:296 got_first_chunk = False297 298 async for chunk in self._atransform_stream_with_config(299 input, identity, config300 ):301 yield chunk302 303 # By definitions, a function will operate on the aggregated304 # input. So we'll aggregate the input until we get to the last305 # chunk.306 # If the input is not addable, then we'll assume that we can307 # only operate on the last chunk.308 if not got_first_chunk:309 final = chunk310 got_first_chunk = True311 else:312 try:313 final = final + chunk # type: ignore[operator]314 except TypeError:315 final = chunk316 317 if got_first_chunk:318 config = ensure_config(config)319 if self.afunc is not None:320 await acall_func_with_variable_args(321 self.afunc, final, config, **kwargs322 )323 elif self.func is not None:324 call_func_with_variable_args(self.func, final, config, **kwargs)325 326 @override327 def stream(328 self,329 input: Other,330 config: RunnableConfig | None = None,331 **kwargs: Any,332 ) -> Iterator[Other]:333 return self.transform(iter([input]), config, **kwargs)334 335 @override336 async def astream(337 self,338 input: Other,339 config: RunnableConfig | None = None,340 **kwargs: Any,341 ) -> AsyncIterator[Other]:342 async def input_aiter() -> AsyncIterator[Other]:343 yield input344 345 async for chunk in self.atransform(input_aiter(), config, **kwargs):346 yield chunk347 348 349_graph_passthrough: RunnablePassthrough = RunnablePassthrough()350 351 352class RunnableAssign(RunnableSerializable[dict[str, Any], dict[str, Any]]):353 """Runnable that assigns key-value pairs to `dict[str, Any]` inputs.354 355 The `RunnableAssign` class takes input dictionaries and, through a356 `RunnableParallel` instance, applies transformations, then combines357 these with the original data, introducing new key-value pairs based358 on the mapper's logic.359 360 Examples:361 ```python362 # This is a RunnableAssign363 from langchain_core.runnables.passthrough import (364 RunnableAssign,365 RunnableParallel,366 )367 from langchain_core.runnables.base import RunnableLambda368 369 370 def add_ten(x: dict[str, int]) -> dict[str, int]:371 return {"added": x["input"] + 10}372 373 374 mapper = RunnableParallel(375 {376 "add_step": RunnableLambda(add_ten),377 }378 )379 380 runnable_assign = RunnableAssign(mapper)381 382 # Synchronous example383 runnable_assign.invoke({"input": 5})384 # returns {'input': 5, 'add_step': {'added': 15}}385 386 # Asynchronous example387 await runnable_assign.ainvoke({"input": 5})388 # returns {'input': 5, 'add_step': {'added': 15}}389 ```390 """391 392 mapper: RunnableParallel393 394 def __init__(self, mapper: RunnableParallel[dict[str, Any]], **kwargs: Any) -> None:395 """Create a `RunnableAssign`.396 397 Args:398 mapper: A `RunnableParallel` instance that will be used to transform the399 input dictionary.400 """401 super().__init__(mapper=mapper, **kwargs)402 403 @classmethod404 @override405 def is_lc_serializable(cls) -> bool:406 """Return `True` as this class is serializable."""407 return True408 409 @classmethod410 @override411 def get_lc_namespace(cls) -> list[str]:412 """Get the namespace of the LangChain object.413 414 Returns:415 `["langchain", "schema", "runnable"]`416 """417 return ["langchain", "schema", "runnable"]418 419 @override420 def get_name(self, suffix: str | None = None, *, name: str | None = None) -> str:421 name = (422 name423 or self.name424 or f"RunnableAssign<{','.join(self.mapper.steps__.keys())}>"425 )426 return super().get_name(suffix, name=name)427 428 @override429 def get_input_schema(self, config: RunnableConfig | None = None) -> type[BaseModel]:430 map_input_schema = self.mapper.get_input_schema(config)431 if not issubclass(map_input_schema, RootModel):432 # ie. it's a dict433 return map_input_schema434 435 return super().get_input_schema(config)436 437 @override438 def get_output_schema(439 self, config: RunnableConfig | None = None440 ) -> type[BaseModel]:441 map_input_schema = self.mapper.get_input_schema(config)442 map_output_schema = self.mapper.get_output_schema(config)443 if not issubclass(map_input_schema, RootModel) and not issubclass(444 map_output_schema, RootModel445 ):446 fields = {}447 448 for name, field_info in map_input_schema.model_fields.items():449 fields[name] = (field_info.annotation, field_info.default)450 451 for name, field_info in map_output_schema.model_fields.items():452 fields[name] = (field_info.annotation, field_info.default)453 454 return create_model_v2("RunnableAssignOutput", field_definitions=fields)455 if not issubclass(map_output_schema, RootModel):456 # ie. only map output is a dict457 # ie. input type is either unknown or inferred incorrectly458 return map_output_schema459 460 return super().get_output_schema(config)461 462 @property463 @override464 def config_specs(self) -> list[ConfigurableFieldSpec]:465 return self.mapper.config_specs466 467 @override468 def get_graph(self, config: RunnableConfig | None = None) -> Graph:469 # get graph from mapper470 graph = self.mapper.get_graph(config)471 # add passthrough node and edges472 input_node = graph.first_node()473 output_node = graph.last_node()474 if input_node is not None and output_node is not None:475 passthrough_node = graph.add_node(_graph_passthrough)476 graph.add_edge(input_node, passthrough_node)477 graph.add_edge(passthrough_node, output_node)478 return graph479 480 def _invoke(481 self,482 value: dict[str, Any],483 run_manager: CallbackManagerForChainRun,484 config: RunnableConfig,485 **kwargs: Any,486 ) -> dict[str, Any]:487 if not isinstance(value, dict):488 msg = "The input to RunnablePassthrough.assign() must be a dict."489 raise ValueError(msg) # noqa: TRY004490 491 return {492 **value,493 **self.mapper.invoke(494 value,495 patch_config(config, callbacks=run_manager.get_child()),496 **kwargs,497 ),498 }499 500 @override501 def invoke(502 self,503 input: dict[str, Any],504 config: RunnableConfig | None = None,505 **kwargs: Any,506 ) -> dict[str, Any]:507 return self._call_with_config(self._invoke, input, config, **kwargs)508 509 async def _ainvoke(510 self,511 value: dict[str, Any],512 run_manager: AsyncCallbackManagerForChainRun,513 config: RunnableConfig,514 **kwargs: Any,515 ) -> dict[str, Any]:516 if not isinstance(value, dict):517 msg = "The input to RunnablePassthrough.assign() must be a dict."518 raise ValueError(msg) # noqa: TRY004519 520 return {521 **value,522 **await self.mapper.ainvoke(523 value,524 patch_config(config, callbacks=run_manager.get_child()),525 **kwargs,526 ),527 }528 529 @override530 async def ainvoke(531 self,532 input: dict[str, Any],533 config: RunnableConfig | None = None,534 **kwargs: Any,535 ) -> dict[str, Any]:536 return await self._acall_with_config(self._ainvoke, input, config, **kwargs)537 538 def _transform(539 self,540 values: Iterator[dict[str, Any]],541 run_manager: CallbackManagerForChainRun,542 config: RunnableConfig,543 **kwargs: Any,544 ) -> Iterator[dict[str, Any]]:545 # collect mapper keys546 mapper_keys = set(self.mapper.steps__.keys())547 # create two streams, one for the map and one for the passthrough548 for_passthrough, for_map = safetee(values, 2, lock=threading.Lock())549 550 # create map output stream551 map_output = self.mapper.transform(552 for_map,553 patch_config(554 config,555 callbacks=run_manager.get_child(),556 ),557 **kwargs,558 )559 560 # get executor to start map output stream in background561 with get_executor_for_config(config) as executor:562 # start map output stream563 first_map_chunk_future = executor.submit(564 next,565 map_output,566 None,567 )568 # consume passthrough stream569 for chunk in for_passthrough:570 if not isinstance(chunk, dict):571 msg = "The input to RunnablePassthrough.assign() must be a dict."572 raise ValueError(msg) # noqa: TRY004573 # remove mapper keys from passthrough chunk, to be overwritten by map574 filtered = AddableDict(575 {k: v for k, v in chunk.items() if k not in mapper_keys}576 )577 if filtered:578 yield filtered579 # yield map output580 yield cast("dict[str, Any]", first_map_chunk_future.result())581 for chunk in map_output:582 yield chunk583 584 @override585 def transform(586 self,587 input: Iterator[dict[str, Any]],588 config: RunnableConfig | None = None,589 **kwargs: Any | None,590 ) -> Iterator[dict[str, Any]]:591 yield from self._transform_stream_with_config(592 input, self._transform, config, **kwargs593 )594 595 async def _atransform(596 self,597 values: AsyncIterator[dict[str, Any]],598 run_manager: AsyncCallbackManagerForChainRun,599 config: RunnableConfig,600 **kwargs: Any,601 ) -> AsyncIterator[dict[str, Any]]:602 # collect mapper keys603 mapper_keys = set(self.mapper.steps__.keys())604 # create two streams, one for the map and one for the passthrough605 for_passthrough, for_map = atee(values, 2, lock=asyncio.Lock())606 # create map output stream607 map_output = self.mapper.atransform(608 for_map,609 patch_config(610 config,611 callbacks=run_manager.get_child(),612 ),613 **kwargs,614 )615 # start map output stream616 first_map_chunk_task: asyncio.Task = asyncio.create_task(617 anext(map_output, None),618 )619 # consume passthrough stream620 async for chunk in for_passthrough:621 if not isinstance(chunk, dict):622 msg = "The input to RunnablePassthrough.assign() must be a dict."623 raise ValueError(msg) # noqa: TRY004624 625 # remove mapper keys from passthrough chunk, to be overwritten by map output626 filtered = AddableDict(627 {k: v for k, v in chunk.items() if k not in mapper_keys}628 )629 if filtered:630 yield filtered631 # yield map output632 yield await first_map_chunk_task633 async for chunk in map_output:634 yield chunk635 636 @override637 async def atransform(638 self,639 input: AsyncIterator[dict[str, Any]],640 config: RunnableConfig | None = None,641 **kwargs: Any,642 ) -> AsyncIterator[dict[str, Any]]:643 async for chunk in self._atransform_stream_with_config(644 input, self._atransform, config, **kwargs645 ):646 yield chunk647 648 @override649 def stream(650 self,651 input: dict[str, Any],652 config: RunnableConfig | None = None,653 **kwargs: Any,654 ) -> Iterator[dict[str, Any]]:655 return self.transform(iter([input]), config, **kwargs)656 657 @override658 async def astream(659 self,660 input: dict[str, Any],661 config: RunnableConfig | None = None,662 **kwargs: Any,663 ) -> AsyncIterator[dict[str, Any]]:664 async def input_aiter() -> AsyncIterator[dict[str, Any]]:665 yield input666 667 async for chunk in self.atransform(input_aiter(), config, **kwargs):668 yield chunk669 670 671class RunnablePick(RunnableSerializable[dict[str, Any], Any]):672 """`Runnable` that picks keys from `dict[str, Any]` inputs.673 674 `RunnablePick` class represents a `Runnable` that selectively picks keys from a675 dictionary input. It allows you to specify one or more keys to extract676 from the input dictionary.677 678 !!! note "Return Type Behavior"679 The return type depends on the `keys` parameter:680 681 - When `keys` is a `str`: Returns the single value associated with that key682 - When `keys` is a `list`: Returns a dictionary containing only the selected683 keys684 685 Example:686 ```python687 from langchain_core.runnables.passthrough import RunnablePick688 689 input_data = {690 "name": "John",691 "age": 30,692 "city": "New York",693 "country": "USA",694 }695 696 # Single key - returns the value directly697 runnable_single = RunnablePick(keys="name")698 result_single = runnable_single.invoke(input_data)699 print(result_single) # Output: "John"700 701 # Multiple keys - returns a dictionary702 runnable_multiple = RunnablePick(keys=["name", "age"])703 result_multiple = runnable_multiple.invoke(input_data)704 print(result_multiple) # Output: {'name': 'John', 'age': 30}705 ```706 """707 708 keys: str | list[str]709 710 def __init__(self, keys: str | list[str], **kwargs: Any) -> None:711 """Create a `RunnablePick`.712 713 Args:714 keys: A single key or a list of keys to pick from the input dictionary.715 """716 super().__init__(keys=keys, **kwargs)717 718 @classmethod719 @override720 def is_lc_serializable(cls) -> bool:721 """Return `True` as this class is serializable."""722 return True723 724 @classmethod725 @override726 def get_lc_namespace(cls) -> list[str]:727 """Get the namespace of the LangChain object.728 729 Returns:730 `["langchain", "schema", "runnable"]`731 """732 return ["langchain", "schema", "runnable"]733 734 @override735 def get_name(self, suffix: str | None = None, *, name: str | None = None) -> str:736 name = (737 name738 or self.name739 or "RunnablePick"740 f"<{','.join([self.keys] if isinstance(self.keys, str) else self.keys)}>"741 )742 return super().get_name(suffix, name=name)743 744 def _pick(self, value: dict[str, Any]) -> Any:745 if not isinstance(value, dict):746 msg = "The input to RunnablePassthrough.assign() must be a dict."747 raise ValueError(msg) # noqa: TRY004748 749 if isinstance(self.keys, str):750 return value.get(self.keys)751 picked = {k: value.get(k) for k in self.keys if k in value}752 if picked:753 return AddableDict(picked)754 return None755 756 @override757 def invoke(758 self,759 input: dict[str, Any],760 config: RunnableConfig | None = None,761 **kwargs: Any,762 ) -> Any:763 return self._call_with_config(self._pick, input, config, **kwargs)764 765 async def _ainvoke(766 self,767 value: dict[str, Any],768 ) -> Any:769 return self._pick(value)770 771 @override772 async def ainvoke(773 self,774 input: dict[str, Any],775 config: RunnableConfig | None = None,776 **kwargs: Any,777 ) -> Any:778 return await self._acall_with_config(self._ainvoke, input, config, **kwargs)779 780 def _transform(781 self,782 chunks: Iterator[dict[str, Any]],783 ) -> Iterator[Any]:784 for chunk in chunks:785 picked = self._pick(chunk)786 if picked is not None:787 yield picked788 789 @override790 def transform(791 self,792 input: Iterator[dict[str, Any]],793 config: RunnableConfig | None = None,794 **kwargs: Any,795 ) -> Iterator[Any]:796 yield from self._transform_stream_with_config(797 input, self._transform, config, **kwargs798 )799 800 async def _atransform(801 self,802 chunks: AsyncIterator[dict[str, Any]],803 ) -> AsyncIterator[Any]:804 async for chunk in chunks:805 picked = self._pick(chunk)806 if picked is not None:807 yield picked808 809 @override810 async def atransform(811 self,812 input: AsyncIterator[dict[str, Any]],813 config: RunnableConfig | None = None,814 **kwargs: Any,815 ) -> AsyncIterator[Any]:816 async for chunk in self._atransform_stream_with_config(817 input, self._atransform, config, **kwargs818 ):819 yield chunk820 821 @override822 def stream(823 self,824 input: dict[str, Any],825 config: RunnableConfig | None = None,826 **kwargs: Any,827 ) -> Iterator[Any]:828 return self.transform(iter([input]), config, **kwargs)829 830 @override831 async def astream(832 self,833 input: dict[str, Any],834 config: RunnableConfig | None = None,835 **kwargs: Any,836 ) -> AsyncIterator[Any]:837 async def input_aiter() -> AsyncIterator[dict[str, Any]]:838 yield input839 840 async for chunk in self.atransform(input_aiter(), config, **kwargs):841 yield chunk842 