save code

This commit is contained in:
Ubuntu
2026-06-14 18:23:08 +08:00
parent 78226ae0b9
commit eaafbc97ea
3 changed files with 324 additions and 54 deletions
+43 -53
View File
@@ -11,10 +11,10 @@ import logging
import re
import uuid
from pathlib import Path
from typing import Any, Dict, List, Optional
from typing import Any, Optional
import httpx
from fastapi import Request, UploadFile
from fastapi import Request
from fastapi.responses import JSONResponse
from gateway.pool import (
@@ -56,16 +56,32 @@ async def close_client():
_DATA_URI_RE = re.compile(r"^data:(image/\w+);base64,(.+)$", re.IGNORECASE)
def _decode_data_uri(value: str) -> Optional[bytes]:
"""析 data URI,返回解码后的字节;不匹配则返回 None。"""
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 not m:
return None
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:
return base64.b64decode(m.group(2))
decoded = base64.b64decode(stripped)
if len(decoded) >= 50: # 最小合法图片大小
return decoded
except Exception:
logger.warning("base64 解码失败,保留原文")
return None
pass
return None
def rewrite_base64_to_url(
@@ -75,7 +91,8 @@ def rewrite_base64_to_url(
) -> Any:
"""递归遍历响应 JSON,将所有 *_base64 字段改写为 *_url。
- 识别 key 以 _base64 结尾、值为 data: URI 的字段
- 识别 key 以 _base64 结尾的字段
- 支持 data: URI 和原始 base64 两种格式
- 解码 base64 → 保存到 static_dir/{uuid}.png
- 删除 *_base64 字段,新增 *_url 字段指向公网 URL
"""
@@ -83,7 +100,7 @@ def rewrite_base64_to_url(
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)
img_bytes = _decode_base64_value(value)
if img_bytes is not None:
# 生成文件名并落盘
filename = f"{uuid.uuid4().hex}.png"
@@ -97,7 +114,7 @@ def rewrite_base64_to_url(
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)
logger.warning("字段 %s 的值无法解码为图片,保留", key)
new_dict[key] = value
else:
new_dict[key] = rewrite_base64_to_url(value, public_base_url, static_dir)
@@ -133,14 +150,8 @@ async def proxy_request(request: Request, path: str) -> JSONResponse:
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
# --- 1. 读取客户端请求体(原始字节,不做解析) ---
body = await request.body()
# --- 2. 获取 worker(最多重试 max_retries+1 次) ---
attempts = max_retries + 1
@@ -165,40 +176,19 @@ async def proxy_request(request: Request, path: str) -> JSONResponse:
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)
# 直接转发原始请求体(不解析、不重建,保证 multipart 原样透传)
headers = {"X-Internal-Token": token}
ct = request.headers.get("content-type", "")
if ct:
headers["Content-Type"] = ct
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,
)
resp = await client.request(
method="POST",
url=f"{worker.url}{path}",
content=body,
headers=headers,
timeout=request_timeout,
)
# --- 判断响应 ---
if resp.status_code == 401: