Files
nekochatapi/utils.py
2025-07-16 15:40:02 +08:00

253 lines
9.3 KiB
Python

import json
import logging
import re
from typing import Any, Dict, List
import httpx
from config import DEBUG
logger = logging.getLogger("proxy")
class UpstreamAPIError(Exception):
"""
自定义异常类,表示与上游 API 通信时发生的错误;继承自 Exception ,增加 status_code 和 details 属性,方便格式统一
"""
def __init__(self, status_code: int, message: str, details: Any = None):
self.status_code = status_code
self.message = message
self.details = details
super().__init__(f"[{self.status_code}] {self.message}: {self.details}")
async def handle_upstream_response(response: httpx.Response):
"""
辅助函数,检查上游 API 的 HTTP 响应;如果响应状态码表示失败(非 2xx),则会解析错误信息并抛出 UpstreamAPIError 异常
"""
if response.is_success:
return
status_code = response.status_code
try:
# 尝试将响应解析为 JSON 格式
details = response.json()
except json.JSONDecodeError:
# 解析失败,直接截取文本作为错误详情
details = response.text[:200]
# 预设常见 HTTP 错误状态码
error_messages = {
400: "Invalid request",
401: "Authentication failed",
403: "Forbidden",
404: "Upstream resource not found",
429: "Too many requests",
500: "Upstream server internal error",
502: "Upstream server gateway error",
503: "Upstream service unavailable",
504: "Upstream server gateway timeout",
}
# 不在预设则使用通用错误信息
message = error_messages.get(status_code, "Unknown upstream API error")
raise UpstreamAPIError(status_code=status_code, message=message, details=details)
async def generate_error_sse(error: UpstreamAPIError, request_id: str):
"""
捕获到 UpstreamAPIError 时,生成符合 SSE 错误响应
"""
error_content = {
"error": {
"message": f"Upstream API Error: {error.message}",
"code": error.status_code,
"details": error.details
}
}
# 使用 json.dumps 将错误详情字典转换为 JSON 字符串
formatted_data = json.dumps(error_content, ensure_ascii=False)
logger.debug(
f"[{request_id}] Yielding Error to Client:\n"
f"event: error\n"
f"data: {formatted_data}\n\n"
)
yield f"event: error\n"
yield f"data: {formatted_data}\n\n"
async def generate_network_error_sse(exc: httpx.RequestError, request_id: str):
"""
捕获到 httpx 网络请求错误时,生成符合 SSE 错误响应
"""
error_content = {
"error": {
"message": f"Network request error: {exc}",
"code": "network_error"
}
}
formatted_data = json.dumps(error_content, ensure_ascii=False)
logger.debug(
f"[{request_id}] Yielding Network Error to Client:\n"
f"event: error\n"
f"data: {formatted_data}\n\n"
)
yield f"event: error\n"
yield f"data: {formatted_data}\n\n"
def convert_prompt_to_messages(prompt: List[str]) -> List[Dict[str, str]]:
"""
将自定义的 prompt 列表格式转换为符合 OpenAI API 标准的 messages 格式
- prompt 列表的第一个元素被视为 system 角色的内容
- prompt 列表的其余所有元素被合并,并视为 user 角色的内容
- 如果 prompt 只有一个元素,则直接视为 user 消息
"""
if not prompt:
return []
if len(prompt) == 1:
return [{"role": "user", "content": prompt[0]}]
# 第一个元素作为 system message
system_message = {"role": "system", "content": prompt[0]}
# 后续所有元素合并成一个 user message
user_content = "\n".join(prompt[1:])
user_message = {"role": "user", "content": user_content}
return [system_message, user_message]
def sanitize_mthreads_response(raw_text: str) -> str:
"""
清理 MThreads Qwen API 返回的无效内容
- 拼接 <|FunctionCallBegin|>...<|FunctionCallEnd|> 中的内容
- 移除各种形式的 <|...|>
- 移除末尾可能残留的特殊字符 `]>`
- 处理 Unicode 转义字符
"""
final_text_parts = []
def _extract_from_call(match):
"""re.sub 的回调函数,用于处理函数调用标记"""
call_content = match.group(1)
if not call_content or call_content == "[]":
return ""
try:
call_json = json.loads(call_content)
if call_json.get("name") == "send_message":
args_str = call_json.get("arguments", "{}")
args_json = json.loads(args_str)
message_text = args_json.get("text", "")
if message_text:
final_text_parts.append(message_text)
except (json.JSONDecodeError, TypeError):
# 解析失败,直接忽略这个函数调用块
pass
return ""
func_call_pattern = re.compile(
r"<\|FunctionCallBegin\|>(.*?)<\|FunctionCallEnd\|>", re.DOTALL
)
main_text = func_call_pattern.sub(_extract_from_call, raw_text)
main_text = re.sub(r"<\|.*?\|>", "", main_text)
main_text = main_text.strip().rstrip(']>')
if main_text.strip():
final_text_parts.insert(0, main_text.strip())
processed_text = "\n".join(final_text_parts).strip()
# 处理双重转义字符;用 latin-1 编码后再用 unicode-escape 解码修复
return processed_text.encode("latin-1", "backslashreplace").decode("unicode-escape")
async def standard_sse_generator(
upstream_url: str,
payload: dict,
headers: dict,
service_name: str,
request_id: str,
):
"""
用于处理与 OpenAI-like API 流式通信的异步生成器;向上游服务发送请求,并以 Server-Sent Events (SSE) 的格式流式返回响应
参数:
- upstream_url:上游 API 的地址
- payload:发送给上游 API 的请求体 (JSON)
- headers:请求头,通常包含认证信息
- service_name:服务名称,日志追踪
- request_id:当前请求的唯一 ID,日志追踪
"""
try:
# 使用 httpx.AsyncClient 以支持异步和连接池
async with httpx.AsyncClient(timeout=httpx.Timeout(600.0)) as client:
# 使用 client.stream 发起流式请求
async with client.stream(
"POST", upstream_url, json=payload, headers=headers
) as response:
# 检查初始响应头,如果状态码非 2xx 则会抛出异常
await handle_upstream_response(response)
accumulated_content = []
# 异步迭代来自服务器的每一行数据
async for line in response.aiter_lines():
# SSE 事件通常以 "data:" 开头
if not line.startswith("data:"):
continue
# 提取 data: 后面的内容
data_str = line.split(":", 1)[1].strip()
# SSE 流以 [DONE] 标记结束
if data_str == "[DONE]":
break
try:
data_json = json.loads(data_str)
# 遵循 OpenAI 的格式,从 choices[0].delta.content 中提取文本块
delta = data_json.get("choices", [{}])[0].get("delta", {})
content_chunk = delta.get("content")
if content_chunk:
accumulated_content.append(content_chunk)
except (json.JSONDecodeError, IndexError):
# 如果某一行不是有效的 JSON 或结构不符合预期,则忽略
pass
# 将所有接收到的文本块拼接成最终的完整文本
final_text = "".join(accumulated_content)
if not final_text:
final_text = "[No content returned from model]"
logger.debug(f"[{request_id}] Final Combined Text: '{final_text}'")
# 将最终文本包装成我们统一的 SSE 'complete' 事件格式
formatted_data = json.dumps([final_text], ensure_ascii=False)
chunk_for_client_event = f"event: complete\n"
chunk_for_client_data = f"data: {formatted_data}\n\n"
logger.debug(
f"[{request_id}] Yielding to Client:\n{chunk_for_client_event}{chunk_for_client_data}"
)
# 产出最终的 SSE 事件
yield chunk_for_client_event
yield chunk_for_client_data
except httpx.RequestError as exc:
# 处理网络层面的错误
async for item in generate_network_error_sse(exc, request_id):
yield item
except UpstreamAPIError as exc:
# 处理应用层面的上游 API 错误
async for item in generate_error_sse(exc, request_id):
yield item
except Exception as exc:
# 处理其他所有未预料到的异常
error = UpstreamAPIError(
status_code=500, message="An unknown error occurred in the proxy server", details=str(exc)
)
async for item in generate_error_sse(error, request_id):
yield item