Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
snowflake.py392 linesDownload Raw Back to chat_models
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 
codekingpro/portable-devtools · Team Ai