Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
coze.py256 linesDownload Raw Back to chat_models
1import json2import logging3from typing import Any, Dict, Iterator, List, Mapping, Optional, Union4 5import requests6from langchain_core.callbacks import CallbackManagerForLLMRun7from langchain_core.language_models.chat_models import (8    BaseChatModel,9    generate_from_stream,10)11from langchain_core.messages import (12    AIMessage,13    AIMessageChunk,14    BaseMessage,15    BaseMessageChunk,16    ChatMessage,17    ChatMessageChunk,18    HumanMessage,19    HumanMessageChunk,20)21from langchain_core.outputs import ChatGeneration, ChatGenerationChunk, ChatResult22from langchain_core.utils import (23    convert_to_secret_str,24    get_from_dict_or_env,25)26from pydantic import ConfigDict, Field, SecretStr, model_validator27 28logger = logging.getLogger(__name__)29 30DEFAULT_API_BASE = "https://api.coze.com"31 32 33def _convert_message_to_dict(message: BaseMessage) -> dict:34    message_dict: Dict[str, Any]35    if isinstance(message, HumanMessage):36        message_dict = {37            "role": "user",38            "content": message.content,39            "content_type": "text",40        }41    else:42        message_dict = {43            "role": "assistant",44            "content": message.content,45            "content_type": "text",46        }47    return message_dict48 49 50def _convert_dict_to_message(_dict: Mapping[str, Any]) -> Union[BaseMessage, None]:51    msg_type = _dict["type"]52    if msg_type != "answer":53        return None54    role = _dict["role"]55    if role == "user":56        return HumanMessage(content=_dict["content"])57    elif role == "assistant":58        return AIMessage(content=_dict.get("content", "") or "")59    else:60        return ChatMessage(content=_dict["content"], role=role)61 62 63def _convert_delta_to_message_chunk(_dict: Mapping[str, Any]) -> BaseMessageChunk:64    role = _dict.get("role")65    content = _dict.get("content") or ""66 67    if role == "user":68        return HumanMessageChunk(content=content)69    elif role == "assistant":70        return AIMessageChunk(content=content)71    else:72        return ChatMessageChunk(content=content, role=role)  # type: ignore[arg-type]73 74 75class ChatCoze(BaseChatModel):76    """ChatCoze chat models API by coze.com77 78    For more information, see https://www.coze.com/open/docs/chat79    """80 81    @property82    def lc_secrets(self) -> Dict[str, str]:83        return {84            "coze_api_key": "COZE_API_KEY",85        }86 87    @property88    def lc_serializable(self) -> bool:89        return True90 91    coze_api_base: str = Field(default=DEFAULT_API_BASE)92    """Coze custom endpoints"""93    coze_api_key: Optional[SecretStr] = None94    """Coze API Key"""95    request_timeout: int = Field(default=60, alias="timeout")96    """request timeout for chat http requests"""97    bot_id: str = Field(default="")98    """The ID of the bot that the API interacts with."""99    conversation_id: str = Field(default="")100    """Indicate which conversation the dialog is taking place in. If there is no need to101    distinguish the context of the conversation(just a question and answer), skip this102    parameter. It will be generated by the system."""103    user: str = Field(default="")104    """The user who calls the API to chat with the bot."""105    streaming: bool = False106    """Whether to stream the response to the client. 107    false: if no value is specified or set to false, a non-streaming response is108    returned. "Non-streaming response" means that all responses will be returned at once109    after they are all ready, and the client does not need to concatenate the content.110    true: set to true, partial message deltas will be sent .111    "Streaming response" will provide real-time response of the model to the client, and112    the client needs to assemble the final reply based on the type of message. """113 114    model_config = ConfigDict(115        populate_by_name=True,116    )117 118    @model_validator(mode="before")119    @classmethod120    def validate_environment(cls, values: Dict) -> Any:121        values["coze_api_base"] = get_from_dict_or_env(122            values,123            "coze_api_base",124            "COZE_API_BASE",125            DEFAULT_API_BASE,126        )127        values["coze_api_key"] = convert_to_secret_str(128            get_from_dict_or_env(129                values,130                "coze_api_key",131                "COZE_API_KEY",132            )133        )134 135        return values136 137    @property138    def _default_params(self) -> Dict[str, Any]:139        """Get the default parameters for calling Coze API."""140        return {141            "bot_id": self.bot_id,142            "conversation_id": self.conversation_id,143            "user": self.user,144            "streaming": self.streaming,145        }146 147    def _generate(148        self,149        messages: List[BaseMessage],150        stop: Optional[List[str]] = None,151        run_manager: Optional[CallbackManagerForLLMRun] = None,152        **kwargs: Any,153    ) -> ChatResult:154        if self.streaming:155            stream_iter = self._stream(156                messages=messages, stop=stop, run_manager=run_manager, **kwargs157            )158            return generate_from_stream(stream_iter)159 160        r = self._chat(messages, **kwargs)161        res = r.json()162        if res["code"] != 0:163            raise ValueError(164                f"Error from Coze api response: {res['code']}: {res['msg']}, "165                f"logid: {r.headers.get('X-Tt-Logid')}"166            )167 168        return self._create_chat_result(res.get("messages") or [])169 170    def _stream(171        self,172        messages: List[BaseMessage],173        stop: Optional[List[str]] = None,174        run_manager: Optional[CallbackManagerForLLMRun] = None,175        **kwargs: Any,176    ) -> Iterator[ChatGenerationChunk]:177        res = self._chat(messages, **kwargs)178        for chunk in res.iter_lines():179            chunk = chunk.decode("utf-8").strip("\r\n")180            parts = chunk.split("data:", 1)181            chunk = parts[1] if len(parts) > 1 else None182            if chunk is None:183                continue184            response = json.loads(chunk)185            if response["event"] == "done":186                break187            elif (188                response["event"] != "message"189                or response["message"]["type"] != "answer"190            ):191                continue192            chunk = _convert_delta_to_message_chunk(response["message"])193            cg_chunk = ChatGenerationChunk(message=chunk)194            if run_manager:195                run_manager.on_llm_new_token(str(chunk.content), chunk=cg_chunk)196            yield cg_chunk197 198    def _chat(self, messages: List[BaseMessage], **kwargs: Any) -> requests.Response:199        parameters = {**self._default_params, **kwargs}200 201        query = ""202        chat_history = []203        for msg in messages:204            if isinstance(msg, HumanMessage):205                query = f"{msg.content}"  # overwrite, to get last user message as query206            chat_history.append(_convert_message_to_dict(msg))207 208        conversation_id = parameters.pop("conversation_id")209        bot_id = parameters.pop("bot_id")210        user = parameters.pop("user")211        streaming = parameters.pop("streaming")212 213        payload = {214            "conversation_id": conversation_id,215            "bot_id": bot_id,216            "user": user,217            "query": query,218            "stream": streaming,219        }220        if chat_history:221            payload["chat_history"] = chat_history222 223        url = self.coze_api_base + "/open_api/v2/chat"224        api_key = ""225        if self.coze_api_key:226            api_key = self.coze_api_key.get_secret_value()227 228        res = requests.post(229            url=url,230            timeout=self.request_timeout,231            headers={232                "Content-Type": "application/json",233                "Authorization": f"Bearer {api_key}",234            },235            json=payload,236            stream=streaming,237        )238        if res.status_code != 200:239            logid = res.headers.get("X-Tt-Logid")240            raise ValueError(f"Error from Coze api response: {res}, logid: {logid}")241        return res242 243    def _create_chat_result(self, messages: List[Mapping[str, Any]]) -> ChatResult:244        generations = []245        for c in messages:246            msg = _convert_dict_to_message(c)247            if msg:248                generations.append(ChatGeneration(message=msg))249 250        llm_output = {"token_usage": "", "model": ""}251        return ChatResult(generations=generations, llm_output=llm_output)252 253    @property254    def _llm_type(self) -> str:255        return "coze-chat"256 
codekingpro/portable-devtools · Team Ai