codekingpro/portable-devtools
114k
1from typing import (2 Any,3 AsyncIterator,4 Callable,5 Dict,6 Iterator,7 List,8 Optional,9 Type,10 Union,11)12 13from langchain_core._api.deprecation import deprecated14from langchain_core.callbacks import (15 AsyncCallbackManagerForLLMRun,16 CallbackManagerForLLMRun,17)18from langchain_core.language_models.chat_models import BaseChatModel19from langchain_core.language_models.llms import create_base_retry_decorator20from langchain_core.messages import (21 AIMessage,22 AIMessageChunk,23 BaseMessage,24 BaseMessageChunk,25 ChatMessage,26 ChatMessageChunk,27 FunctionMessage,28 FunctionMessageChunk,29 HumanMessage,30 HumanMessageChunk,31 SystemMessage,32 SystemMessageChunk,33)34from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult35from langchain_core.utils import convert_to_secret_str36from langchain_core.utils.env import get_from_dict_or_env37from pydantic import Field, SecretStr, model_validator38 39from langchain_community.adapters.openai import convert_message_to_dict40 41 42def _convert_delta_to_message_chunk(43 _dict: Any, default_class: Type[BaseMessageChunk]44) -> BaseMessageChunk:45 """Convert a delta response to a message chunk."""46 role = _dict.role47 content = _dict.content or ""48 additional_kwargs: Dict = {}49 50 if role == "user" or default_class == HumanMessageChunk:51 return HumanMessageChunk(content=content)52 elif role == "assistant" or default_class == AIMessageChunk:53 return AIMessageChunk(content=content, additional_kwargs=additional_kwargs)54 elif role == "system" or default_class == SystemMessageChunk:55 return SystemMessageChunk(content=content)56 elif role == "function" or default_class == FunctionMessageChunk:57 return FunctionMessageChunk(content=content, name=_dict.name)58 elif role or default_class == ChatMessageChunk:59 return ChatMessageChunk(content=content, role=role)60 else:61 return default_class(content=content) # type: ignore[call-arg]62 63 64def convert_dict_to_message(_dict: Any) -> BaseMessage:65 """Convert a dict response to a message."""66 role = _dict.role67 content = _dict.content or ""68 if role == "user":69 return HumanMessage(content=content)70 elif role == "assistant":71 content = _dict.content72 additional_kwargs: Dict = {}73 return AIMessage(content=content, additional_kwargs=additional_kwargs)74 elif role == "system":75 return SystemMessage(content=content)76 elif role == "function":77 return FunctionMessage(content=content, name=_dict.name)78 else:79 return ChatMessage(content=content, role=role)80 81 82@deprecated(83 since="0.0.26",84 removal="1.0",85 alternative_import="langchain_fireworks.ChatFireworks",86)87class ChatFireworks(BaseChatModel):88 """Fireworks Chat models."""89 90 model: str = "accounts/fireworks/models/llama-v2-7b-chat"91 model_kwargs: dict = Field(92 default_factory=lambda: {93 "temperature": 0.7,94 "max_tokens": 512,95 "top_p": 1,96 }.copy()97 )98 fireworks_api_key: Optional[SecretStr] = None99 max_retries: int = 20100 use_retry: bool = True101 102 @property103 def lc_secrets(self) -> Dict[str, str]:104 return {"fireworks_api_key": "FIREWORKS_API_KEY"}105 106 @classmethod107 def is_lc_serializable(cls) -> bool:108 return True109 110 @classmethod111 def get_lc_namespace(cls) -> List[str]:112 """Get the namespace of the langchain object."""113 return ["langchain", "chat_models", "fireworks"]114 115 @model_validator(mode="before")116 @classmethod117 def validate_environment(cls, values: Dict) -> Any:118 """Validate that api key in environment."""119 try:120 import fireworks.client121 except ImportError as e:122 raise ImportError(123 "Could not import fireworks-ai python package. "124 "Please install it with `pip install fireworks-ai`."125 ) from e126 fireworks_api_key = convert_to_secret_str(127 get_from_dict_or_env(values, "fireworks_api_key", "FIREWORKS_API_KEY")128 )129 fireworks.client.api_key = fireworks_api_key.get_secret_value()130 return values131 132 @property133 def _llm_type(self) -> str:134 """Return type of llm."""135 return "fireworks-chat"136 137 def _generate(138 self,139 messages: List[BaseMessage],140 stop: Optional[List[str]] = None,141 run_manager: Optional[CallbackManagerForLLMRun] = None,142 **kwargs: Any,143 ) -> ChatResult:144 message_dicts = self._create_message_dicts(messages)145 146 params = {147 "model": self.model,148 "messages": message_dicts,149 **self.model_kwargs,150 **kwargs,151 }152 response = completion_with_retry(153 self,154 self.use_retry,155 run_manager=run_manager,156 stop=stop,157 **params,158 )159 return self._create_chat_result(response)160 161 async def _agenerate(162 self,163 messages: List[BaseMessage],164 stop: Optional[List[str]] = None,165 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,166 **kwargs: Any,167 ) -> ChatResult:168 message_dicts = self._create_message_dicts(messages)169 params = {170 "model": self.model,171 "messages": message_dicts,172 **self.model_kwargs,173 **kwargs,174 }175 response = await acompletion_with_retry(176 self, self.use_retry, run_manager=run_manager, stop=stop, **params177 )178 return self._create_chat_result(response)179 180 def _combine_llm_outputs(self, llm_outputs: List[Optional[dict]]) -> dict:181 if llm_outputs[0] is None:182 return {}183 return llm_outputs[0]184 185 def _create_chat_result(self, response: Any) -> ChatResult:186 generations = []187 for res in response.choices:188 message = convert_dict_to_message(res.message)189 gen = ChatGeneration(190 message=message,191 generation_info=dict(finish_reason=res.finish_reason),192 )193 generations.append(gen)194 llm_output = {"model": self.model}195 return ChatResult(generations=generations, llm_output=llm_output)196 197 def _create_message_dicts(198 self, messages: List[BaseMessage]199 ) -> List[Dict[str, Any]]:200 message_dicts = [convert_message_to_dict(m) for m in messages]201 return message_dicts202 203 def _stream(204 self,205 messages: List[BaseMessage],206 stop: Optional[List[str]] = None,207 run_manager: Optional[CallbackManagerForLLMRun] = None,208 **kwargs: Any,209 ) -> Iterator[ChatGenerationChunk]:210 message_dicts = self._create_message_dicts(messages)211 default_chunk_class: Type[BaseMessageChunk] = AIMessageChunk212 params = {213 "model": self.model,214 "messages": message_dicts,215 "stream": True,216 **self.model_kwargs,217 **kwargs,218 }219 for chunk in completion_with_retry(220 self, self.use_retry, run_manager=run_manager, stop=stop, **params221 ):222 choice = chunk.choices[0]223 chunk = _convert_delta_to_message_chunk(choice.delta, default_chunk_class)224 finish_reason = choice.finish_reason225 generation_info = (226 dict(finish_reason=finish_reason) if finish_reason is not None else None227 )228 default_chunk_class = chunk.__class__229 cg_chunk = ChatGenerationChunk(230 message=chunk, generation_info=generation_info231 )232 if run_manager:233 run_manager.on_llm_new_token(cg_chunk.text, chunk=cg_chunk)234 yield cg_chunk235 236 async def _astream(237 self,238 messages: List[BaseMessage],239 stop: Optional[List[str]] = None,240 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,241 **kwargs: Any,242 ) -> AsyncIterator[ChatGenerationChunk]:243 message_dicts = self._create_message_dicts(messages)244 default_chunk_class: Type[BaseMessageChunk] = AIMessageChunk245 params = {246 "model": self.model,247 "messages": message_dicts,248 "stream": True,249 **self.model_kwargs,250 **kwargs,251 }252 async for chunk in await acompletion_with_retry_streaming(253 self, self.use_retry, run_manager=run_manager, stop=stop, **params254 ):255 choice = chunk.choices[0]256 chunk = _convert_delta_to_message_chunk(choice.delta, default_chunk_class)257 finish_reason = choice.finish_reason258 generation_info = (259 dict(finish_reason=finish_reason) if finish_reason is not None else None260 )261 default_chunk_class = chunk.__class__262 cg_chunk = ChatGenerationChunk(263 message=chunk, generation_info=generation_info264 )265 if run_manager:266 await run_manager.on_llm_new_token(token=cg_chunk.text, chunk=cg_chunk)267 yield cg_chunk268 269 270def conditional_decorator(271 condition: bool, decorator: Callable[[Any], Any]272) -> Callable[[Any], Any]:273 """Define conditional decorator.274 275 Args:276 condition: The condition.277 decorator: The decorator.278 279 Returns:280 The decorated function.281 """282 283 def actual_decorator(func: Callable[[Any], Any]) -> Callable[[Any], Any]:284 if condition:285 return decorator(func)286 return func287 288 return actual_decorator289 290 291def completion_with_retry(292 llm: ChatFireworks,293 use_retry: bool,294 *,295 run_manager: Optional[CallbackManagerForLLMRun] = None,296 **kwargs: Any,297) -> Any:298 """Use tenacity to retry the completion call."""299 import fireworks.client300 301 retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)302 303 @conditional_decorator(use_retry, retry_decorator)304 def _completion_with_retry(**kwargs: Any) -> Any:305 """Use tenacity to retry the completion call."""306 return fireworks.client.ChatCompletion.create(307 **kwargs,308 )309 310 return _completion_with_retry(**kwargs)311 312 313async def acompletion_with_retry(314 llm: ChatFireworks,315 use_retry: bool,316 *,317 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,318 **kwargs: Any,319) -> Any:320 """Use tenacity to retry the async completion call."""321 import fireworks.client322 323 retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)324 325 @conditional_decorator(use_retry, retry_decorator)326 async def _completion_with_retry(**kwargs: Any) -> Any:327 return await fireworks.client.ChatCompletion.acreate(328 **kwargs,329 )330 331 return await _completion_with_retry(**kwargs)332 333 334async def acompletion_with_retry_streaming(335 llm: ChatFireworks,336 use_retry: bool,337 *,338 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,339 **kwargs: Any,340) -> Any:341 """Use tenacity to retry the completion call for streaming."""342 import fireworks.client343 344 retry_decorator = _create_retry_decorator(llm, run_manager=run_manager)345 346 @conditional_decorator(use_retry, retry_decorator)347 async def _completion_with_retry(**kwargs: Any) -> Any:348 return fireworks.client.ChatCompletion.acreate(349 **kwargs,350 )351 352 return await _completion_with_retry(**kwargs)353 354 355def _create_retry_decorator(356 llm: ChatFireworks,357 run_manager: Optional[358 Union[AsyncCallbackManagerForLLMRun, CallbackManagerForLLMRun]359 ] = None,360) -> Callable[[Any], Any]:361 """Define retry mechanism."""362 import fireworks.client363 364 errors = [365 fireworks.client.error.RateLimitError,366 fireworks.client.error.InternalServerError,367 fireworks.client.error.BadGatewayError,368 fireworks.client.error.ServiceUnavailableError,369 ]370 return create_base_retry_decorator(371 error_types=errors, max_retries=llm.max_retries, run_manager=run_manager372 )373 