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