codekingpro/portable-devtools
114k
1import re2from collections import defaultdict3from typing import Any, Dict, Iterator, List, Optional, Tuple, Union4 5from langchain_core._api.deprecation import deprecated6from langchain_core.callbacks import (7 CallbackManagerForLLMRun,8)9from langchain_core.language_models.chat_models import BaseChatModel10from langchain_core.messages import (11 AIMessage,12 AIMessageChunk,13 BaseMessage,14 ChatMessage,15 HumanMessage,16 SystemMessage,17)18from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult19from pydantic import ConfigDict20 21from langchain_community.chat_models.anthropic import (22 convert_messages_to_prompt_anthropic,23)24from langchain_community.chat_models.meta import convert_messages_to_prompt_llama25from langchain_community.llms.bedrock import BedrockBase26from langchain_community.utilities.anthropic import (27 get_num_tokens_anthropic,28 get_token_ids_anthropic,29)30 31 32def _convert_one_message_to_text_mistral(message: BaseMessage) -> str:33 if isinstance(message, ChatMessage):34 message_text = f"\n\n{message.role.capitalize()}: {message.content}"35 elif isinstance(message, HumanMessage):36 message_text = f"[INST] {message.content} [/INST]"37 elif isinstance(message, AIMessage):38 message_text = f"{message.content}"39 elif isinstance(message, SystemMessage):40 message_text = f"<<SYS>> {message.content} <</SYS>>"41 else:42 raise ValueError(f"Got unknown type {message}")43 return message_text44 45 46def convert_messages_to_prompt_mistral(messages: List[BaseMessage]) -> str:47 """Convert a list of messages to a prompt for mistral."""48 return "\n".join(49 [_convert_one_message_to_text_mistral(message) for message in messages]50 )51 52 53def _format_image(image_url: str) -> Dict:54 """55 Formats an image of format data:image/jpeg;base64,{b64_string}56 to a dict for anthropic api57 58 {59 "type": "base64",60 "media_type": "image/jpeg",61 "data": "/9j/4AAQSkZJRg...",62 }63 64 And throws an error if it's not a b64 image65 """66 regex = r"^data:(?P<media_type>image/.+);base64,(?P<data>.+)$"67 match = re.match(regex, image_url)68 if match is None:69 raise ValueError(70 "Anthropic only supports base64-encoded images currently."71 " Example: data:image/png;base64,'/9j/4AAQSk'..."72 )73 return {74 "type": "base64",75 "media_type": match.group("media_type"),76 "data": match.group("data"),77 }78 79 80def _format_anthropic_messages(81 messages: List[BaseMessage],82) -> Tuple[Optional[str], List[Dict]]:83 """Format messages for anthropic."""84 85 """86 [87 {88 "role": _message_type_lookups[m.type],89 "content": [_AnthropicMessageContent(text=m.content).dict()],90 }91 for m in messages92 ]93 """94 system: Optional[str] = None95 formatted_messages: List[Dict] = []96 for i, message in enumerate(messages):97 if message.type == "system":98 if i != 0:99 raise ValueError("System message must be at beginning of message list.")100 if not isinstance(message.content, str):101 raise ValueError(102 "System message must be a string, "103 f"instead was: {type(message.content)}"104 )105 system = message.content106 continue107 108 role = _message_type_lookups[message.type]109 content: Union[str, List[Dict]]110 111 if not isinstance(message.content, str):112 # parse as dict113 assert isinstance(message.content, list), (114 "Anthropic message content must be str or list of dicts"115 )116 117 # populate content118 content = []119 for item in message.content:120 if isinstance(item, str):121 content.append(122 {123 "type": "text",124 "text": item,125 }126 )127 elif isinstance(item, dict):128 if "type" not in item:129 raise ValueError("Dict content item must have a type key")130 if item["type"] == "image_url":131 # convert format132 source = _format_image(item["image_url"]["url"])133 content.append(134 {135 "type": "image",136 "source": source,137 }138 )139 else:140 content.append(item)141 else:142 raise ValueError(143 f"Content items must be str or dict, instead was: {type(item)}"144 )145 else:146 content = message.content147 148 formatted_messages.append(149 {150 "role": role,151 "content": content,152 }153 )154 return system, formatted_messages155 156 157class ChatPromptAdapter:158 """Adapter class to prepare the inputs from Langchain to prompt format159 that Chat model expects.160 """161 162 @classmethod163 def convert_messages_to_prompt(164 cls, provider: str, messages: List[BaseMessage]165 ) -> str:166 if provider == "anthropic":167 prompt = convert_messages_to_prompt_anthropic(messages=messages)168 elif provider == "meta":169 prompt = convert_messages_to_prompt_llama(messages=messages)170 elif provider == "mistral":171 prompt = convert_messages_to_prompt_mistral(messages=messages)172 elif provider == "amazon":173 prompt = convert_messages_to_prompt_anthropic(174 messages=messages,175 human_prompt="\n\nUser:",176 ai_prompt="\n\nBot:",177 )178 else:179 raise NotImplementedError(180 f"Provider {provider} model does not support chat."181 )182 return prompt183 184 @classmethod185 def format_messages(186 cls, provider: str, messages: List[BaseMessage]187 ) -> Tuple[Optional[str], List[Dict]]:188 if provider == "anthropic":189 return _format_anthropic_messages(messages)190 191 raise NotImplementedError(192 f"Provider {provider} not supported for format_messages"193 )194 195 196_message_type_lookups = {197 "human": "user",198 "ai": "assistant",199 "AIMessageChunk": "assistant",200 "HumanMessageChunk": "user",201 "function": "user",202}203 204 205@deprecated(206 since="0.0.34", removal="1.0", alternative_import="langchain_aws.ChatBedrock"207)208class BedrockChat(BaseChatModel, BedrockBase):209 """Chat model that uses the Bedrock API."""210 211 @property212 def _llm_type(self) -> str:213 """Return type of chat model."""214 return "amazon_bedrock_chat"215 216 @classmethod217 def is_lc_serializable(cls) -> bool:218 """Return whether this model can be serialized by Langchain."""219 return True220 221 @classmethod222 def get_lc_namespace(cls) -> List[str]:223 """Get the namespace of the langchain object."""224 return ["langchain", "chat_models", "bedrock"]225 226 @property227 def lc_attributes(self) -> Dict[str, Any]:228 attributes: Dict[str, Any] = {}229 230 if self.region_name:231 attributes["region_name"] = self.region_name232 233 return attributes234 235 model_config = ConfigDict(236 extra="forbid",237 )238 239 def _stream(240 self,241 messages: List[BaseMessage],242 stop: Optional[List[str]] = None,243 run_manager: Optional[CallbackManagerForLLMRun] = None,244 **kwargs: Any,245 ) -> Iterator[ChatGenerationChunk]:246 provider = self._get_provider()247 prompt, system, formatted_messages = None, None, None248 249 if provider == "anthropic":250 system, formatted_messages = ChatPromptAdapter.format_messages(251 provider, messages252 )253 else:254 prompt = ChatPromptAdapter.convert_messages_to_prompt(255 provider=provider, messages=messages256 )257 258 for chunk in self._prepare_input_and_invoke_stream(259 prompt=prompt,260 system=system,261 messages=formatted_messages,262 stop=stop,263 run_manager=run_manager,264 **kwargs,265 ):266 delta = chunk.text267 yield ChatGenerationChunk(message=AIMessageChunk(content=delta))268 269 def _generate(270 self,271 messages: List[BaseMessage],272 stop: Optional[List[str]] = None,273 run_manager: Optional[CallbackManagerForLLMRun] = None,274 **kwargs: Any,275 ) -> ChatResult:276 completion = ""277 llm_output: Dict[str, Any] = {"model_id": self.model_id}278 279 if self.streaming:280 for chunk in self._stream(messages, stop, run_manager, **kwargs):281 completion += chunk.text282 else:283 provider = self._get_provider()284 prompt, system, formatted_messages = None, None, None285 params: Dict[str, Any] = {**kwargs}286 287 if provider == "anthropic":288 system, formatted_messages = ChatPromptAdapter.format_messages(289 provider, messages290 )291 else:292 prompt = ChatPromptAdapter.convert_messages_to_prompt(293 provider=provider, messages=messages294 )295 296 if stop:297 params["stop_sequences"] = stop298 299 completion, usage_info = self._prepare_input_and_invoke(300 prompt=prompt,301 stop=stop,302 run_manager=run_manager,303 system=system,304 messages=formatted_messages,305 **params,306 )307 308 llm_output["usage"] = usage_info309 310 return ChatResult(311 generations=[ChatGeneration(message=AIMessage(content=completion))],312 llm_output=llm_output,313 )314 315 def _combine_llm_outputs(self, llm_outputs: List[Optional[dict]]) -> dict:316 final_usage: Dict[str, int] = defaultdict(int)317 final_output = {}318 for output in llm_outputs:319 output = output or {}320 usage = output.get("usage", {})321 for token_type, token_count in usage.items():322 final_usage[token_type] += token_count323 final_output.update(output)324 final_output["usage"] = final_usage325 return final_output326 327 def get_num_tokens(self, text: str) -> int:328 if self._model_is_anthropic:329 return get_num_tokens_anthropic(text)330 else:331 return super().get_num_tokens(text)332 333 def get_token_ids(self, text: str) -> List[int]:334 if self._model_is_anthropic:335 return get_token_ids_anthropic(text)336 else:337 return super().get_token_ids(text)338 