mirror of
https://github.com/zhayujie/chatgpt-on-wechat.git
synced 2026-07-20 05:27:59 +08:00
refactor(openai): drop SDK dependency and switch to native HTTP client
This commit is contained in:
@@ -3,8 +3,15 @@
|
||||
import time
|
||||
import json
|
||||
|
||||
import openai
|
||||
from models.openai.openai_compat import error as openai_error, RateLimitError, Timeout, APIError, APIConnectionError
|
||||
from models.openai.openai_compat import (
|
||||
error as openai_error,
|
||||
RateLimitError,
|
||||
Timeout,
|
||||
APIError,
|
||||
APIConnectionError,
|
||||
wrap_http_error,
|
||||
)
|
||||
from models.openai.openai_http_client import OpenAIHTTPClient, OpenAIHTTPError
|
||||
import requests
|
||||
from common import const
|
||||
from models.bot import Bot
|
||||
@@ -23,18 +30,19 @@ from models.baidu.baidu_wenxin_session import BaiduWenxinSession
|
||||
class ChatGPTBot(Bot, OpenAIImage, OpenAICompatibleBot):
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
# set the default api_key / api_base based on bot_type
|
||||
# Resolve api key / base from config (no global SDK state anymore).
|
||||
if conf().get("bot_type") == "custom":
|
||||
openai.api_key = conf().get("custom_api_key", "")
|
||||
if conf().get("custom_api_base"):
|
||||
openai.api_base = conf().get("custom_api_base")
|
||||
self._api_key = conf().get("custom_api_key", "")
|
||||
self._api_base = conf().get("custom_api_base") or None
|
||||
else:
|
||||
openai.api_key = conf().get("open_ai_api_key")
|
||||
if conf().get("open_ai_api_base"):
|
||||
openai.api_base = conf().get("open_ai_api_base")
|
||||
proxy = conf().get("proxy")
|
||||
if proxy:
|
||||
openai.proxy = proxy
|
||||
self._api_key = conf().get("open_ai_api_key")
|
||||
self._api_base = conf().get("open_ai_api_base") or None
|
||||
self._proxy = conf().get("proxy") or None
|
||||
self._http_client = OpenAIHTTPClient(
|
||||
api_key=self._api_key,
|
||||
api_base=self._api_base,
|
||||
proxy=self._proxy,
|
||||
)
|
||||
if conf().get("rate_limit_chatgpt"):
|
||||
self.tb4chatgpt = TokenBucket(conf().get("rate_limit_chatgpt", 20))
|
||||
conf_model = conf().get("model") or "gpt-3.5-turbo"
|
||||
@@ -71,6 +79,10 @@ class ChatGPTBot(Bot, OpenAIImage, OpenAICompatibleBot):
|
||||
'default_frequency_penalty': conf().get("frequency_penalty", 0.0),
|
||||
'default_presence_penalty': conf().get("presence_penalty", 0.0),
|
||||
}
|
||||
|
||||
def _get_http_client(self) -> OpenAIHTTPClient:
|
||||
"""Override the default HTTP client to reuse our pre-configured one."""
|
||||
return self._http_client
|
||||
|
||||
def reply(self, query, context=None):
|
||||
# acquire reply content
|
||||
@@ -195,20 +207,16 @@ class ChatGPTBot(Bot, OpenAIImage, OpenAICompatibleBot):
|
||||
|
||||
logger.info(f"[CHATGPT] Calling vision API with model: {model}")
|
||||
|
||||
# Call OpenAI API
|
||||
kwargs = {
|
||||
"model": model,
|
||||
"messages": messages,
|
||||
"max_tokens": 1000
|
||||
}
|
||||
if api_key:
|
||||
kwargs["api_key"] = api_key
|
||||
if api_base:
|
||||
kwargs["api_base"] = api_base
|
||||
|
||||
response = openai.ChatCompletion.create(**kwargs)
|
||||
|
||||
content = response.choices[0]["message"]["content"]
|
||||
# Call OpenAI-compatible API via HTTP
|
||||
response = self._http_client.chat_completions(
|
||||
api_key=api_key or None,
|
||||
api_base=api_base or None,
|
||||
model=model,
|
||||
messages=messages,
|
||||
max_tokens=1000,
|
||||
)
|
||||
|
||||
content = response["choices"][0]["message"]["content"]
|
||||
logger.info(f"[CHATGPT] Vision API response: {content[:100]}...")
|
||||
|
||||
# Clean up temp file
|
||||
@@ -237,57 +245,100 @@ class ChatGPTBot(Bot, OpenAIImage, OpenAICompatibleBot):
|
||||
try:
|
||||
if conf().get("rate_limit_chatgpt") and not self.tb4chatgpt.get_token():
|
||||
raise RateLimitError("RateLimitError: rate limit exceeded")
|
||||
# if api_key == None, the default openai.api_key will be used
|
||||
# If api_key is None, the per-instance default key will be used.
|
||||
if args is None:
|
||||
args = self.args
|
||||
response = openai.ChatCompletion.create(api_key=api_key, messages=session.messages, **args)
|
||||
# logger.debug("[CHATGPT] response={}".format(response))
|
||||
logger.info("[ChatGPT] reply={}, total_tokens={}".format(response.choices[0]['message']['content'], response["usage"]["total_tokens"]))
|
||||
# Translate old SDK kwargs to HTTP client params:
|
||||
# - request_timeout / timeout -> per-call timeout
|
||||
call_args = dict(args)
|
||||
timeout = call_args.pop("request_timeout", None) or call_args.pop("timeout", None)
|
||||
response = self._http_client.chat_completions(
|
||||
api_key=api_key or None,
|
||||
timeout=timeout,
|
||||
messages=session.messages,
|
||||
**call_args,
|
||||
)
|
||||
logger.info("[ChatGPT] reply={}, total_tokens={}".format(
|
||||
response["choices"][0]["message"]["content"],
|
||||
response["usage"]["total_tokens"]
|
||||
))
|
||||
return {
|
||||
"total_tokens": response["usage"]["total_tokens"],
|
||||
"completion_tokens": response["usage"]["completion_tokens"],
|
||||
"content": response.choices[0]["message"]["content"],
|
||||
"content": response["choices"][0]["message"]["content"],
|
||||
}
|
||||
except OpenAIHTTPError as http_err:
|
||||
return self._handle_reply_error(
|
||||
wrap_http_error(http_err), session, api_key, args, retry_count
|
||||
)
|
||||
except Exception as e:
|
||||
need_retry = retry_count < 2
|
||||
result = {"completion_tokens": 0, "content": "我现在有点累了,等会再来吧"}
|
||||
if isinstance(e, RateLimitError):
|
||||
logger.warn("[CHATGPT] RateLimitError: {}".format(e))
|
||||
result["content"] = "提问太快啦,请休息一下再问我吧"
|
||||
if need_retry:
|
||||
time.sleep(20)
|
||||
elif isinstance(e, Timeout):
|
||||
logger.warn("[CHATGPT] Timeout: {}".format(e))
|
||||
result["content"] = "我没有收到你的消息"
|
||||
if need_retry:
|
||||
time.sleep(5)
|
||||
elif isinstance(e, APIError):
|
||||
logger.warn("[CHATGPT] Bad Gateway: {}".format(e))
|
||||
result["content"] = "请再问我一次"
|
||||
if need_retry:
|
||||
time.sleep(10)
|
||||
elif isinstance(e, APIConnectionError):
|
||||
logger.warn("[CHATGPT] APIConnectionError: {}".format(e))
|
||||
result["content"] = "我连接不到你的网络"
|
||||
if need_retry:
|
||||
time.sleep(5)
|
||||
else:
|
||||
logger.exception("[CHATGPT] Exception: {}".format(e))
|
||||
need_retry = False
|
||||
self.sessions.clear_session(session.session_id)
|
||||
return self._handle_reply_error(e, session, api_key, args, retry_count)
|
||||
|
||||
def _handle_reply_error(self, e, session, api_key, args, retry_count):
|
||||
"""Map exception to user-facing reply with retry/backoff (mirrors SDK behavior)."""
|
||||
need_retry = retry_count < 2
|
||||
result = {"completion_tokens": 0, "content": "我现在有点累了,等会再来吧"}
|
||||
if isinstance(e, RateLimitError):
|
||||
logger.warn("[CHATGPT] RateLimitError: {}".format(e))
|
||||
result["content"] = "提问太快啦,请休息一下再问我吧"
|
||||
if need_retry:
|
||||
logger.warn("[CHATGPT] 第{}次重试".format(retry_count + 1))
|
||||
return self.reply_text(session, api_key, args, retry_count + 1)
|
||||
else:
|
||||
return result
|
||||
time.sleep(20)
|
||||
elif isinstance(e, Timeout):
|
||||
logger.warn("[CHATGPT] Timeout: {}".format(e))
|
||||
result["content"] = "我没有收到你的消息"
|
||||
if need_retry:
|
||||
time.sleep(5)
|
||||
elif isinstance(e, APIConnectionError):
|
||||
logger.warn("[CHATGPT] APIConnectionError: {}".format(e))
|
||||
result["content"] = "我连接不到你的网络"
|
||||
if need_retry:
|
||||
time.sleep(5)
|
||||
elif isinstance(e, APIError):
|
||||
logger.warn("[CHATGPT] Bad Gateway: {}".format(e))
|
||||
result["content"] = "请再问我一次"
|
||||
if need_retry:
|
||||
time.sleep(10)
|
||||
else:
|
||||
logger.exception("[CHATGPT] Exception: {}".format(e))
|
||||
need_retry = False
|
||||
self.sessions.clear_session(session.session_id)
|
||||
|
||||
if need_retry:
|
||||
logger.warn("[CHATGPT] 第{}次重试".format(retry_count + 1))
|
||||
return self.reply_text(session, api_key, args, retry_count + 1)
|
||||
return result
|
||||
|
||||
class AzureChatGPTBot(ChatGPTBot):
|
||||
"""Azure OpenAI variant.
|
||||
|
||||
Azure's HTTP shape differs from public OpenAI:
|
||||
URL : {endpoint}/openai/deployments/{deployment}/chat/completions
|
||||
Auth : api-key header (not Bearer)
|
||||
Query : ?api-version={version}
|
||||
We model that with a dedicated HTTP client and override _get_http_client
|
||||
so the OpenAICompatibleBot streaming/tool path uses it transparently.
|
||||
"""
|
||||
|
||||
def __init__(self):
|
||||
super().__init__()
|
||||
openai.api_type = "azure"
|
||||
openai.api_version = conf().get("azure_api_version", "2023-06-01-preview")
|
||||
self.args["deployment_id"] = conf().get("azure_deployment_id")
|
||||
self._azure_api_version = conf().get("azure_api_version", "2023-06-01-preview")
|
||||
self._azure_deployment_id = conf().get("azure_deployment_id")
|
||||
# Drop legacy SDK kwarg; Azure deployment is encoded in the URL now.
|
||||
self.args.pop("deployment_id", None)
|
||||
|
||||
endpoint = (self._api_base or "").rstrip("/")
|
||||
deployment = self._azure_deployment_id or ""
|
||||
# Build a base that already includes /openai/deployments/{deployment}.
|
||||
# /chat/completions will be appended by the client.
|
||||
azure_base = (
|
||||
f"{endpoint}/openai/deployments/{deployment}" if endpoint and deployment else endpoint
|
||||
)
|
||||
self._http_client = _AzureChatHTTPClient(
|
||||
api_key=self._api_key,
|
||||
api_base=azure_base,
|
||||
api_version=self._azure_api_version,
|
||||
proxy=self._proxy,
|
||||
)
|
||||
|
||||
def create_img(self, query, retry_count=0, api_key=None):
|
||||
text_to_image_model = conf().get("text_to_image")
|
||||
@@ -357,3 +408,35 @@ class AzureChatGPTBot(ChatGPTBot):
|
||||
return False, "图片生成失败"
|
||||
else:
|
||||
return False, "图片生成失败,未配置text_to_image参数"
|
||||
|
||||
|
||||
class _AzureChatHTTPClient(OpenAIHTTPClient):
|
||||
"""Subclass that injects Azure's ``api-version`` query param and ``api-key``
|
||||
header on every chat-completion request, and accepts the deployment-scoped
|
||||
base URL set by :class:`AzureChatGPTBot`.
|
||||
"""
|
||||
|
||||
def __init__(self, api_key, api_base, api_version, proxy=None, timeout=None):
|
||||
super().__init__(
|
||||
api_key=api_key, api_base=api_base, proxy=proxy, timeout=timeout
|
||||
)
|
||||
self._api_version = api_version
|
||||
|
||||
def _build_headers(self, api_key, extra_headers):
|
||||
# Azure uses api-key header, not Bearer token.
|
||||
key = api_key if api_key is not None else self.api_key
|
||||
headers = {"Content-Type": "application/json"}
|
||||
if key:
|
||||
headers["api-key"] = key
|
||||
if self.extra_headers:
|
||||
headers.update(self.extra_headers)
|
||||
if extra_headers:
|
||||
headers.update(extra_headers)
|
||||
return headers
|
||||
|
||||
def chat_completions(self, **kwargs):
|
||||
# Always force api-version query param for Azure.
|
||||
eq = dict(kwargs.get("extra_query") or {})
|
||||
eq.setdefault("api-version", self._api_version)
|
||||
kwargs["extra_query"] = eq
|
||||
return super().chat_completions(**kwargs)
|
||||
|
||||
Reference in New Issue
Block a user