Files
hair/gateway/forward.py
T
2026-06-14 15:04:34 +08:00

308 lines
11 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""请求转发 + base64→URL 改写。
- 用 httpx.AsyncClient 把客户端请求原样转发到选中的 worker
- 附加 X-Internal-Token 头
- 失败重试(换 worker
- 响应中含 *_base64 图片字段 → 解码落盘 → 改写为 *_url
"""
import base64
import logging
import re
import uuid
from pathlib import Path
from typing import Any, Dict, List, Optional
import httpx
from fastapi import Request, UploadFile
from fastapi.responses import JSONResponse
from gateway.pool import (
NoWorkerAvailable,
acquire_worker,
mark_worker_unhealthy,
release_worker,
)
logger = logging.getLogger("gateway.forward")
# ---------------------------------------------------------------------------
# 全局 httpx client(连接复用)
# ---------------------------------------------------------------------------
_client: Optional[httpx.AsyncClient] = None
def get_client() -> httpx.AsyncClient:
global _client
if _client is None:
_client = httpx.AsyncClient()
return _client
async def close_client():
global _client
if _client:
await _client.aclose()
_client = None
# ---------------------------------------------------------------------------
# base64 → URL 改写
# ---------------------------------------------------------------------------
# data URI 正则:data:image/png;base64,xxxx
_DATA_URI_RE = re.compile(r"^data:(image/\w+);base64,(.+)$", re.IGNORECASE)
def _decode_data_uri(value: str) -> Optional[bytes]:
"""解析 data URI,返回解码后的字节;不匹配则返回 None。"""
m = _DATA_URI_RE.match(value)
if not m:
return None
try:
return base64.b64decode(m.group(2))
except Exception:
logger.warning("base64 解码失败,保留原文")
return None
def rewrite_base64_to_url(
obj: Any,
public_base_url: str,
static_dir: str,
) -> Any:
"""递归遍历响应 JSON,将所有 *_base64 字段改写为 *_url。
- 识别 key 以 _base64 结尾、值为 data: URI 的字段
- 解码 base64 → 保存到 static_dir/{uuid}.png
- 删除 *_base64 字段,新增 *_url 字段指向公网 URL
"""
if isinstance(obj, dict):
new_dict: Dict[str, Any] = {}
for key, value in obj.items():
if key.endswith("_base64") and isinstance(value, str):
img_bytes = _decode_data_uri(value)
if img_bytes is not None:
# 生成文件名并落盘
filename = f"{uuid.uuid4().hex}.png"
filepath = Path(static_dir) / filename
filepath.write_bytes(img_bytes)
# 构造对外 URL
url_key = key[:-7] + "_url" # "xxx_base64" → "xxx_url"
new_dict[url_key] = f"{public_base_url}/static/annotations/{filename}"
logger.info("base64→URL: %s%s (%d bytes)", key, new_dict[url_key], len(img_bytes))
continue # 跳过原 _base64 key
else:
logger.warning("字段 %s 的值不是有效的 data URI,保留", key)
new_dict[key] = value
else:
new_dict[key] = rewrite_base64_to_url(value, public_base_url, static_dir)
return new_dict
elif isinstance(obj, list):
return [rewrite_base64_to_url(item, public_base_url, static_dir) for item in obj]
else:
return obj
# ---------------------------------------------------------------------------
# 请求转发
# ---------------------------------------------------------------------------
async def proxy_request(request: Request, path: str) -> JSONResponse:
"""代理一次请求到 worker,处理重试与 base64 改写。
流程:
1. acquire worker(排队等待空闲)
2. 重构 multipart/form 请求,加 X-Internal-Token
3. 转发到 worker
4. 成功 → release worker → 改写 base64 → 返回 JSONResponse
5. 失败(连接/超时/5xx/401 → mark unhealthy → retry
6. retry 耗尽 / acquire 失败 → 返回 1007
"""
from gateway.config import get_config
cfg = get_config()
dispatch_cfg = cfg["dispatch"]
token = cfg["shared_password"]
request_timeout = dispatch_cfg["request_timeout_seconds"]
max_retries = dispatch_cfg.get("max_retries", 1)
retry_on_failure = dispatch_cfg.get("retry_on_failure", True)
public_base_url = cfg["public_base_url"]
static_dir = cfg["static_dir"]
# --- 1. 读取客户端请求体 ---
# 获取 form 数据(multipart/form-data
try:
form = await request.form()
except Exception:
# 非 form 请求:尝试读 JSON
body = await request.body()
form = None
# --- 2. 获取 worker(最多重试 max_retries+1 次) ---
attempts = max_retries + 1
last_error_response = None
for attempt in range(attempts):
worker = None
try:
worker = await acquire_worker(cfg)
except NoWorkerAvailable:
logger.warning("无可用 worker,返回 1007")
return JSONResponse(
status_code=503,
content={
"code": 1007,
"message": "后端服务暂不可用,请稍后重试",
"request_id": f"gw-{uuid.uuid4().hex[:8]}",
"data": None,
},
)
try:
client = get_client()
# 构建转发请求
if form is not None:
# multipart/form-data:重构文件和表单字段
req_files: List = []
req_data: Dict[str, str] = {}
for field_name, field_value in form.items():
if isinstance(field_value, UploadFile):
content = await field_value.read()
req_files.append(
(field_name, (field_value.filename or "file", content, field_value.content_type or "application/octet-stream"))
)
else:
req_data[field_name] = str(field_value)
resp = await client.post(
f"{worker.url}{path}",
data=req_data or None,
files=req_files or None,
headers={"X-Internal-Token": token},
timeout=request_timeout,
)
else:
# 回退:直接转发 body + content-type
headers = {"X-Internal-Token": token}
ct = request.headers.get("content-type", "")
if ct:
headers["Content-Type"] = ct
resp = await client.request(
method=request.method,
url=f"{worker.url}{path}",
content=body,
headers=headers,
timeout=request_timeout,
)
# --- 判断响应 ---
if resp.status_code == 401:
# worker 鉴权失败 → 视为 worker 异常
logger.warning("Worker %s 返回 401(鉴权失败),标记不健康", worker.url)
await mark_worker_unhealthy(worker)
await release_worker(worker)
if attempt < attempts - 1:
continue # retry
last_error_response = JSONResponse(
status_code=503,
content={
"code": 1007,
"message": "后端服务暂不可用,请稍后重试",
"request_id": f"gw-{uuid.uuid4().hex[:8]}",
"data": None,
},
)
break
if resp.status_code >= 500:
# worker 内部错误 → 视为该 worker 异常
logger.warning("Worker %s 返回 %d,标记不健康", worker.url, resp.status_code)
await mark_worker_unhealthy(worker)
await release_worker(worker)
if attempt < attempts - 1:
continue
last_error_response = JSONResponse(
status_code=503,
content={
"code": 1007,
"message": "后端服务暂不可用,请稍后重试",
"request_id": f"gw-{uuid.uuid4().hex[:8]}",
"data": None,
},
)
break
# --- 成功:改写 base64 → URL ---
await release_worker(worker)
try:
worker_json = resp.json()
except Exception:
logger.warning("Worker %s 返回非 JSON 响应", worker.url)
return JSONResponse(
status_code=502,
content={
"code": 1007,
"message": "后端服务响应异常",
"request_id": f"gw-{uuid.uuid4().hex[:8]}",
"data": None,
},
)
# 递归改写 base64 图片字段
rewritten = rewrite_base64_to_url(worker_json, public_base_url, static_dir)
return JSONResponse(
status_code=resp.status_code,
content=rewritten,
)
except (httpx.ConnectError, httpx.ConnectTimeout, httpx.ReadTimeout, httpx.RemoteProtocolError) as exc:
logger.warning("Worker %s 连接失败: %s", worker.url, exc)
await mark_worker_unhealthy(worker)
await release_worker(worker)
if attempt < attempts - 1:
continue
last_error_response = JSONResponse(
status_code=503,
content={
"code": 1007,
"message": "后端服务暂不可用,请稍后重试",
"request_id": f"gw-{uuid.uuid4().hex[:8]}",
"data": None,
},
)
break
except Exception as exc:
logger.error("转发异常: %s", exc)
await release_worker(worker)
if attempt < attempts - 1:
continue
last_error_response = JSONResponse(
status_code=503,
content={
"code": 1007,
"message": "后端服务暂不可用,请稍后重试",
"request_id": f"gw-{uuid.uuid4().hex[:8]}",
"data": None,
},
)
break
# --- 所有尝试耗尽 ---
return last_error_response or JSONResponse(
status_code=503,
content={
"code": 1007,
"message": "后端服务暂不可用,请稍后重试",
"request_id": f"gw-{uuid.uuid4().hex[:8]}",
"data": None,
},
)