codekingpro/portable-devtools
114k
1"""Chain that runs an arbitrary python function."""2 3import functools4import logging5from collections.abc import Awaitable, Callable6from typing import Any7 8from langchain_core.callbacks import (9 AsyncCallbackManagerForChainRun,10 CallbackManagerForChainRun,11)12from pydantic import Field13from typing_extensions import override14 15from langchain_classic.chains.base import Chain16 17logger = logging.getLogger(__name__)18 19 20class TransformChain(Chain):21 """Chain that transforms the chain output.22 23 Example:24 ```python25 from langchain_classic.chains import TransformChain26 transform_chain = TransformChain(input_variables=["text"],27 output_variables["entities"], transform=func())28 29 ```30 """31 32 input_variables: list[str]33 """The keys expected by the transform's input dictionary."""34 output_variables: list[str]35 """The keys returned by the transform's output dictionary."""36 transform_cb: Callable[[dict[str, str]], dict[str, str]] = Field(alias="transform")37 """The transform function."""38 atransform_cb: Callable[[dict[str, Any]], Awaitable[dict[str, Any]]] | None = Field(39 None, alias="atransform"40 )41 """The async coroutine transform function."""42 43 @staticmethod44 @functools.lru_cache45 def _log_once(msg: str) -> None:46 """Log a message once."""47 logger.warning(msg)48 49 @property50 def input_keys(self) -> list[str]:51 """Expect input keys."""52 return self.input_variables53 54 @property55 def output_keys(self) -> list[str]:56 """Return output keys."""57 return self.output_variables58 59 @override60 def _call(61 self,62 inputs: dict[str, str],63 run_manager: CallbackManagerForChainRun | None = None,64 ) -> dict[str, str]:65 return self.transform_cb(inputs)66 67 @override68 async def _acall(69 self,70 inputs: dict[str, Any],71 run_manager: AsyncCallbackManagerForChainRun | None = None,72 ) -> dict[str, Any]:73 if self.atransform_cb is not None:74 return await self.atransform_cb(inputs)75 self._log_once(76 "TransformChain's atransform is not provided, falling"77 " back to synchronous transform",78 )79 return self.transform_cb(inputs)80 