import json import logging import uuid import httpx import asyncio 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 后,轮询流式数据;总共尝试3次,每次间隔15秒 """ stream_url = f"{MTHREADS_UPSTREAM_URL}/{event_id}" logger.debug(f"[{request_id}] Streaming from URL: {stream_url}") for attempt in range(1, 4): logger.info(f"[{request_id}] Attempt {attempt}/3 to stream from MThreads.") try: async with httpx.AsyncClient(timeout=150.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: error"): try: data_line = await line_iterator.__anext__() if 'data: "404: Session not found."' in data_line: logger.warning( f"[{request_id}] MThreads session not found on attempt {attempt}. " f"{'Will retry.' if attempt < 3 else 'This was the last attempt.'}" ) break except StopAsyncIteration: logger.warning(f"[{request_id}] Stream ended after 'event: error' on attempt {attempt}.") break 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}'") 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 return except httpx.RequestError as e: logger.error(f"[{request_id}] Network error on attempt {attempt}: {e}") if attempt == 3: async for item in generate_network_error_sse(e, request_id): yield item return except (UpstreamAPIError, StopAsyncIteration, json.JSONDecodeError) as e: logger.warning(f"[{request_id}] Handled error on attempt {attempt}: {type(e).__name__} - {e}") if isinstance(e, StopAsyncIteration): logger.warning(f"[{request_id}] MThreads stream was interrupted on attempt {attempt}. Retrying...") if attempt == 3: 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 unexpectedly on final attempt.") else: err = e async for item in generate_error_sse(err, request_id): yield item return if attempt < 3: logger.info(f"[{request_id}] Waiting 15 seconds before next attempt.") await asyncio.sleep(15) logger.error(f"[{request_id}] Failed to get response from MThreads after 3 attempts. Returning 'No content' message.") error_message = "No content returned from model" error_payload = json.dumps({"error": {"message": error_message, "type": "proxy_error", "code": 504}}, ensure_ascii=False) yield f"event: error\n" yield f"data: {error_payload}\n\n" return StreamingResponse(stream_generator(), media_type="text/event-stream")