"""请求转发 + 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, }, )