save code
This commit is contained in:
+43
-53
@@ -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 两种格式)。
|
||||
|
||||
格式1(data 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:
|
||||
|
||||
Reference in New Issue
Block a user