253 lines
9.3 KiB
Python
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
|