Team Ai
Datasetpublic

codekingpro/portable-devtools

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