308 lines
11 KiB
Python
308 lines
11 KiB
Python
"""请求转发 + 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,
|
||
},
|
||
)
|