160 lines
7.2 KiB
Python
160 lines
7.2 KiB
Python
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")
|