Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
passthrough.py842 linesDownload Raw Back to runnables
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 
codekingpro/portable-devtools · Team Ai