codekingpro/portable-devtools
114k
1import base642import hashlib3import hmac4import json5import logging6from datetime import datetime7from time import mktime8from typing import Any, Dict, List, Literal, Optional9from urllib.parse import urlencode10from wsgiref.handlers import format_date_time11 12import numpy as np13import requests14from langchain_core.embeddings import Embeddings15from langchain_core.utils import (16 secret_from_env,17)18from numpy import ndarray19from pydantic import BaseModel, ConfigDict, Field, SecretStr20 21# SparkLLMTextEmbeddings is an embedding model provided by iFLYTEK Co., Ltd.. (https://iflytek.com/en/).22 23# Official Website: https://www.xfyun.cn/doc/spark/Embedding_api.html24# Developers need to create an application in the console first, use the appid, APIKey,25# and APISecret provided in the application for authentication,26# and generate an authentication URL for handshake.27# You can get one by registering at https://console.xfyun.cn/services/bm3.28# SparkLLMTextEmbeddings support 2K token window and preduces vectors with29# 2560 dimensions.30 31logger = logging.getLogger(__name__)32 33 34class Url:35 """URL class for parsing the URL."""36 37 def __init__(self, host: str, path: str, schema: str) -> None:38 self.host = host39 self.path = path40 self.schema = schema41 pass42 43 44class SparkLLMTextEmbeddings(BaseModel, Embeddings):45 """SparkLLM embedding model integration.46 47 Setup:48 To use, you should have the environment variable "SPARK_APP_ID","SPARK_API_KEY"49 and "SPARK_API_SECRET" set your APP_ID, API_KEY and API_SECRET or pass it50 as a name parameter to the constructor.51 52 .. code-block:: bash53 54 export SPARK_APP_ID="your-api-id"55 export SPARK_API_KEY="your-api-key"56 export SPARK_API_SECRET="your-api-secret"57 58 Key init args — completion params:59 api_key: Optional[str]60 Automatically inferred from env var `SPARK_API_KEY` if not provided.61 app_id: Optional[str]62 Automatically inferred from env var `SPARK_APP_ID` if not provided.63 api_secret: Optional[str]64 Automatically inferred from env var `SPARK_API_SECRET` if not provided.65 base_url: Optional[str]66 Base URL path for API requests.67 68 See full list of supported init args and their descriptions in the params section.69 70 Instantiate:71 72 .. code-block:: python73 74 from langchain_community.embeddings import SparkLLMTextEmbeddings75 76 embed = SparkLLMTextEmbeddings(77 api_key="...",78 app_id="...",79 api_secret="...",80 # other81 )82 83 Embed single text:84 .. code-block:: python85 86 input_text = "The meaning of life is 42"87 embed.embed_query(input_text)88 89 .. code-block:: python90 91 [-0.4912109375, 0.60595703125, 0.658203125, 0.3037109375, 0.6591796875, 0.60302734375, ...]92 93 Embed multiple text:94 .. code-block:: python95 96 input_texts = ["This is a test query1.", "This is a test query2."]97 embed.embed_documents(input_texts)98 99 .. code-block:: python100 101 [102 [-0.1962890625, 0.94677734375, 0.7998046875, -0.1971435546875, 0.445556640625, 0.54638671875, ...],103 [ -0.44970703125, 0.06585693359375, 0.7421875, -0.474609375, 0.62353515625, 1.0478515625, ...],104 ]105 """ # noqa: E501106 107 spark_app_id: SecretStr = Field(108 alias="app_id", default_factory=secret_from_env("SPARK_APP_ID")109 )110 """Automatically inferred from env var `SPARK_APP_ID` if not provided."""111 spark_api_key: Optional[SecretStr] = Field(112 alias="api_key", default_factory=secret_from_env("SPARK_API_KEY", default=None)113 )114 """Automatically inferred from env var `SPARK_API_KEY` if not provided."""115 spark_api_secret: Optional[SecretStr] = Field(116 alias="api_secret",117 default_factory=secret_from_env("SPARK_API_SECRET", default=None),118 )119 """Automatically inferred from env var `SPARK_API_SECRET` if not provided."""120 base_url: str = Field(default="https://emb-cn-huabei-1.xf-yun.com/")121 """Base URL path for API requests"""122 domain: Literal["para", "query"] = Field(default="para")123 """This parameter is used for which Embedding this time belongs to.124 If "para"(default), it belongs to document Embedding. 125 If "query", it belongs to query Embedding."""126 127 model_config = ConfigDict(128 populate_by_name=True,129 )130 131 def _embed(self, texts: List[str], host: str) -> Optional[List[List[float]]]:132 """Internal method to call Spark Embedding API and return embeddings.133 134 Args:135 texts: A list of texts to embed.136 host: Base URL path for API requests137 138 Returns:139 A list of list of floats representing the embeddings,140 or list with value None if an error occurs.141 """142 app_id = ""143 api_key = ""144 api_secret = ""145 if self.spark_app_id:146 app_id = self.spark_app_id.get_secret_value()147 if self.spark_api_key:148 api_key = self.spark_api_key.get_secret_value()149 if self.spark_api_secret:150 api_secret = self.spark_api_secret.get_secret_value()151 url = self._assemble_ws_auth_url(152 request_url=host,153 method="POST",154 api_key=api_key,155 api_secret=api_secret,156 )157 embed_result: list = []158 for text in texts:159 query_context = {"messages": [{"content": text, "role": "user"}]}160 content = self._get_body(app_id, query_context)161 response = requests.post(162 url, json=content, headers={"content-type": "application/json"}163 ).text164 res_arr = self._parser_message(response)165 if res_arr is not None:166 embed_result.append(res_arr.tolist())167 else:168 embed_result.append(None)169 return embed_result170 171 def embed_documents(self, texts: List[str]) -> Optional[List[List[float]]]: # type: ignore[override]172 """Public method to get embeddings for a list of documents.173 174 Args:175 texts: The list of texts to embed.176 177 Returns:178 A list of embeddings, one for each text, or None if an error occurs.179 """180 return self._embed(texts, self.base_url)181 182 def embed_query(self, text: str) -> Optional[List[float]]: # type: ignore[override]183 """Public method to get embedding for a single query text.184 185 Args:186 text: The text to embed.187 188 Returns:189 Embeddings for the text, or None if an error occurs.190 """191 result = self._embed([text], self.base_url)192 return result[0] if result is not None else None193 194 @staticmethod195 def _assemble_ws_auth_url(196 request_url: str, method: str = "GET", api_key: str = "", api_secret: str = ""197 ) -> str:198 u = SparkLLMTextEmbeddings._parse_url(request_url)199 host = u.host200 path = u.path201 now = datetime.now()202 date = format_date_time(mktime(now.timetuple()))203 signature_origin = "host: {}\ndate: {}\n{} {} HTTP/1.1".format(204 host, date, method, path205 )206 signature_sha = hmac.new(207 api_secret.encode("utf-8"),208 signature_origin.encode("utf-8"),209 digestmod=hashlib.sha256,210 ).digest()211 signature_sha_str = base64.b64encode(signature_sha).decode(encoding="utf-8")212 authorization_origin = (213 'api_key="%s", algorithm="%s", headers="%s", signature="%s"'214 % (api_key, "hmac-sha256", "host date request-line", signature_sha_str)215 )216 authorization = base64.b64encode(authorization_origin.encode("utf-8")).decode(217 encoding="utf-8"218 )219 values = {"host": host, "date": date, "authorization": authorization}220 221 return request_url + "?" + urlencode(values)222 223 @staticmethod224 def _parse_url(request_url: str) -> Url:225 stidx = request_url.index("://")226 host = request_url[stidx + 3 :]227 schema = request_url[: stidx + 3]228 edidx = host.index("/")229 if edidx <= 0:230 raise AssembleHeaderException("invalid request url:" + request_url)231 path = host[edidx:]232 host = host[:edidx]233 u = Url(host, path, schema)234 return u235 236 def _get_body(self, appid: str, text: dict) -> Dict[str, Any]:237 body = {238 "header": {"app_id": appid, "uid": "39769795890", "status": 3},239 "parameter": {240 "emb": {"domain": self.domain, "feature": {"encoding": "utf8"}}241 },242 "payload": {243 "messages": {244 "text": base64.b64encode(json.dumps(text).encode("utf-8")).decode()245 }246 },247 }248 return body249 250 @staticmethod251 def _parser_message(252 message: str,253 ) -> Optional[ndarray]:254 data = json.loads(message)255 code = data["header"]["code"]256 if code != 0:257 logger.warning(f"Request error: {code}, {data}")258 return None259 else:260 text_base = data["payload"]["feature"]["text"]261 text_data = base64.b64decode(text_base)262 dt = np.dtype(np.float32)263 dt = dt.newbyteorder("<")264 text = np.frombuffer(text_data, dtype=dt)265 if len(text) > 2560:266 array = text[:2560]267 else:268 array = text269 return array270 271 272class AssembleHeaderException(Exception):273 """Exception raised for errors in the header assembly."""274 275 def __init__(self, msg: str) -> None:276 self.message = msg277 