feat(relay): 中继新增 /v1/messages(Anthropic 格式)通道与记账

- 显式 relay_messages 路由,支持 x-api-key 鉴权
- _bill_usage_anthropic 按 Anthropic 用量格式记账
This commit is contained in:
Pine
2026-09-13 11:14:20 +08:00
parent f88f4a4e13
commit 9f6c4fc21f
+105
View File
@@ -214,6 +214,111 @@ async def relay_chat_completions(request: Request, db: Database = Depends(get_db
)
@router.post("/messages")
async def relay_messages(request: Request, db: Database = Depends(get_db)):
"""Anthropic 格式 POST /v1/messages:转发到 compute-engine。
鉴权与 chat 同规:Authorization: Bearer 或 x-api-key 均可(_auth_headers 已兼容);
平台 JWT 登录态调用:按 Anthropic usage(input/output_tokens) 记账;
引擎 PAT 调用(第三方 Anthropic 客户端):不额外记账,仅引擎计量。
流式:透传 SSE 并解析 message_delta.usage 记账。
"""
try:
body: dict[str, Any] = await request.json()
except Exception as exc: # noqa: BLE001
raise HTTPException(status_code=400, detail="Invalid JSON body") from exc
messages_url = f"{config.COMPUTE_BASE_URL.rstrip('/')}/v1/messages"
stream = bool(body.get("stream", False))
model = str(body.get("model") or "")
billing_user_id = await _resolve_billing_user(request, db)
headers = {**(await _auth_headers(request)), "Content-Type": "application/json"}
logger.info("[relay] messages model=%s stream=%s bytes=%d bill=%s", model, stream, len(json.dumps(body)), bool(billing_user_id))
client = httpx.AsyncClient(timeout=None)
upstream_request = client.build_request("POST", messages_url, json=body, headers=headers)
if not stream:
try:
upstream = await client.send(upstream_request)
payload = upstream.json() if upstream.content else None
except Exception: # noqa: BLE001
await client.aclose()
raise HTTPException(status_code=502, detail="Upstream relay failed") from None
await client.aclose()
# Anthropic usageinput_tokens/output_tokens)→ 平台记账
if billing_user_id and isinstance(payload, dict) and payload.get("usage"):
await _bill_usage_anthropic(db, billing_user_id, model, payload.get("usage"))
return JSONResponse(status_code=upstream.status_code, content=payload)
try:
upstream = await client.send(upstream_request, stream=True)
except Exception: # noqa: BLE001
await client.aclose()
raise HTTPException(status_code=502, detail="Upstream relay failed") from None
resp_headers: dict[str, str] = {}
if "content-type" in upstream.headers:
resp_headers["content-type"] = upstream.headers["content-type"]
if "x-request-id" in upstream.headers:
resp_headers["x-request-id"] = upstream.headers["x-request-id"]
if not billing_user_id:
return StreamingResponse(
upstream.aiter_raw(),
status_code=upstream.status_code,
media_type="text/event-stream",
headers=resp_headers,
background=BackgroundTask(client.aclose),
)
# 流式:透传 SSE 同时解析 usagemessage_start.input_tokens + message_delta.output_tokens),结束后记账
async def _stream_with_billing():
usage: dict | None = None
try:
async for line in upstream.aiter_lines():
yield (line + "\n").encode("utf-8")
if line.startswith("data:"):
payload = line[5:].strip()
if payload:
try:
obj = json.loads(payload)
if isinstance(obj, dict):
if obj.get("type") == "message_start" and isinstance(obj.get("message"), dict):
usage = dict(obj["message"].get("usage") or {})
elif obj.get("type") == "message_delta" and isinstance(obj.get("usage"), dict):
usage = {**(usage or {}), **obj["usage"]}
except Exception: # noqa: BLE001
pass
finally:
try:
await client.aclose()
except Exception: # noqa: BLE001
pass
if usage:
await _bill_usage_anthropic(db, billing_user_id, model, usage)
return StreamingResponse(
_stream_with_billing(),
status_code=upstream.status_code,
media_type="text/event-stream",
headers=resp_headers,
)
async def _bill_usage_anthropic(db: Database, user_id: str, model: str, usage: dict) -> None:
"""Anthropic usageinput_tokens/output_tokens)→ 平台记账(best-effort)。"""
try:
from ...services.compute_pricing_service import deduct_usage_post
await deduct_usage_post(
db, user_id, model,
int(usage.get("input_tokens") or usage.get("prompt_tokens") or 0),
int(usage.get("output_tokens") or usage.get("completion_tokens") or 0),
)
except Exception as exc: # noqa: BLE001
logger.warning("[relay] messages billing failed user=%s model=%s err=%s", user_id, model, exc)
async def _resolve_billing_user(request: Request, db: Database) -> str:
"""解析平台 JWT 登录态 → 平台 user_id;非 JWT(引擎 PAT)返回空串不记账。"""
client_auth = (request.headers.get("Authorization", "") or "").strip()