336 lines
12 KiB
Python
336 lines
12 KiB
Python
"""外网网关 — 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",
|
||
"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,≤ 1 MB)"},
|
||
"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):
|
||
"""接口2:C端生发"""
|
||
return await _proxy(request, "/api/v1/hair/grow")
|
||
|
||
|
||
@app.post("/api/v1/hair/grow-b", tags=["生发"])
|
||
async def hair_grow_b(request: Request):
|
||
"""接口3:B端生发"""
|
||
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,≤ 1 MB)"),
|
||
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:
|
||
raw = await image_file.read()
|
||
# TODO: 临时取消限制,后续恢复
|
||
# if len(raw) > 1_000_000:
|
||
# return JSONResponse(status_code=200, content={
|
||
# "code": 1006, "message": "文件超出 1 MB 限制",
|
||
# "request_id": f"gw-{_uuid.uuid4().hex[:8]}", "data": None,
|
||
# })
|
||
img_bytes = raw
|
||
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):
|
||
"""接口7:C端生发 v2(add_hair2 工作流)"""
|
||
return await _proxy(request, "/api/v1/hair/grow-v2")
|