Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
sparkllm.py653 linesDownload Raw Back to chat_models
1import base642import hashlib3import hmac4import json5import logging6import queue7import threading8from datetime import datetime9from queue import Queue10from time import mktime11from typing import Any, Dict, Generator, Iterator, List, Mapping, Optional, Type, cast12from urllib.parse import urlencode, urlparse, urlunparse13from wsgiref.handlers import format_date_time14 15from langchain_core.callbacks import (16    CallbackManagerForLLMRun,17)18from langchain_core.language_models.chat_models import (19    BaseChatModel,20    generate_from_stream,21)22from langchain_core.messages import (23    AIMessage,24    AIMessageChunk,25    BaseMessage,26    BaseMessageChunk,27    ChatMessage,28    ChatMessageChunk,29    FunctionMessageChunk,30    HumanMessage,31    HumanMessageChunk,32    SystemMessage,33    ToolMessageChunk,34)35from langchain_core.output_parsers.openai_tools import (36    make_invalid_tool_call,37    parse_tool_call,38)39from langchain_core.outputs import (40    ChatGeneration,41    ChatGenerationChunk,42    ChatResult,43)44from langchain_core.utils import (45    get_from_dict_or_env,46    get_pydantic_field_names,47)48from langchain_core.utils.pydantic import get_fields49from pydantic import ConfigDict, Field, model_validator50 51logger = logging.getLogger(__name__)52 53SPARK_API_URL = "wss://spark-api.xf-yun.com/v3.5/chat"54SPARK_LLM_DOMAIN = "generalv3.5"55 56 57def convert_message_to_dict(message: BaseMessage) -> dict:58    message_dict: Dict[str, Any]59    if isinstance(message, ChatMessage):60        message_dict = {"role": "user", "content": message.content}61    elif isinstance(message, HumanMessage):62        message_dict = {"role": "user", "content": message.content}63    elif isinstance(message, AIMessage):64        message_dict = {"role": "assistant", "content": message.content}65        if "function_call" in message.additional_kwargs:66            message_dict["function_call"] = message.additional_kwargs["function_call"]67            # If function call only, content is None not empty string68            if message_dict["content"] == "":69                message_dict["content"] = None70        if "tool_calls" in message.additional_kwargs:71            message_dict["tool_calls"] = message.additional_kwargs["tool_calls"]72            # If tool calls only, content is None not empty string73            if message_dict["content"] == "":74                message_dict["content"] = None75    elif isinstance(message, SystemMessage):76        message_dict = {"role": "system", "content": message.content}77    else:78        raise ValueError(f"Got unknown type {message}")79 80    return message_dict81 82 83def convert_dict_to_message(_dict: Mapping[str, Any]) -> BaseMessage:84    msg_role = _dict["role"]85    msg_content = _dict["content"]86    if msg_role == "user":87        return HumanMessage(content=msg_content)88    elif msg_role == "assistant":89        invalid_tool_calls = []90        additional_kwargs: Dict = {}91        if function_call := _dict.get("function_call"):92            additional_kwargs["function_call"] = dict(function_call)93        tool_calls = []94        if raw_tool_calls := _dict.get("tool_calls"):95            additional_kwargs["tool_calls"] = raw_tool_calls96            for raw_tool_call in _dict["tool_calls"]:97                try:98                    tool_calls.append(parse_tool_call(raw_tool_call, return_id=True))99                except Exception as e:100                    invalid_tool_calls.append(101                        make_invalid_tool_call(raw_tool_call, str(e))102                    )103        else:104            additional_kwargs = {}105        content = msg_content or ""106        return AIMessage(107            content=content,108            additional_kwargs=additional_kwargs,109            tool_calls=tool_calls,110            invalid_tool_calls=invalid_tool_calls,111        )112    elif msg_role == "system":113        return SystemMessage(content=msg_content)114    else:115        return ChatMessage(content=msg_content, role=msg_role)116 117 118def _convert_delta_to_message_chunk(119    _dict: Mapping[str, Any], default_class: Type[BaseMessageChunk]120) -> BaseMessageChunk:121    msg_role = cast(str, _dict.get("role"))122    msg_content = cast(str, _dict.get("content") or "")123    additional_kwargs: Dict = {}124    if _dict.get("function_call"):125        function_call = dict(_dict["function_call"])126        if "name" in function_call and function_call["name"] is None:127            function_call["name"] = ""128        additional_kwargs["function_call"] = function_call129    if _dict.get("tool_calls"):130        additional_kwargs["tool_calls"] = _dict["tool_calls"]131    if msg_role == "user" or default_class == HumanMessageChunk:132        return HumanMessageChunk(content=msg_content)133    elif msg_role == "assistant" or default_class == AIMessageChunk:134        return AIMessageChunk(content=msg_content, additional_kwargs=additional_kwargs)135    elif msg_role == "function" or default_class == FunctionMessageChunk:136        return FunctionMessageChunk(content=msg_content, name=_dict["name"])137    elif msg_role == "tool" or default_class == ToolMessageChunk:138        return ToolMessageChunk(content=msg_content, tool_call_id=_dict["tool_call_id"])139    elif msg_role or default_class == ChatMessageChunk:140        return ChatMessageChunk(content=msg_content, role=msg_role)141    else:142        return default_class(content=msg_content)  # type: ignore[call-arg]143 144 145class ChatSparkLLM(BaseChatModel):146    """IFlyTek Spark chat model integration.147 148    Setup:149        To use, you should have the environment variable``IFLYTEK_SPARK_API_KEY``,150        ``IFLYTEK_SPARK_API_SECRET`` and ``IFLYTEK_SPARK_APP_ID``.151 152    Key init args — completion params:153        model: Optional[str]154            Name of IFLYTEK SPARK model to use.155        temperature: Optional[float]156            Sampling temperature.157        top_k: Optional[float]158            What search sampling control to use.159        streaming: Optional[bool]160             Whether to stream the results or not.161 162    Key init args — client params:163        api_key: Optional[str]164            IFLYTEK SPARK API KEY. If not passed in will be read from env var IFLYTEK_SPARK_API_KEY.165        api_secret: Optional[str]166            IFLYTEK SPARK API SECRET. If not passed in will be read from env var IFLYTEK_SPARK_API_SECRET.167        api_url: Optional[str]168            Base URL for API requests.169        timeout: Optional[int]170            Timeout for requests.171 172    See full list of supported init args and their descriptions in the params section.173 174    Instantiate:175        .. code-block:: python176 177            from langchain_community.chat_models import ChatSparkLLM178 179            chat = ChatSparkLLM(180                api_key="your-api-key",181                api_secret="your-api-secret",182                model='Spark4.0 Ultra',183                # temperature=...,184                # other params...185            )186 187    Invoke:188        .. code-block:: python189 190            messages = [191                ("system", "你是一名专业的翻译家,可以将用户的中文翻译为英文。"),192                ("human", "我喜欢编程。"),193            ]194            chat.invoke(messages)195 196        .. code-block:: python197 198            AIMessage(199                content='I like programming.',200                response_metadata={201                    'token_usage': {202                        'question_tokens': 3,203                        'prompt_tokens': 16,204                        'completion_tokens': 4,205                        'total_tokens': 20206                    }207                },208                id='run-af8b3531-7bf7-47f0-bfe8-9262cb2a9d47-0'209            )210 211    Stream:212        .. code-block:: python213 214            for chunk in chat.stream(messages):215                print(chunk)216 217        .. code-block:: python218 219            content='I' id='run-fdbb57c2-2d32-4516-b894-6c5a67605d83'220            content=' like programming' id='run-fdbb57c2-2d32-4516-b894-6c5a67605d83'221            content='.' id='run-fdbb57c2-2d32-4516-b894-6c5a67605d83'222 223        .. code-block:: python224 225            stream = chat.stream(messages)226            full = next(stream)227            for chunk in stream:228                full += chunk229            full230 231        .. code-block:: python232 233            AIMessageChunk(234                content='I like programming.',235                id='run-aca2fa82-c2e4-4835-b7e2-865ddd3c46cb'236            )237 238    Response metadata239        .. code-block:: python240 241            ai_msg = chat.invoke(messages)242            ai_msg.response_metadata243 244        .. code-block:: python245 246            {247                'token_usage': {248                    'question_tokens': 3,249                    'prompt_tokens': 16,250                    'completion_tokens': 4,251                    'total_tokens': 20252                }253            }254 255    """  # noqa: E501256 257    @classmethod258    def is_lc_serializable(cls) -> bool:259        """Return whether this model can be serialized by Langchain."""260        return False261 262    @property263    def lc_secrets(self) -> Dict[str, str]:264        return {265            "spark_app_id": "IFLYTEK_SPARK_APP_ID",266            "spark_api_key": "IFLYTEK_SPARK_API_KEY",267            "spark_api_secret": "IFLYTEK_SPARK_API_SECRET",268            "spark_api_url": "IFLYTEK_SPARK_API_URL",269            "spark_llm_domain": "IFLYTEK_SPARK_LLM_DOMAIN",270        }271 272    client: Any = None  #: :meta private:273    spark_app_id: Optional[str] = Field(default=None, alias="app_id")274    """Automatically inferred from env var `IFLYTEK_SPARK_APP_ID` 275        if not provided."""276    spark_api_key: Optional[str] = Field(default=None, alias="api_key")277    """Automatically inferred from env var `IFLYTEK_SPARK_API_KEY` 278        if not provided."""279    spark_api_secret: Optional[str] = Field(default=None, alias="api_secret")280    """Automatically inferred from env var `IFLYTEK_SPARK_API_SECRET` 281        if not provided."""282    spark_api_url: Optional[str] = Field(default=None, alias="api_url")283    """Base URL path for API requests, leave blank if not using a proxy or service 284        emulator."""285    spark_llm_domain: Optional[str] = Field(default=None, alias="model")286    """Model name to use."""287    spark_user_id: str = "lc_user"288    streaming: bool = False289    """Whether to stream the results or not."""290    request_timeout: int = Field(30, alias="timeout")291    """request timeout for chat http requests"""292    temperature: float = Field(default=0.5)293    """What sampling temperature to use."""294    top_k: int = 4295    """What search sampling control to use."""296    model_kwargs: Dict[str, Any] = Field(default_factory=dict)297    """Holds any model parameters valid for API call not explicitly specified."""298 299    model_config = ConfigDict(300        populate_by_name=True,301    )302 303    @model_validator(mode="before")304    @classmethod305    def validate_environment(cls, values: Dict) -> Any:306        values["spark_app_id"] = get_from_dict_or_env(307            values,308            ["spark_app_id", "app_id"],309            "IFLYTEK_SPARK_APP_ID",310        )311        values["spark_api_key"] = get_from_dict_or_env(312            values,313            ["spark_api_key", "api_key"],314            "IFLYTEK_SPARK_API_KEY",315        )316        values["spark_api_secret"] = get_from_dict_or_env(317            values,318            ["spark_api_secret", "api_secret"],319            "IFLYTEK_SPARK_API_SECRET",320        )321        values["spark_api_url"] = get_from_dict_or_env(322            values,323            "spark_api_url",324            "IFLYTEK_SPARK_API_URL",325            SPARK_API_URL,326        )327        values["spark_llm_domain"] = get_from_dict_or_env(328            values,329            "spark_llm_domain",330            "IFLYTEK_SPARK_LLM_DOMAIN",331            SPARK_LLM_DOMAIN,332        )333 334        # put extra params into model_kwargs335        default_values = {336            name: field.default337            for name, field in get_fields(cls).items()338            if field.default is not None339        }340        values["model_kwargs"]["temperature"] = default_values.get("temperature")341        values["model_kwargs"]["top_k"] = default_values.get("top_k")342 343        values["client"] = _SparkLLMClient(344            app_id=values["spark_app_id"],345            api_key=values["spark_api_key"],346            api_secret=values["spark_api_secret"],347            api_url=values["spark_api_url"],348            spark_domain=values["spark_llm_domain"],349            model_kwargs=values["model_kwargs"],350        )351        return values352 353    # When using Pydantic V2354    # The execution order of multiple @model_validator decorators is opposite to355    # their declaration order. https://github.com/pydantic/pydantic/discussions/7434356 357    @model_validator(mode="before")358    @classmethod359    def build_extra(cls, values: Dict[str, Any]) -> Any:360        """Build extra kwargs from additional params that were passed in."""361        all_required_field_names = get_pydantic_field_names(cls)362        extra = values.get("model_kwargs", {})363        for field_name in list(values):364            if field_name in extra:365                raise ValueError(f"Found {field_name} supplied twice.")366            if field_name not in all_required_field_names:367                logger.warning(368                    f"""WARNING! {field_name} is not default parameter.369                    {field_name} was transferred to model_kwargs.370                    Please confirm that {field_name} is what you intended."""371                )372                extra[field_name] = values.pop(field_name)373 374        invalid_model_kwargs = all_required_field_names.intersection(extra.keys())375        if invalid_model_kwargs:376            raise ValueError(377                f"Parameters {invalid_model_kwargs} should be specified explicitly. "378                f"Instead they were passed in as part of `model_kwargs` parameter."379            )380 381        values["model_kwargs"] = extra382 383        return values384 385    def _stream(386        self,387        messages: List[BaseMessage],388        stop: Optional[List[str]] = None,389        run_manager: Optional[CallbackManagerForLLMRun] = None,390        **kwargs: Any,391    ) -> Iterator[ChatGenerationChunk]:392        default_chunk_class = AIMessageChunk393 394        self.client.arun(395            [convert_message_to_dict(m) for m in messages],396            self.spark_user_id,397            self.model_kwargs,398            streaming=True,399        )400        for content in self.client.subscribe(timeout=self.request_timeout):401            if "data" not in content:402                continue403            delta = content["data"]404            chunk = _convert_delta_to_message_chunk(delta, default_chunk_class)405            cg_chunk = ChatGenerationChunk(message=chunk)406            if run_manager:407                run_manager.on_llm_new_token(str(chunk.content), chunk=cg_chunk)408            yield cg_chunk409 410    def _generate(411        self,412        messages: List[BaseMessage],413        stop: Optional[List[str]] = None,414        run_manager: Optional[CallbackManagerForLLMRun] = None,415        stream: Optional[bool] = None,416        **kwargs: Any,417    ) -> ChatResult:418        if stream or self.streaming:419            stream_iter = self._stream(420                messages=messages, stop=stop, run_manager=run_manager, **kwargs421            )422            return generate_from_stream(stream_iter)423 424        self.client.arun(425            [convert_message_to_dict(m) for m in messages],426            self.spark_user_id,427            self.model_kwargs,428            False,429        )430        completion = {}431        llm_output = {}432        for content in self.client.subscribe(timeout=self.request_timeout):433            if "usage" in content:434                llm_output["token_usage"] = content["usage"]435            if "data" not in content:436                continue437            completion = content["data"]438        message = convert_dict_to_message(completion)439        generations = [ChatGeneration(message=message)]440        return ChatResult(generations=generations, llm_output=llm_output)441 442    @property443    def _llm_type(self) -> str:444        return "spark-llm-chat"445 446 447class _SparkLLMClient:448    """449    Use websocket-client to call the SparkLLM interface provided by Xfyun,450    which is the iFlyTek's open platform for AI capabilities451    """452 453    def __init__(454        self,455        app_id: str,456        api_key: str,457        api_secret: str,458        api_url: Optional[str] = None,459        spark_domain: Optional[str] = None,460        model_kwargs: Optional[dict] = None,461    ):462        try:463            import websocket464 465            self.websocket_client = websocket466        except ImportError:467            raise ImportError(468                "Could not import websocket client python package. "469                "Please install it with `pip install websocket-client`."470            )471 472        self.api_url = SPARK_API_URL if not api_url else api_url473        self.app_id = app_id474        self.model_kwargs = model_kwargs475        self.spark_domain = spark_domain or SPARK_LLM_DOMAIN476        self.queue: Queue[Dict] = Queue()477        self.blocking_message = {"content": "", "role": "assistant"}478        self.api_key = api_key479        self.api_secret = api_secret480 481    @staticmethod482    def _create_url(api_url: str, api_key: str, api_secret: str) -> str:483        """484        Generate a request url with an api key and an api secret.485        """486        # generate timestamp by RFC1123487        date = format_date_time(mktime(datetime.now().timetuple()))488 489        # urlparse490        parsed_url = urlparse(api_url)491        host = parsed_url.netloc492        path = parsed_url.path493 494        signature_origin = f"host: {host}\ndate: {date}\nGET {path} HTTP/1.1"495 496        # encrypt using hmac-sha256497        signature_sha = hmac.new(498            api_secret.encode("utf-8"),499            signature_origin.encode("utf-8"),500            digestmod=hashlib.sha256,501        ).digest()502 503        signature_sha_base64 = base64.b64encode(signature_sha).decode(encoding="utf-8")504 505        authorization_origin = f'api_key="{api_key}", algorithm="hmac-sha256", \506        headers="host date request-line", signature="{signature_sha_base64}"'507        authorization = base64.b64encode(authorization_origin.encode("utf-8")).decode(508            encoding="utf-8"509        )510 511        # generate url512        params_dict = {"authorization": authorization, "date": date, "host": host}513        encoded_params = urlencode(params_dict)514        url = urlunparse(515            (516                parsed_url.scheme,517                parsed_url.netloc,518                parsed_url.path,519                parsed_url.params,520                encoded_params,521                parsed_url.fragment,522            )523        )524        return url525 526    def run(527        self,528        messages: List[Dict],529        user_id: str,530        model_kwargs: Optional[dict] = None,531        streaming: bool = False,532    ) -> None:533        self.websocket_client.enableTrace(False)534        ws = self.websocket_client.WebSocketApp(535            _SparkLLMClient._create_url(536                self.api_url,537                self.api_key,538                self.api_secret,539            ),540            on_message=self.on_message,541            on_error=self.on_error,542            on_close=self.on_close,543            on_open=self.on_open,544        )545        ws.messages = messages  # type: ignore[attr-defined]546        ws.user_id = user_id  # type: ignore[attr-defined]547        ws.model_kwargs = self.model_kwargs if model_kwargs is None else model_kwargs  # type: ignore[attr-defined]548        ws.streaming = streaming  # type: ignore[attr-defined]549        ws.run_forever()550 551    def arun(552        self,553        messages: List[Dict],554        user_id: str,555        model_kwargs: Optional[dict] = None,556        streaming: bool = False,557    ) -> threading.Thread:558        ws_thread = threading.Thread(559            target=self.run,560            args=(561                messages,562                user_id,563                model_kwargs,564                streaming,565            ),566        )567        ws_thread.start()568        return ws_thread569 570    def on_error(self, ws: Any, error: Optional[Any]) -> None:571        self.queue.put({"error": error})572        ws.close()573 574    def on_close(self, ws: Any, close_status_code: int, close_reason: str) -> None:575        logger.debug(576            {577                "log": {578                    "close_status_code": close_status_code,579                    "close_reason": close_reason,580                }581            }582        )583        self.queue.put({"done": True})584 585    def on_open(self, ws: Any) -> None:586        self.blocking_message = {"content": "", "role": "assistant"}587        data = json.dumps(588            self.gen_params(589                messages=ws.messages, user_id=ws.user_id, model_kwargs=ws.model_kwargs590            )591        )592        ws.send(data)593 594    def on_message(self, ws: Any, message: str) -> None:595        data = json.loads(message)596        code = data["header"]["code"]597        if code != 0:598            self.queue.put(599                {"error": f"Code: {code}, Error: {data['header']['message']}"}600            )601            ws.close()602        else:603            choices = data["payload"]["choices"]604            status = choices["status"]605            content = choices["text"][0]["content"]606            if ws.streaming:607                self.queue.put({"data": choices["text"][0]})608            else:609                self.blocking_message["content"] += content610            if status == 2:611                if not ws.streaming:612                    self.queue.put({"data": self.blocking_message})613                usage_data = (614                    data.get("payload", {}).get("usage", {}).get("text", {})615                    if data616                    else {}617                )618                self.queue.put({"usage": usage_data})619                ws.close()620 621    def gen_params(622        self, messages: list, user_id: str, model_kwargs: Optional[dict] = None623    ) -> dict:624        data: Dict = {625            "header": {"app_id": self.app_id, "uid": user_id},626            "parameter": {"chat": {"domain": self.spark_domain}},627            "payload": {"message": {"text": messages}},628        }629 630        if model_kwargs:631            data["parameter"]["chat"].update(model_kwargs)632        logger.debug(f"Spark Request Parameters: {data}")633        return data634 635    def subscribe(self, timeout: Optional[int] = 30) -> Generator[Dict, None, None]:636        while True:637            try:638                content = self.queue.get(timeout=timeout)639            except queue.Empty as _:640                raise TimeoutError(641                    f"SparkLLMClient wait LLM api response timeout {timeout} seconds"642                )643            if "error" in content:644                raise ConnectionError(content["error"])645            if "usage" in content:646                yield content647                continue648            if "done" in content:649                break650            if "data" not in content:651                break652            yield content653 
codekingpro/portable-devtools · Team Ai