Files
hair/gateway/app.py
T
xslandClaude Opus 4.8 a177dc2583 去除上传图片的分辨率与文件大小限制
- app.py: 删除全部 MAX_FILE_BYTES(≤1MB→1006) 与 MIN_SHORT_SIDE/MIN_LONG_SIDE
  (分辨率→1002) 校验及对应常量; 同步清理 File 描述、图片要求说明、
  错误码表(移除1002/1006)与过时示例
- gateway/app.py: 删除注释掉的 1006 大小校验块与描述里的 ≤1MB
- run_worker.sh / hair-worker.service: 删除临时放开限制的环境变量
- tests: 移除已过时的 test_oversize_1006 / test_lowres_1002 及 oversize_file fixture

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-06-24 22:21:39 +08:00

330 lines
11 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.
"""外网网关 — FastAPI 应用。
薄反向代理层:对外保持 HTTPS 接口不变,对内转发到 worker 池。
不跑任何算法(无 torch/mediapipe/opencv 依赖)。
"""
import asyncio
import base64
import json
import logging
import time
from contextlib import asynccontextmanager
from io import BytesIO
from pathlib import Path
from typing import Optional
from fastapi import FastAPI, File, Form, Request, UploadFile
from fastapi.responses import JSONResponse
from fastapi.staticfiles import StaticFiles
from gateway.config import load_config
# ---------------------------------------------------------------------------
# 日志
# ---------------------------------------------------------------------------
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s [%(levelname)s] %(name)s: %(message)s",
datefmt="%Y-%m-%d %H:%M:%S",
)
logger = logging.getLogger("gateway")
# ---------------------------------------------------------------------------
# 应用生命周期
# ---------------------------------------------------------------------------
@asynccontextmanager
async def lifespan(app: FastAPI):
"""启动:加载配置、初始化健康池;关闭:清理资源。"""
# 启动
cfg = load_config()
logger.info("网关启动中... workers=%s", cfg["workers"])
# 初始化健康池(阶段二实现)
try:
from gateway.pool import init_pool, shutdown_pool as _pool_shutdown
_has_pool = True
except ImportError:
logger.warning("pool 模块未就绪,跳过健康池初始化")
_has_pool = False
if _has_pool:
await init_pool(cfg)
app.state._pool_shutdown = _pool_shutdown
else:
app.state._pool_shutdown = None
# 确保标注图目录存在
static_dir = Path(cfg["static_dir"])
static_dir.mkdir(parents=True, exist_ok=True)
logger.info("标注图目录: %s", static_dir)
# 启动定期清理任务(阶段四)
cleanup_shutdown = asyncio.Event()
cleanup_task = asyncio.create_task(
_cleanup_loop(static_dir, cfg, cleanup_shutdown)
)
app.state._cleanup_shutdown = cleanup_shutdown
app.state._cleanup_task = cleanup_task
yield
# 关闭
logger.info("网关关闭中...")
# 停止清理任务
if app.state._cleanup_shutdown:
app.state._cleanup_shutdown.set()
if app.state._cleanup_task:
app.state._cleanup_task.cancel()
try:
await app.state._cleanup_task
except asyncio.CancelledError:
pass
if app.state._pool_shutdown:
await app.state._pool_shutdown()
logger.info("网关已关闭")
# ---------------------------------------------------------------------------
# 创建应用
# ---------------------------------------------------------------------------
app = FastAPI(
title="旷视五接口 — 网关",
version="0.1.0",
description="外网网关:反向代理 5 个接口到高性能 worker 池。",
lifespan=lifespan,
)
# 静态文件托管(阶段四完善)
static_root = Path(__file__).resolve().parent.parent / "static"
static_root.mkdir(parents=True, exist_ok=True)
(static_root / "annotations").mkdir(parents=True, exist_ok=True)
app.mount("/static", StaticFiles(directory=str(static_root)), name="static")
# ---------------------------------------------------------------------------
# 健康检查(网关自身)
# ---------------------------------------------------------------------------
def _get_pool_status_safe():
"""安全获取池状态(pool 未就绪时返回占位值)。"""
try:
from gateway.pool import get_pool_status
return get_pool_status()
except ImportError:
return {"total": 0, "healthy": 0, "busy": 0}
async def _cleanup_loop(annotations_dir: Path, cfg: dict, shutdown: asyncio.Event):
"""定期清理 static/annotations/ 中过期的标注图文件。
配置项(可选,在 config.json 中设定):
- cleanup.interval_minutes: 清理间隔,默认 60
- cleanup.max_age_hours: 文件保留时长(小时),默认 24
"""
cleanup_cfg = cfg.get("cleanup", {})
interval_s = cleanup_cfg.get("interval_minutes", 60) * 60
max_age_s = cleanup_cfg.get("max_age_hours", 24) * 3600
logger.info(
"清理任务启动 | 间隔=%dmin | 保留=%dh | 目录=%s",
interval_s // 60, max_age_s // 3600, annotations_dir,
)
while not shutdown.is_set():
try:
await asyncio.wait_for(shutdown.wait(), timeout=interval_s)
break # shutdown
except asyncio.TimeoutError:
pass # 正常到时,执行清理
now = time.time()
deleted = 0
for f in annotations_dir.iterdir():
if f.name == ".gitkeep":
continue
if not f.is_file():
continue
try:
age_s = now - f.stat().st_mtime
if age_s > max_age_s:
f.unlink()
deleted += 1
logger.debug("清理过期文件: %s (age=%.1fh)", f.name, age_s / 3600)
except Exception:
logger.warning("清理文件失败: %s", f.name, exc_info=True)
if deleted:
logger.info("清理完成: 删除 %d 个过期文件", deleted)
logger.info("清理任务已停止")
@app.get("/gateway-health", include_in_schema=False)
async def gateway_health():
"""网关自身健康检查(区别于 worker 的 /health)。"""
status = _get_pool_status_safe()
return {
"status": "ok",
"service": "gateway",
"workers_total": status["total"],
"workers_healthy": status["healthy"],
"workers_busy": status["busy"],
}
@app.get("/health", include_in_schema=False)
async def health():
"""兼容旧 /health 路径,返回网关状态。"""
return await gateway_health()
@app.get("/", include_in_schema=False)
async def index():
return {
"service": "旷视五接口 — 网关",
"version": "0.1.0",
"docs": "/docs",
"integration_guide": "/static/integration.html",
"test_pages": {
"if1_measure": "/static/test_interface1.html",
"if2_hair_grow": "/static/test_interface2.html",
"if3_hair_grow_b": "/static/test_interface3.html",
"if4_features": "/static/test_interface4.html",
"if5_hairline": "/static/test_interface5.html",
"if6_measure_v2": "/static/test_interface6.html",
"if7_hair_grow_v2": "/static/test_interface7.html",
},
}
# ---------------------------------------------------------------------------
# 代理路由
# ---------------------------------------------------------------------------
# 所有接口统一走「选 worker → 转发 → 改写 base64 → 返回」链路。
# 使用 Request 对象直接读取并转发,不做业务入参解析(解析在 worker 侧完成)。
# Form/File 声明保留在 OpenAPI extra 中以便文档生成。
def _proxy(request: Request, path: str):
"""延迟导入 proxy_request。"""
from gateway.forward import proxy_request
return proxy_request(request, path)
# 声明各接口的 form 参数用于 OpenAPI schema(实际转发直接读 Request
_MEASURE_FORMS = {
"image_file": {"type": "file", "description": "上传图片文件(JPG/PNG"},
"image_url": {"type": "string", "description": "图片 URL"},
"image_base64": {"type": "string", "description": "图片 base64(需带前缀)"},
}
_GROW_FORMS = {
**_MEASURE_FORMS,
"beauty_enabled": {"type": "boolean", "description": "是否开启美颜效果"},
}
_GROW_B_FORMS = {
"marked_image_file": {"type": "file", "description": "划线图片文件"},
"marked_image_url": {"type": "string", "description": "划线图片 URL"},
"marked_image_base64": {"type": "string", "description": "划线图片 base64"},
}
@app.post("/api/v1/face/measure", tags=["人脸分析"])
async def face_measure(request: Request):
"""接口1:四庭七眼测量标注"""
return await _proxy(request, "/api/v1/face/measure")
@app.post("/api/v1/face/measure-v2", tags=["人脸分析"])
async def face_measure_v2(request: Request):
"""接口6:四庭七眼测量标注 v2(去顶庭 + 去头部端线)"""
return await _proxy(request, "/api/v1/face/measure-v2")
@app.post("/api/v1/hair/grow", tags=["生发"])
async def hair_grow(request: Request):
"""接口2C端生发"""
return await _proxy(request, "/api/v1/hair/grow")
@app.post("/api/v1/hair/grow-b", tags=["生发"])
async def hair_grow_b(request: Request):
"""接口3B端生发"""
return await _proxy(request, "/api/v1/hair/grow-b")
@app.post("/api/v1/face/features", tags=["人脸分析"])
async def face_features(
image_file: Optional[UploadFile] = File(default=None, description="上传图片文件(JPG/PNG"),
image_url: Optional[str] = Form(default=None, description="图片 URL"),
image_base64: Optional[str] = Form(default=None, description="图片 base64(需带前缀)"),
):
"""接口4:用户特征分析 — 本机直接调豆包视觉模型,不经过 worker。"""
import uuid as _uuid
# 三选一校验
provided = [x for x in (image_file, image_url, image_base64) if x]
if len(provided) != 1:
return JSONResponse(status_code=200, content={
"code": 1007, "message": "图片参数错误:必须且只能传 image_file / image_url / image_base64 其中一个",
"request_id": f"gw-{_uuid.uuid4().hex[:8]}", "data": None,
})
img_bytes = None
if image_file:
img_bytes = await image_file.read()
elif image_base64:
b64 = image_base64
if "," in b64:
b64 = b64.split(",", 1)[1]
try:
img_bytes = base64.b64decode(b64)
except Exception:
return JSONResponse(status_code=200, content={
"code": 1008, "message": "图片格式不支持(base64 解码失败)",
"request_id": f"gw-{_uuid.uuid4().hex[:8]}", "data": None,
})
from fastapi.concurrency import run_in_threadpool
from face_features import analyze_features
try:
feats = await run_in_threadpool(analyze_features, img_bytes, image_url)
except Exception as ex:
logger.exception("接口4 豆包调用失败")
return JSONResponse(status_code=200, content={
"code": 1007, "message": f"分析服务异常:{ex}",
"request_id": f"gw-{_uuid.uuid4().hex[:8]}", "data": None,
})
if feats is None:
return JSONResponse(status_code=200, content={
"code": 1001, "message": "无法识别人像",
"request_id": f"gw-{_uuid.uuid4().hex[:8]}", "data": None,
})
return JSONResponse(status_code=200, content={
"code": 0,
"message": "success",
"request_id": f"gw-{_uuid.uuid4().hex[:8]}",
"data": {"features": json.dumps(feats, ensure_ascii=False)},
})
@app.post("/api/v1/hairline/generate", tags=["人脸分析"])
async def hairline_generate(request: Request):
"""接口5:发际线PNG生成"""
return await _proxy(request, "/api/v1/hairline/generate")
@app.post("/api/v1/hair/grow-v2", tags=["生发"])
async def hair_grow_v2(request: Request):
"""接口7C端生发 v2add_hair2 工作流)"""
return await _proxy(request, "/api/v1/hair/grow-v2")