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