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

123 lines
5.2 KiB
Python

import json
import logging
import uuid
import httpx
from fastapi import APIRouter
from fastapi.responses import StreamingResponse
from config import MTHREADS_UPSTREAM_URL
from models import ChatRequest
from utils import (
UpstreamAPIError,
sanitize_mthreads_response,
handle_upstream_response,
generate_error_sse,
generate_network_error_sse,
)
router = APIRouter()
logger = logging.getLogger("proxy")
@router.post("/mthreads/chat")
async def proxy_mthreads_chat(request: ChatRequest):
request_id = f"mthreads-{uuid.uuid4()}"
client_request_body = request.model_dump()
logger.debug(
f"[{request_id}] === MThreads New Request ===\n"
f"[{request_id}] Client Request Body:\n"
f"{json.dumps(client_request_body, indent=2, ensure_ascii=False)}"
)
prompt_str = "\n".join(request.prompt)
mthreads_payload = {
"data": [prompt_str, request.temperature, request.top_p, request.max_tokens]
}
logger.debug(
f"[{request_id}] Agent Transformed Payload (to be sent to upstream):\n"
f"{json.dumps(mthreads_payload, indent=2, ensure_ascii=False)}"
)
try:
# --- 第一步: POST 请求获取 event_id ---
async with httpx.AsyncClient(timeout=60.0) as client:
response = await client.post(
MTHREADS_UPSTREAM_URL,
json=mthreads_payload,
headers={"Content-Type": "application/json"},
)
await handle_upstream_response(response)
response_json = response.json()
logger.debug(
f"[{request_id}] Upstream Initial Response:\n"
f"{json.dumps(response_json, indent=2, ensure_ascii=False)}"
)
event_id = response_json.get("event_id")
if not event_id:
raise UpstreamAPIError(
status_code=502,
message="Failed to get event_id from MThreads response.",
details=response_json,
)
except httpx.RequestError as e:
return StreamingResponse(
generate_network_error_sse(e, request_id),
media_type="text/event-stream",
)
except (UpstreamAPIError, json.JSONDecodeError) as e:
if isinstance(e, json.JSONDecodeError):
e = UpstreamAPIError(status_code=502, message="MThreads returned a non-JSON response.")
return StreamingResponse(generate_error_sse(e, request_id), media_type="text/event-stream")
async def stream_generator():
"""
第一次请求获取 event_id 后,轮询流式数据
"""
# --- 第二步: GET 请求流式数据 ---
stream_url = f"{MTHREADS_UPSTREAM_URL}/{event_id}"
logger.debug(f"[{request_id}] Streaming from URL: {stream_url}")
try:
async with httpx.AsyncClient(timeout=300.0) as client:
async with client.stream("GET", stream_url) as response:
await handle_upstream_response(response)
line_iterator = response.aiter_lines()
async for line in line_iterator:
if line.startswith("event: complete"):
data_line = await line_iterator.__anext__()
logger.debug(f"[{request_id}] Upstream Raw Data Line: {data_line}")
raw_text = json.loads(data_line.split(":", 1)[1].strip())[0]
logger.debug(f"[{request_id}] Upstream Parsed Raw Text: '{raw_text}'")
# 清理 MThreads 返回的非标准文本
final_text = sanitize_mthreads_response(raw_text)
logger.debug(f"[{request_id}] Agent Cleaned Text: '{final_text}'")
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}"
)
yield chunk_for_client_event
yield chunk_for_client_data
# 收到 complete 事件后,认为流已结束
break
except httpx.RequestError as e:
async for item in generate_network_error_sse(e, request_id):
yield item
except (UpstreamAPIError, StopAsyncIteration, json.JSONDecodeError) as e:
if isinstance(e, json.JSONDecodeError):
err = UpstreamAPIError(status_code=502, message="Failed to parse MThreads stream response.")
elif isinstance(e, StopAsyncIteration):
err = UpstreamAPIError(status_code=502, message="MThreads stream was interrupted.")
else:
err = e
async for item in generate_error_sse(err, request_id):
yield item
return StreamingResponse(stream_generator(), media_type="text/event-stream")