Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
fallbacks.py665 linesDownload Raw Back to runnables
1"""`Runnable` that can fallback to other `Runnable` objects if it fails."""2 3import asyncio4import inspect5import typing6from collections.abc import AsyncIterator, Iterator, Sequence7from functools import wraps8from typing import TYPE_CHECKING, Any, cast9 10from pydantic import BaseModel, ConfigDict11from typing_extensions import override12 13from langchain_core.callbacks.manager import AsyncCallbackManager, CallbackManager14from langchain_core.runnables.base import Runnable, RunnableSerializable15from langchain_core.runnables.config import (16    RunnableConfig,17    ensure_config,18    get_async_callback_manager_for_config,19    get_callback_manager_for_config,20    get_config_list,21    patch_config,22    set_config_context,23)24from langchain_core.runnables.utils import (25    ConfigurableFieldSpec,26    Input,27    Output,28    coro_with_context,29    get_unique_config_specs,30)31 32if TYPE_CHECKING:33    from langchain_core.callbacks.manager import AsyncCallbackManagerForChainRun34 35 36class RunnableWithFallbacks(RunnableSerializable[Input, Output]):37    """`Runnable` that can fallback to other `Runnable` objects if it fails.38 39    External APIs (e.g., APIs for a language model) may at times experience40    degraded performance or even downtime.41 42    In these cases, it can be useful to have a fallback `Runnable` that can be43    used in place of the original `Runnable` (e.g., fallback to another LLM provider).44 45    Fallbacks can be defined at the level of a single `Runnable`, or at the level46    of a chain of `Runnable`s. Fallbacks are tried in order until one succeeds or47    all fail.48 49    While you can instantiate a `RunnableWithFallbacks` directly, it is usually50    more convenient to use the `with_fallbacks` method on a `Runnable`.51 52    Example:53        ```python54        from langchain_core.chat_models.openai import ChatOpenAI55        from langchain_core.chat_models.anthropic import ChatAnthropic56 57        model = ChatAnthropic(model="claude-sonnet-4-6").with_fallbacks(58            [ChatOpenAI(model="gpt-5.4-mini")]59        )60        # Will usually use ChatAnthropic, but fallback to ChatOpenAI61        # if ChatAnthropic fails.62        model.invoke("hello")63 64        # And you can also use fallbacks at the level of a chain.65        # Here if both LLM providers fail, we'll fallback to a good hardcoded66        # response.67 68        from langchain_core.prompts import PromptTemplate69        from langchain_core.output_parser import StrOutputParser70        from langchain_core.runnables import RunnableLambda71 72 73        def when_all_is_lost(inputs):74            return (75                "Looks like our LLM providers are down. "76                "Here's a nice 🦜️ emoji for you instead."77            )78 79 80        chain_with_fallback = (81            PromptTemplate.from_template("Tell me a joke about {topic}")82            | model83            | StrOutputParser()84        ).with_fallbacks([RunnableLambda(when_all_is_lost)])85        ```86    """87 88    runnable: Runnable[Input, Output]89    """The `Runnable` to run first."""90    fallbacks: Sequence[Runnable[Input, Output]]91    """A sequence of fallbacks to try."""92    exceptions_to_handle: tuple[type[BaseException], ...] = (Exception,)93    """The exceptions on which fallbacks should be tried.94 95    Any exception that is not a subclass of these exceptions will be raised immediately.96    """97    exception_key: str | None = None98    """If `string` is specified then handled exceptions will be passed to fallbacks as99    part of the input under the specified key.100 101    If `None`, exceptions will not be passed to fallbacks.102 103    If used, the base `Runnable` and its fallbacks must accept a dictionary as input.104    """105 106    model_config = ConfigDict(107        arbitrary_types_allowed=True,108    )109 110    @property111    @override112    def InputType(self) -> type[Input]:113        return self.runnable.InputType114 115    @property116    @override117    def OutputType(self) -> type[Output]:118        return self.runnable.OutputType119 120    @override121    def get_input_schema(self, config: RunnableConfig | None = None) -> type[BaseModel]:122        return self.runnable.get_input_schema(config)123 124    @override125    def get_output_schema(126        self, config: RunnableConfig | None = None127    ) -> type[BaseModel]:128        return self.runnable.get_output_schema(config)129 130    @property131    @override132    def config_specs(self) -> list[ConfigurableFieldSpec]:133        return get_unique_config_specs(134            spec135            for step in [self.runnable, *self.fallbacks]136            for spec in step.config_specs137        )138 139    @classmethod140    @override141    def is_lc_serializable(cls) -> bool:142        """Return `True` as this class is serializable."""143        return True144 145    @classmethod146    @override147    def get_lc_namespace(cls) -> list[str]:148        """Get the namespace of the LangChain object.149 150        Returns:151            `["langchain", "schema", "runnable"]`152        """153        return ["langchain", "schema", "runnable"]154 155    @property156    def runnables(self) -> Iterator[Runnable[Input, Output]]:157        """Iterator over the `Runnable` and its fallbacks.158 159        Yields:160            The `Runnable` then its fallbacks.161        """162        yield self.runnable163        yield from self.fallbacks164 165    @override166    def invoke(167        self, input: Input, config: RunnableConfig | None = None, **kwargs: Any168    ) -> Output:169        if self.exception_key is not None and not isinstance(input, dict):170            msg = (171                "If 'exception_key' is specified then input must be a dictionary."172                f"However found a type of {type(input)} for input"173            )174            raise ValueError(msg)175        # setup callbacks176        config = ensure_config(config)177        callback_manager = get_callback_manager_for_config(config)178        # start the root run179        run_manager = callback_manager.on_chain_start(180            None,181            input,182            name=config.get("run_name") or self.get_name(),183            run_id=config.pop("run_id", None),184        )185        first_error = None186        last_error = None187        for runnable in self.runnables:188            try:189                if self.exception_key and last_error is not None:190                    input[self.exception_key] = last_error  # type: ignore[index]191                child_config = patch_config(config, callbacks=run_manager.get_child())192                with set_config_context(child_config) as context:193                    output = context.run(194                        runnable.invoke,195                        input,196                        config,197                        **kwargs,198                    )199            except self.exceptions_to_handle as e:200                if first_error is None:201                    first_error = e202                last_error = e203            except BaseException as e:204                run_manager.on_chain_error(e)205                raise206            else:207                run_manager.on_chain_end(output)208                return output209        if first_error is None:210            msg = "No error stored at end of fallbacks."211            raise ValueError(msg)212        run_manager.on_chain_error(first_error)213        raise first_error214 215    @override216    async def ainvoke(217        self,218        input: Input,219        config: RunnableConfig | None = None,220        **kwargs: Any | None,221    ) -> Output:222        if self.exception_key is not None and not isinstance(input, dict):223            msg = (224                "If 'exception_key' is specified then input must be a dictionary."225                f"However found a type of {type(input)} for input"226            )227            raise ValueError(msg)228        # setup callbacks229        config = ensure_config(config)230        callback_manager = get_async_callback_manager_for_config(config)231        # start the root run232        run_manager = await callback_manager.on_chain_start(233            None,234            input,235            name=config.get("run_name") or self.get_name(),236            run_id=config.pop("run_id", None),237        )238 239        first_error = None240        last_error = None241        for runnable in self.runnables:242            try:243                if self.exception_key and last_error is not None:244                    input[self.exception_key] = last_error  # type: ignore[index]245                child_config = patch_config(config, callbacks=run_manager.get_child())246                with set_config_context(child_config) as context:247                    coro = context.run(runnable.ainvoke, input, config, **kwargs)248                    output = await coro_with_context(coro, context)249            except self.exceptions_to_handle as e:250                if first_error is None:251                    first_error = e252                last_error = e253            except BaseException as e:254                await run_manager.on_chain_error(e)255                raise256            else:257                await run_manager.on_chain_end(output)258                return output259        if first_error is None:260            msg = "No error stored at end of fallbacks."261            raise ValueError(msg)262        await run_manager.on_chain_error(first_error)263        raise first_error264 265    @override266    def batch(267        self,268        inputs: list[Input],269        config: RunnableConfig | list[RunnableConfig] | None = None,270        *,271        return_exceptions: bool = False,272        **kwargs: Any | None,273    ) -> list[Output]:274        if self.exception_key is not None and not all(275            isinstance(input_, dict) for input_ in inputs276        ):277            msg = (278                "If 'exception_key' is specified then inputs must be dictionaries."279                f"However found a type of {type(inputs[0])} for input"280            )281            raise ValueError(msg)282 283        if not inputs:284            return []285 286        # setup callbacks287        configs = get_config_list(config, len(inputs))288        callback_managers = [289            CallbackManager.configure(290                inheritable_callbacks=config.get("callbacks"),291                local_callbacks=None,292                verbose=False,293                inheritable_tags=config.get("tags"),294                local_tags=None,295                inheritable_metadata=config.get("metadata"),296                local_metadata=None,297            )298            for config in configs299        ]300        # start the root runs, one per input301        run_managers = [302            cm.on_chain_start(303                None,304                input_ if isinstance(input_, dict) else {"input": input_},305                name=config.get("run_name") or self.get_name(),306                run_id=config.pop("run_id", None),307            )308            for cm, input_, config in zip(309                callback_managers, inputs, configs, strict=False310            )311        ]312 313        to_return: dict[int, Any] = {}314        run_again = dict(enumerate(inputs))315        handled_exceptions: dict[int, BaseException] = {}316        first_to_raise = None317        for runnable in self.runnables:318            outputs = runnable.batch(319                [input_ for _, input_ in sorted(run_again.items())],320                [321                    # each step a child run of the corresponding root run322                    patch_config(configs[i], callbacks=run_managers[i].get_child())323                    for i in sorted(run_again)324                ],325                return_exceptions=True,326                **kwargs,327            )328            for (i, input_), output in zip(329                sorted(run_again.copy().items()), outputs, strict=False330            ):331                if isinstance(output, BaseException) and not isinstance(332                    output, self.exceptions_to_handle333                ):334                    if not return_exceptions:335                        first_to_raise = first_to_raise or output336                    else:337                        handled_exceptions[i] = output338                    run_again.pop(i)339                elif isinstance(output, self.exceptions_to_handle):340                    if self.exception_key:341                        input_[self.exception_key] = output  # type: ignore[index]342                    handled_exceptions[i] = output343                else:344                    run_managers[i].on_chain_end(output)345                    to_return[i] = output346                    run_again.pop(i)347                    handled_exceptions.pop(i, None)348            if first_to_raise:349                raise first_to_raise350            if not run_again:351                break352 353        sorted_handled_exceptions = sorted(handled_exceptions.items())354        for i, error in sorted_handled_exceptions:355            run_managers[i].on_chain_error(error)356        if not return_exceptions and sorted_handled_exceptions:357            raise sorted_handled_exceptions[0][1]358        to_return.update(handled_exceptions)359        return [output for _, output in sorted(to_return.items())]360 361    @override362    async def abatch(363        self,364        inputs: list[Input],365        config: RunnableConfig | list[RunnableConfig] | None = None,366        *,367        return_exceptions: bool = False,368        **kwargs: Any | None,369    ) -> list[Output]:370        if self.exception_key is not None and not all(371            isinstance(input_, dict) for input_ in inputs372        ):373            msg = (374                "If 'exception_key' is specified then inputs must be dictionaries."375                f"However found a type of {type(inputs[0])} for input"376            )377            raise ValueError(msg)378 379        if not inputs:380            return []381 382        # setup callbacks383        configs = get_config_list(config, len(inputs))384        callback_managers = [385            AsyncCallbackManager.configure(386                inheritable_callbacks=config.get("callbacks"),387                local_callbacks=None,388                verbose=False,389                inheritable_tags=config.get("tags"),390                local_tags=None,391                inheritable_metadata=config.get("metadata"),392                local_metadata=None,393            )394            for config in configs395        ]396        # start the root runs, one per input397        run_managers: list[AsyncCallbackManagerForChainRun] = await asyncio.gather(398            *(399                cm.on_chain_start(400                    None,401                    input_,402                    name=config.get("run_name") or self.get_name(),403                    run_id=config.pop("run_id", None),404                )405                for cm, input_, config in zip(406                    callback_managers, inputs, configs, strict=False407                )408            )409        )410 411        to_return: dict[int, Output | BaseException] = {}412        run_again = dict(enumerate(inputs))413        handled_exceptions: dict[int, BaseException] = {}414        first_to_raise = None415        for runnable in self.runnables:416            outputs = await runnable.abatch(417                [input_ for _, input_ in sorted(run_again.items())],418                [419                    # each step a child run of the corresponding root run420                    patch_config(configs[i], callbacks=run_managers[i].get_child())421                    for i in sorted(run_again)422                ],423                return_exceptions=True,424                **kwargs,425            )426 427            for (i, input_), output in zip(428                sorted(run_again.copy().items()), outputs, strict=False429            ):430                if isinstance(output, BaseException) and not isinstance(431                    output, self.exceptions_to_handle432                ):433                    if not return_exceptions:434                        first_to_raise = first_to_raise or output435                    else:436                        handled_exceptions[i] = output437                    run_again.pop(i)438                elif isinstance(output, self.exceptions_to_handle):439                    if self.exception_key:440                        input_[self.exception_key] = output  # type: ignore[index]441                    handled_exceptions[i] = output442                else:443                    to_return[i] = output444                    await run_managers[i].on_chain_end(output)445                    run_again.pop(i)446                    handled_exceptions.pop(i, None)447 448            if first_to_raise:449                raise first_to_raise450            if not run_again:451                break452 453        sorted_handled_exceptions = sorted(handled_exceptions.items())454        await asyncio.gather(455            *(456                run_managers[i].on_chain_error(error)457                for i, error in sorted_handled_exceptions458            )459        )460        if not return_exceptions and sorted_handled_exceptions:461            raise sorted_handled_exceptions[0][1]462        to_return.update(handled_exceptions)463        return [cast("Output", output) for _, output in sorted(to_return.items())]464 465    @override466    def stream(467        self,468        input: Input,469        config: RunnableConfig | None = None,470        **kwargs: Any | None,471    ) -> Iterator[Output]:472        if self.exception_key is not None and not isinstance(input, dict):473            msg = (474                "If 'exception_key' is specified then input must be a dictionary."475                f"However found a type of {type(input)} for input"476            )477            raise ValueError(msg)478        # setup callbacks479        config = ensure_config(config)480        callback_manager = get_callback_manager_for_config(config)481        # start the root run482        run_manager = callback_manager.on_chain_start(483            None,484            input,485            name=config.get("run_name") or self.get_name(),486            run_id=config.pop("run_id", None),487        )488        first_error = None489        last_error = None490        for runnable in self.runnables:491            try:492                if self.exception_key and last_error is not None:493                    input[self.exception_key] = last_error  # type: ignore[index]494                child_config = patch_config(config, callbacks=run_manager.get_child())495                with set_config_context(child_config) as context:496                    stream = context.run(497                        runnable.stream,498                        input,499                        **kwargs,500                    )501                    chunk: Output = context.run(next, stream)502            except self.exceptions_to_handle as e:503                first_error = e if first_error is None else first_error504                last_error = e505            except BaseException as e:506                run_manager.on_chain_error(e)507                raise508            else:509                first_error = None510                break511        if first_error:512            run_manager.on_chain_error(first_error)513            raise first_error514 515        yield chunk516        output: Output | None = chunk517        try:518            for chunk in stream:519                yield chunk520                try:521                    output = output + chunk  # type: ignore[operator]522                except TypeError:523                    output = None524        except BaseException as e:525            run_manager.on_chain_error(e)526            raise527        run_manager.on_chain_end(output)528 529    @override530    async def astream(531        self,532        input: Input,533        config: RunnableConfig | None = None,534        **kwargs: Any | None,535    ) -> AsyncIterator[Output]:536        if self.exception_key is not None and not isinstance(input, dict):537            msg = (538                "If 'exception_key' is specified then input must be a dictionary."539                f"However found a type of {type(input)} for input"540            )541            raise ValueError(msg)542        # setup callbacks543        config = ensure_config(config)544        callback_manager = get_async_callback_manager_for_config(config)545        # start the root run546        run_manager = await callback_manager.on_chain_start(547            None,548            input,549            name=config.get("run_name") or self.get_name(),550            run_id=config.pop("run_id", None),551        )552        first_error = None553        last_error = None554        for runnable in self.runnables:555            try:556                if self.exception_key and last_error is not None:557                    input[self.exception_key] = last_error  # type: ignore[index]558                child_config = patch_config(config, callbacks=run_manager.get_child())559                with set_config_context(child_config) as context:560                    stream = runnable.astream(561                        input,562                        child_config,563                        **kwargs,564                    )565                    chunk = await coro_with_context(anext(stream), context)566            except self.exceptions_to_handle as e:567                first_error = e if first_error is None else first_error568                last_error = e569            except BaseException as e:570                await run_manager.on_chain_error(e)571                raise572            else:573                first_error = None574                break575        if first_error:576            await run_manager.on_chain_error(first_error)577            raise first_error578 579        yield chunk580        output: Output | None = chunk581        try:582            async for chunk in stream:583                yield chunk584                try:585                    output = output + chunk  # type: ignore[operator]586                except TypeError:587                    output = None588        except BaseException as e:589            await run_manager.on_chain_error(e)590            raise591        await run_manager.on_chain_end(output)592 593    def __getattr__(self, name: str) -> Any:594        """Get an attribute from the wrapped `Runnable` and its fallbacks.595 596        Returns:597            If the attribute is anything other than a method that outputs a `Runnable`,598            returns `getattr(self.runnable, name)`. If the attribute is a method that599            does return a new `Runnable` (e.g. `model.bind_tools([...])` outputs a new600            `RunnableBinding`) then `self.runnable` and each of the runnables in601            `self.fallbacks` is replaced with `getattr(x, name)`.602 603        Example:604            ```python605            from langchain_openai import ChatOpenAI606            from langchain_anthropic import ChatAnthropic607 608            gpt_4o = ChatOpenAI(model="gpt-4o")609            claude_3_sonnet = ChatAnthropic(model="claude-sonnet-4-5-20250929")610            model = gpt_4o.with_fallbacks([claude_3_sonnet])611 612            model.model_name613            # -> "gpt-4o"614 615            # .bind_tools() is called on both ChatOpenAI and ChatAnthropic616            # Equivalent to:617            # gpt_4o.bind_tools([...]).with_fallbacks([claude_3_sonnet.bind_tools([...])])618            model.bind_tools([...])619            # -> RunnableWithFallbacks(620                runnable=RunnableBinding(bound=ChatOpenAI(...), kwargs={"tools": [...]}),621                fallbacks=[RunnableBinding(bound=ChatAnthropic(...), kwargs={"tools": [...]})],622            )623            ```624        """  # noqa: E501625        attr = getattr(self.runnable, name)626        if _returns_runnable(attr):627 628            @wraps(attr)629            def wrapped(*args: Any, **kwargs: Any) -> Any:630                new_runnable = attr(*args, **kwargs)631                new_fallbacks = []632                for fallback in self.fallbacks:633                    fallback_attr = getattr(fallback, name)634                    new_fallbacks.append(fallback_attr(*args, **kwargs))635 636                return self.__class__(637                    **{638                        **self.model_dump(),639                        "runnable": new_runnable,640                        "fallbacks": new_fallbacks,641                    }642                )643 644            return wrapped645 646        return attr647 648 649def _returns_runnable(attr: Any) -> bool:650    if not callable(attr):651        return False652    return_type = typing.get_type_hints(attr).get("return")653    return bool(return_type and _is_runnable_type(return_type))654 655 656def _is_runnable_type(type_: Any) -> bool:657    if inspect.isclass(type_):658        return issubclass(type_, Runnable)659    origin = getattr(type_, "__origin__", None)660    if inspect.isclass(origin):661        return issubclass(origin, Runnable)662    if origin is typing.Union:663        return all(_is_runnable_type(t) for t in type_.__args__)664    return False665 
codekingpro/portable-devtools · Team Ai