Files
hair/gateway/forward.py
T
xslandClaude Opus 4.8 4c7681b338 perf(图片): 接口2/3/5 返回 JPG(体积~9×↓),接口1 标注图仍 PNG(透明)
- app.py: 接口2(预览+生发)/3(生发)/5(发际线叠图) 编码改 JPG(质量90,env JPG_QUALITY);
  接口1 annotated_image 含透明仍 PNG。_png_to_jpg_b64 把 ComfyUI 的 PNG 重编码为 JPG(无法解码则透传)
- gateway/forward.py: 落盘按内容嗅探扩展名(PNG头→.png 否则.jpg),原先硬编码 .png
- 测试/文档同步;实测接口5 一张 59KB(JPG) vs 548KB(PNG)。pytest 44 全绿

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-06-15 23:35:24 +08:00

299 lines
10 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, Optional
import httpx
from fastapi import Request
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_base64_value(value: str) -> Optional[bytes]:
"""解码 base64 值(支持 data URI 和原始 base64 两种格式)。
格式1data URI: data:image/png;base64,xxxx
格式2(原始base64): xxxx(无前缀,自动尝试解码)
"""
# 优先匹配 data URI
m = _DATA_URI_RE.match(value)
if m:
try:
return base64.b64decode(m.group(2))
except Exception:
logger.warning("data URI base64 解码失败")
return None
# 尝试当作原始 base64 解码(排除明显不是 base64 的短字符串)
stripped = value.strip()
if len(stripped) < 20:
return None # 太短,不可能是图片
try:
decoded = base64.b64decode(stripped)
if len(decoded) >= 50: # 最小合法图片大小
return decoded
except Exception:
pass
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 两种格式
- 解码 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_base64_value(value)
if img_bytes is not None:
# 按内容嗅探扩展名:PNG(接口1标注图,含透明) / JPEG(接口2/3/5 照片)
ext = "png" if img_bytes[:8] == b"\x89PNG\r\n\x1a\n" else "jpg"
filename = f"{uuid.uuid4().hex}.{ext}"
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 的值无法解码为图片,保留", 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. 读取客户端请求体(原始字节,不做解析) ---
body = await request.body()
# --- 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()
# 直接转发原始请求体(不解析、不重建,保证 multipart 原样透传)
headers = {"X-Internal-Token": token}
ct = request.headers.get("content-type", "")
if ct:
headers["Content-Type"] = ct
resp = await client.request(
method="POST",
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,
},
)