codekingpro/portable-devtools
114k
1from typing import Any, Dict, List, Optional, Union, cast2 3from langchain_core.callbacks import CallbackManagerForLLMRun4from langchain_core.language_models.llms import LLM5from langchain_core.utils import convert_to_secret_str, get_from_dict_or_env6from pydantic import ConfigDict, SecretStr, model_validator7 8from langchain_community.utilities.arcee import ArceeWrapper, DALMFilter9 10 11class Arcee(LLM):12 """Arcee's Domain Adapted Language Models (DALMs).13 14 To use, set the ``ARCEE_API_KEY`` environment variable with your Arcee API key,15 or pass ``arcee_api_key`` as a named parameter.16 17 Example:18 .. code-block:: python19 20 from langchain_community.llms import Arcee21 22 arcee = Arcee(23 model="DALM-PubMed",24 arcee_api_key="ARCEE-API-KEY"25 )26 27 response = arcee("AI-driven music therapy")28 """29 30 _client: Optional[ArceeWrapper] = None #: :meta private:31 """Arcee _client."""32 33 arcee_api_key: Union[SecretStr, str, None] = None34 """Arcee API Key"""35 36 model: str37 """Arcee DALM name"""38 39 arcee_api_url: str = "https://api.arcee.ai"40 """Arcee API URL"""41 42 arcee_api_version: str = "v2"43 """Arcee API Version"""44 45 arcee_app_url: str = "https://app.arcee.ai"46 """Arcee App URL"""47 48 model_id: str = ""49 """Arcee Model ID"""50 51 model_kwargs: Optional[Dict[str, Any]] = None52 """Keyword arguments to pass to the model."""53 54 model_config = ConfigDict(55 extra="forbid",56 )57 58 @property59 def _llm_type(self) -> str:60 """Return type of llm."""61 return "arcee"62 63 def __init__(self, **data: Any) -> None:64 """Initializes private fields."""65 66 super().__init__(**data)67 api_key = cast(SecretStr, self.arcee_api_key)68 self._client = ArceeWrapper(69 arcee_api_key=api_key,70 arcee_api_url=self.arcee_api_url,71 arcee_api_version=self.arcee_api_version,72 model_kwargs=self.model_kwargs,73 model_name=self.model,74 )75 76 @model_validator(mode="before")77 @classmethod78 def validate_environments(cls, values: Dict) -> Any:79 """Validate Arcee environment variables."""80 81 # validate env vars82 values["arcee_api_key"] = convert_to_secret_str(83 get_from_dict_or_env(84 values,85 "arcee_api_key",86 "ARCEE_API_KEY",87 )88 )89 90 values["arcee_api_url"] = get_from_dict_or_env(91 values,92 "arcee_api_url",93 "ARCEE_API_URL",94 )95 96 values["arcee_app_url"] = get_from_dict_or_env(97 values,98 "arcee_app_url",99 "ARCEE_APP_URL",100 )101 102 values["arcee_api_version"] = get_from_dict_or_env(103 values,104 "arcee_api_version",105 "ARCEE_API_VERSION",106 )107 108 # validate model kwargs109 if values.get("model_kwargs"):110 kw = values["model_kwargs"]111 112 # validate size113 if kw.get("size") is not None:114 if not kw.get("size") >= 0:115 raise ValueError("`size` must be positive")116 117 # validate filters118 if kw.get("filters") is not None:119 if not isinstance(kw.get("filters"), List):120 raise ValueError("`filters` must be a list")121 for f in kw.get("filters"):122 DALMFilter(**f)123 return values124 125 def _call(126 self,127 prompt: str,128 stop: Optional[List[str]] = None,129 run_manager: Optional[CallbackManagerForLLMRun] = None,130 **kwargs: Any,131 ) -> str:132 """Generate text from Arcee DALM.133 134 Args:135 prompt: Prompt to generate text from.136 size: The max number of context results to retrieve.137 Defaults to 3. (Can be less if filters are provided).138 filters: Filters to apply to the context dataset.139 """140 141 try:142 if not self._client:143 raise ValueError("Client is not initialized.")144 return self._client.generate(prompt=prompt, **kwargs)145 except Exception as e:146 raise Exception(f"Failed to generate text: {e}") from e147 