codekingpro/portable-devtools
114k
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 