codekingpro/portable-devtools
114k
1import json2from typing import (3 Any,4 Callable,5 Dict,6 Iterator,7 List,8 Literal,9 Optional,10 Sequence,11 Type,12 Union,13)14 15from langchain_core.callbacks.manager import CallbackManagerForLLMRun16from langchain_core.language_models import BaseChatModel17from langchain_core.messages import (18 AIMessage,19 AIMessageChunk,20 BaseMessage,21 ChatMessage,22 HumanMessage,23 SystemMessage,24)25from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult26from langchain_core.tools import BaseTool27from langchain_core.utils import (28 convert_to_secret_str,29 get_from_dict_or_env,30 get_pydantic_field_names,31)32from langchain_core.utils.function_calling import convert_to_openai_tool33from langchain_core.utils.utils import _build_model_kwargs34from pydantic import Field, SecretStr, model_validator35 36SUPPORTED_ROLES: List[str] = [37 "system",38 "user",39 "assistant",40]41 42 43class ChatSnowflakeCortexError(Exception):44 """Error with Snowpark client."""45 46 47def _convert_message_to_dict(message: BaseMessage) -> dict:48 """Convert a LangChain message to a dictionary.49 50 Args:51 message: The LangChain message.52 53 Returns:54 The dictionary.55 """56 message_dict: Dict[str, Any] = {57 "content": message.content,58 }59 60 # Populate role and additional message data61 if isinstance(message, ChatMessage) and message.role in SUPPORTED_ROLES:62 message_dict["role"] = message.role63 elif isinstance(message, SystemMessage):64 message_dict["role"] = "system"65 elif isinstance(message, HumanMessage):66 message_dict["role"] = "user"67 elif isinstance(message, AIMessage):68 message_dict["role"] = "assistant"69 else:70 raise TypeError(f"Got unknown type {message}")71 return message_dict72 73 74def _truncate_at_stop_tokens(75 text: str,76 stop: Optional[List[str]],77) -> str:78 """Truncates text at the earliest stop token found."""79 if stop is None:80 return text81 82 for stop_token in stop:83 stop_token_idx = text.find(stop_token)84 if stop_token_idx != -1:85 text = text[:stop_token_idx]86 return text87 88 89class ChatSnowflakeCortex(BaseChatModel):90 """Snowflake Cortex based Chat model91 92 To use the chat model, you must have the ``snowflake-snowpark-python`` Python93 package installed and either:94 95 1. environment variables set with your snowflake credentials or96 2. directly passed in as kwargs to the ChatSnowflakeCortex constructor.97 98 Example:99 .. code-block:: python100 101 from langchain_community.chat_models import ChatSnowflakeCortex102 chat = ChatSnowflakeCortex()103 """104 105 # test_tools: Dict[str, Any] = Field(default_factory=dict)106 test_tools: Dict[str, Union[Dict[str, Any], Type, Callable, BaseTool]] = Field(107 default_factory=dict108 )109 110 session: Any = None111 """Snowpark session object."""112 113 model: str = "mistral-large"114 """Snowflake cortex hosted LLM model name, defaulted to `mistral-large`.115 Refer to docs for more options. Also note, not all models support 116 agentic workflows."""117 118 cortex_function: str = "complete"119 """Cortex function to use, defaulted to `complete`.120 Refer to docs for more options."""121 122 temperature: float = 0123 """Model temperature. Value should be >= 0 and <= 1.0"""124 125 max_tokens: Optional[int] = None126 """The maximum number of output tokens in the response."""127 128 top_p: Optional[float] = 0129 """top_p adjusts the number of choices for each predicted tokens based on130 cumulative probabilities. Value should be ranging between 0.0 and 1.0. 131 """132 133 snowflake_username: Optional[str] = Field(default=None, alias="username")134 """Automatically inferred from env var `SNOWFLAKE_USERNAME` if not provided."""135 snowflake_password: Optional[SecretStr] = Field(default=None, alias="password")136 """Automatically inferred from env var `SNOWFLAKE_PASSWORD` if not provided."""137 snowflake_account: Optional[str] = Field(default=None, alias="account")138 """Automatically inferred from env var `SNOWFLAKE_ACCOUNT` if not provided."""139 snowflake_database: Optional[str] = Field(default=None, alias="database")140 """Automatically inferred from env var `SNOWFLAKE_DATABASE` if not provided."""141 snowflake_schema: Optional[str] = Field(default=None, alias="schema")142 """Automatically inferred from env var `SNOWFLAKE_SCHEMA` if not provided."""143 snowflake_warehouse: Optional[str] = Field(default=None, alias="warehouse")144 """Automatically inferred from env var `SNOWFLAKE_WAREHOUSE` if not provided."""145 snowflake_role: Optional[str] = Field(default=None, alias="role")146 """Automatically inferred from env var `SNOWFLAKE_ROLE` if not provided."""147 148 def bind_tools(149 self,150 tools: Sequence[Union[Dict[str, Any], Type, Callable, BaseTool]],151 *,152 tool_choice: Optional[153 Union[dict, str, Literal["auto", "any", "none"], bool]154 ] = "auto",155 **kwargs: Any,156 ) -> "ChatSnowflakeCortex":157 """Bind tool-like objects to this chat model, ensuring they conform to158 expected formats."""159 160 formatted_tools = [convert_to_openai_tool(tool) for tool in tools]161 # self.test_tools.update(formatted_tools)162 formatted_tools_dict = {163 tool["name"]: tool for tool in formatted_tools if "name" in tool164 }165 self.test_tools.update(formatted_tools_dict)166 167 return self168 169 @model_validator(mode="before")170 @classmethod171 def build_extra(cls, values: Dict[str, Any]) -> Any:172 """Build extra kwargs from additional params that were passed in."""173 all_required_field_names = get_pydantic_field_names(cls)174 values = _build_model_kwargs(values, all_required_field_names)175 return values176 177 @model_validator(mode="before")178 def validate_environment(cls, values: Dict) -> Dict:179 try:180 from snowflake.snowpark import Session181 except ImportError:182 raise ImportError(183 """`snowflake-snowpark-python` package not found, please install:184 `pip install snowflake-snowpark-python`185 """186 )187 188 values["snowflake_username"] = get_from_dict_or_env(189 values, "snowflake_username", "SNOWFLAKE_USERNAME"190 )191 values["snowflake_password"] = convert_to_secret_str(192 get_from_dict_or_env(values, "snowflake_password", "SNOWFLAKE_PASSWORD")193 )194 values["snowflake_account"] = get_from_dict_or_env(195 values, "snowflake_account", "SNOWFLAKE_ACCOUNT"196 )197 values["snowflake_database"] = get_from_dict_or_env(198 values, "snowflake_database", "SNOWFLAKE_DATABASE"199 )200 values["snowflake_schema"] = get_from_dict_or_env(201 values, "snowflake_schema", "SNOWFLAKE_SCHEMA"202 )203 values["snowflake_warehouse"] = get_from_dict_or_env(204 values, "snowflake_warehouse", "SNOWFLAKE_WAREHOUSE"205 )206 values["snowflake_role"] = get_from_dict_or_env(207 values, "snowflake_role", "SNOWFLAKE_ROLE"208 )209 210 connection_params = {211 "account": values["snowflake_account"],212 "user": values["snowflake_username"],213 "password": values["snowflake_password"].get_secret_value(),214 "database": values["snowflake_database"],215 "schema": values["snowflake_schema"],216 "warehouse": values["snowflake_warehouse"],217 "role": values["snowflake_role"],218 "client_session_keep_alive": "True",219 }220 221 try:222 values["session"] = Session.builder.configs(connection_params).create()223 except Exception as e:224 raise ChatSnowflakeCortexError(f"Failed to create session: {e}")225 226 return values227 228 def __del__(self) -> None:229 if getattr(self, "session", None) is not None:230 self.session.close()231 232 @property233 def _llm_type(self) -> str:234 """Get the type of language model used by this chat model."""235 return f"snowflake-cortex-{self.model}"236 237 def _generate(238 self,239 messages: List[BaseMessage],240 stop: Optional[List[str]] = None,241 run_manager: Optional[CallbackManagerForLLMRun] = None,242 **kwargs: Any,243 ) -> ChatResult:244 message_dicts = [_convert_message_to_dict(m) for m in messages]245 246 # Check for tool invocation in the messages and prepare for tool use247 tool_output = None248 for message in messages:249 if (250 isinstance(message.content, dict)251 and isinstance(message, SystemMessage)252 and "invoke_tool" in message.content253 ):254 tool_info = json.loads(message.content.get("invoke_tool"))255 tool_name = tool_info.get("tool_name")256 if tool_name in self.test_tools:257 tool_args = tool_info.get("args", {})258 tool_output = self.test_tools[tool_name](**tool_args)259 break260 261 # Prepare messages for SQL query262 if tool_output:263 message_dicts.append(264 {"tool_output": str(tool_output)}265 ) # Ensure tool_output is a string266 267 # JSON dump the message_dicts and options without additional escaping268 message_json = json.dumps(message_dicts)269 options = {270 "temperature": self.temperature,271 "top_p": self.top_p if self.top_p is not None else 1.0,272 "max_tokens": self.max_tokens if self.max_tokens is not None else 2048,273 }274 options_json = json.dumps(options) # JSON string of options275 276 # Form the SQL statement using JSON literals277 sql_stmt = f"""278 select snowflake.cortex.{self.cortex_function}(279 '{self.model}',280 parse_json($${message_json}$$),281 parse_json($${options_json}$$)282 ) as llm_response;283 """284 285 try:286 # Use the Snowflake Cortex Complete function287 self.session.sql(288 f"USE WAREHOUSE {self.session.get_current_warehouse()};"289 ).collect()290 l_rows = self.session.sql(sql_stmt).collect()291 except Exception as e:292 raise ChatSnowflakeCortexError(293 f"Error while making request to Snowflake Cortex: {e}"294 )295 296 response = json.loads(l_rows[0]["LLM_RESPONSE"])297 ai_message_content = response["choices"][0]["messages"]298 299 content = _truncate_at_stop_tokens(ai_message_content, stop)300 message = AIMessage(301 content=content,302 response_metadata=response["usage"],303 )304 generation = ChatGeneration(message=message)305 return ChatResult(generations=[generation])306 307 def _stream_content(308 self, content: str, stop: Optional[List[str]]309 ) -> Iterator[ChatGenerationChunk]:310 """311 Stream the output of the model in chunks to return ChatGenerationChunk.312 """313 chunk_size = 50 # Define a reasonable chunk size for streaming314 truncated_content = _truncate_at_stop_tokens(content, stop)315 316 for i in range(0, len(truncated_content), chunk_size):317 chunk_content = truncated_content[i : i + chunk_size]318 319 # Create and yield a ChatGenerationChunk with partial content320 yield ChatGenerationChunk(message=AIMessageChunk(content=chunk_content))321 322 def _stream(323 self,324 messages: List[BaseMessage],325 stop: Optional[List[str]] = None,326 run_manager: Optional[CallbackManagerForLLMRun] = None,327 **kwargs: Any,328 ) -> Iterator[ChatGenerationChunk]:329 """Stream the output of the model in chunks to return ChatGenerationChunk."""330 message_dicts = [_convert_message_to_dict(m) for m in messages]331 332 # Check for and potentially use a tool before streaming333 for message in messages:334 if (335 isinstance(message, str)336 and isinstance(message, SystemMessage)337 and "invoke_tool" in message.content338 ):339 tool_info = json.loads(message.content)340 tool_list = tool_info.get("invoke_tools", [])341 for tool in tool_list:342 tool_name = tool.get("tool_name")343 tool_args = tool.get("args", {})344 345 if tool_name in self.test_tools:346 tool_args = tool_info.get("args", {})347 tool_result = self.test_tools[tool_name](**tool_args)348 additional_context = {"tool_output": tool_result}349 message_dicts.append(350 additional_context351 ) # Append tool result to message dicts352 353 # JSON dump the message_dicts and options without additional escaping354 message_json = json.dumps(message_dicts)355 options = {356 "temperature": self.temperature,357 "top_p": self.top_p if self.top_p is not None else 1.0,358 "max_tokens": self.max_tokens if self.max_tokens is not None else 2048,359 # "stream": True,360 }361 options_json = json.dumps(options) # JSON string of options362 363 # Form the SQL statement using JSON literals364 sql_stmt = f"""365 select snowflake.cortex.{self.cortex_function}(366 '{self.model}',367 parse_json($${message_json}$$),368 parse_json($${options_json}$$)369 ) as llm_stream_response;370 """371 372 try:373 # Use the Snowflake Cortex Complete function374 self.session.sql(375 f"USE WAREHOUSE {self.session.get_current_warehouse()};"376 ).collect()377 result = self.session.sql(sql_stmt).collect()378 379 # Iterate over the generator to yield streaming responses380 for row in result:381 response = json.loads(row["LLM_STREAM_RESPONSE"])382 ai_message_content = response["choices"][0]["messages"]383 384 # Stream response content in chunks385 for chunk in self._stream_content(ai_message_content, stop):386 yield chunk387 388 except Exception as e:389 raise ChatSnowflakeCortexError(390 f"Error while making request to Snowflake Cortex stream: {e}"391 )392 