Files
pifu/api_server.py
2026-08-17 10:03:37 +08:00

429 lines
14 KiB
Python
Raw Permalink 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.
#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
皮肤分析 API 接口服务
"""
import sys
import os
import threading
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from fastapi import FastAPI, HTTPException, BackgroundTasks, Query
from pydantic import BaseModel, Field
import uvicorn
import logging
from datetime import datetime
from typing import Optional
import subprocess
from pathlib import Path
from config.settings import (
ANALYSIS_ROOT,
API_HOST,
API_CALLBACK,
MODEL_PATH,
PORE_ENABLED,
PORE_PREDICTOR_PATH,
)
# ==================== 日志配置 ====================
logging.basicConfig(
level=logging.INFO,
format='%(asctime)s | %(levelname)s | %(message)s'
)
logger = logging.getLogger(__name__)
# ==================== FastAPI 应用 ====================
app = FastAPI(
title="皮肤分析任务接口",
description="接收文件夹名称,触发皮肤分析任务",
version="1.0.0"
)
# ==================== 数据存储 ====================
task_store = {}
active_tasks = set()
task_store_lock = threading.RLock()
# ==================== 数据模型 ====================
class TaskRequest(BaseModel):
folder_name: str = Field(..., description="文件夹全路径")
task_id: str = Field(..., description="目标文件夹名称(任务 ID")
task_type: Optional[str] = Field(
default=None,
description="任务类型,pore 表示仅执行大毛孔检测",
)
class TaskResponse(BaseModel):
code: int = Field(..., description="状态码")
message: str = Field(..., description="消息")
task_id: str = Field(..., description="任务 ID")
status: str = Field(..., description="任务状态")
class TaskStatus(BaseModel):
task_id: str = Field(..., description="任务 ID")
status: str = Field(..., description="状态")
created_at: str = Field(..., description="创建时间")
completed_at: Optional[str] = Field(default=None, description="完成时间")
error: Optional[str] = Field(default=None, description="错误信息")
class CallbackRequest(BaseModel):
task_id: str = Field(..., description="任务 ID")
status: int = Field(..., ge=0, le=1, description="处理状态,1=成功,0=失败")
def validate_task_id(task_id: str) -> str:
"""校验会进入文件路径、日志和命令参数的任务 ID。"""
if not task_id or len(task_id) < 3 or len(task_id) > 128:
raise ValueError("任务 ID 格式错误,长度必须为 3 到 128 个字符")
if task_id in {".", ".."} or task_id.startswith("-"):
raise ValueError("任务 ID 格式错误")
if any(char in task_id for char in ("/", "\\")):
raise ValueError("任务 ID 不能包含路径分隔符")
if any(ord(char) < 32 or ord(char) == 127 for char in task_id):
raise ValueError("任务 ID 不能包含控制字符")
return task_id
def normalize_task_type(task_type: Optional[str]) -> Optional[str]:
"""规范化可选任务类型;空值保持现有完整分析流程。"""
normalized = (task_type or "").strip().lower()
if not normalized:
return None
if normalized != "pore":
raise ValueError(f"不支持的任务类型:{task_type}")
return normalized
def resolve_input_dir(folder_name: str, task_id: str, now: Optional[datetime] = None) -> Path:
"""解析输入目录,并确保目录始终位于分析数据根目录内。"""
root = Path(ANALYSIS_ROOT).resolve()
candidate = Path(folder_name)
if not candidate.is_absolute():
candidate = root / (now or datetime.now()).strftime("%Y%m%d") / task_id
candidate = candidate.resolve()
try:
candidate.relative_to(root)
except ValueError as exc:
raise ValueError(f"输入目录必须位于 {root}") from exc
if candidate == root:
raise ValueError("输入目录不能是分析数据根目录本身")
return candidate
def update_task(task_id: str, **updates) -> None:
"""线程安全地更新仍存在的任务。"""
with task_store_lock:
task = task_store.get(task_id)
if task is not None:
task.update(updates)
def finish_task(task_id: str, status: str, error: Optional[str]) -> None:
"""原子写入最终状态并解除运行标记,避免迟到回调覆盖结果。"""
with task_store_lock:
task = task_store.get(task_id)
if task is not None:
task.update(
status=status,
error=error,
completed_at=datetime.now().isoformat(),
)
active_tasks.discard(task_id)
# ==================== 后台任务 ====================
def run_analysis_task(
task_id: str,
request: TaskRequest,
input_dir: Path,
task_type: Optional[str] = None,
):
"""后台运行分析任务"""
logger.info(f"🚀 开始执行任务:{task_id}")
try:
update_task(task_id, status="running")
# 再次解析路径,防止任务排队期间目录被替换为越界符号链接。
revalidated_input_dir = resolve_input_dir(str(input_dir), task_id)
if revalidated_input_dir != input_dir:
raise Exception("输入目录在任务执行前发生变化")
input_dir = revalidated_input_dir
# 检查目录是否存在
if not input_dir.is_dir():
raise Exception(f"输入目录不存在或不是目录: {input_dir}")
# ==================== 脚本路径 ====================
script_path = Path(__file__).parent / "wanzheng.py"
if not script_path.exists():
raise Exception(f"分析脚本不存在: {script_path}")
# ==================== 模型路径 ====================
pore_only = task_type == "pore"
model_path = Path(MODEL_PATH)
if not pore_only and not model_path.is_file():
raise Exception(f"模型文件不存在: {model_path}")
# ==================== 构建命令 ====================
cmd = [
sys.executable,
str(script_path),
str(input_dir),
"--workers", "4",
"--analysis-root", str(Path(ANALYSIS_ROOT).resolve()),
"--api-url", f"{API_HOST.rstrip('/')}/{API_CALLBACK.lstrip('/')}",
"--task-id", task_id,
]
if pore_only:
cmd.extend([
"--task-type", "pore",
"--pore-predictor", PORE_PREDICTOR_PATH,
])
else:
cmd.extend(["--model", str(model_path)])
if PORE_ENABLED:
cmd.extend(["--pore-predictor", PORE_PREDICTOR_PATH])
else:
cmd.append("--disable-pores")
logger.info(f"🔧 命令参数:{cmd}")
logger.info(f"📋 执行命令:{' '.join(cmd)}")
logger.info(f"📁 输入目录:{input_dir}")
# ==================== 执行脚本 ====================
result = subprocess.run(
cmd,
capture_output=True,
text=True,
encoding="utf-8",
errors="replace",
timeout=3600
)
# ==================== 输出日志 ====================
if result.stdout:
logger.info(f"📋 标准输出:\n{result.stdout}")
if result.stderr:
if result.returncode == 0:
logger.info(f"📋 进程日志:\n{result.stderr}")
else:
logger.error(f"📋 错误输出:\n{result.stderr}")
# ==================== 判断执行结果 ====================
if result.returncode == 0:
finish_task(task_id, status="completed", error=None)
logger.info(f"✅ 任务完成:{task_id}")
else:
finish_task(
task_id,
status="failed",
error=result.stderr or f"分析进程退出码:{result.returncode}"
)
logger.error(
f"❌ 任务失败:{task_id} "
f"(code={result.returncode})"
)
except subprocess.TimeoutExpired:
finish_task(
task_id,
status="failed",
error="任务超时(超过1小时)"
)
logger.error(f"⏰ 任务超时:{task_id}")
except Exception as e:
finish_task(
task_id,
status="failed",
error=str(e)
)
logger.exception(f"❌ 任务异常:{task_id}")
# ==================== API 接口 ====================
@app.get("/")
async def root():
return {
"service": "皮肤分析任务接口",
"version": "1.0.0",
"status": "running"
}
@app.get("/health")
async def health_check():
with task_store_lock:
statuses = [task["status"] for task in task_store.values()]
return {
"status": "healthy",
"timestamp": datetime.now().isoformat(),
"tasks": {
"total": len(statuses),
"pending": statuses.count("pending"),
"running": statuses.count("running"),
"completed": statuses.count("completed"),
"failed": statuses.count("failed")
}
}
@app.post("/api/v1/analysis/submit", response_model=TaskResponse)
async def submit_task(request: TaskRequest, background_tasks: BackgroundTasks):
# 使用传入的 task_id,如果没有则从 folder_name 提取
task_id = request.task_id if request.task_id else Path(request.folder_name).name
try:
validate_task_id(task_id)
task_type = normalize_task_type(request.task_type)
input_dir = resolve_input_dir(request.folder_name, task_id)
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
with task_store_lock:
if task_id in active_tasks:
existing_task = task_store[task_id]
return TaskResponse(
code=200,
message="任务已存在,正在运行中",
task_id=task_id,
status=existing_task["status"]
)
task_store[task_id] = {
"task_id": task_id,
"status": "pending",
"created_at": datetime.now().isoformat(),
"completed_at": None,
"error": None,
"request": {
"folder_name": request.folder_name,
"task_id": request.task_id,
"task_type": task_type,
}
}
active_tasks.add(task_id)
background_tasks.add_task(
run_analysis_task,
task_id,
request,
input_dir,
task_type,
)
logger.info(f"📝 任务已提交:{task_id}")
return TaskResponse(
code=200,
message="任务提交成功",
task_id=task_id,
status="pending"
)
@app.post("/api/v1/analysis/callback")
async def analysis_callback(request: CallbackRequest):
try:
task_id = validate_task_id(request.task_id)
except ValueError as exc:
raise HTTPException(status_code=400, detail=str(exc)) from exc
callback_status = "completed" if request.status == 1 else "failed"
with task_store_lock:
if task_id not in task_store:
task_store[task_id] = {
"task_id": task_id,
"status": callback_status,
"created_at": datetime.now().isoformat(),
"completed_at": datetime.now().isoformat(),
"error": None if request.status == 1 else "分析任务回调失败",
"request": None
}
else:
current_status = task_store[task_id]["status"]
is_finished = current_status in {"completed", "failed"} and task_id not in active_tasks
if not is_finished:
task_store[task_id]["status"] = callback_status
task_store[task_id]["completed_at"] = datetime.now().isoformat()
task_store[task_id]["error"] = (
None if request.status == 1 else "分析任务回调失败"
)
status = task_store[task_id]["status"]
logger.info(f"📩 收到分析回调:{task_id} status={request.status}")
return {"code": 200, "message": "回调接收成功", "task_id": task_id, "status": status}
@app.get("/api/v1/analysis/status/{task_id}", response_model=TaskStatus)
async def get_task_status(task_id: str):
with task_store_lock:
task = task_store.get(task_id)
if task is not None:
task = dict(task)
if task is None:
raise HTTPException(status_code=404, detail=f"任务不存在:{task_id}")
return TaskStatus(
task_id=task["task_id"],
status=task["status"],
created_at=task["created_at"],
completed_at=task.get("completed_at"),
error=task.get("error")
)
@app.get("/api/v1/analysis/list")
async def list_tasks(
status: Optional[str] = Query(default=None),
limit: int = Query(default=20, ge=1, le=100)
):
with task_store_lock:
tasks = [dict(task) for task in task_store.values()]
if status:
tasks = [t for t in tasks if t["status"] == status]
tasks = sorted(tasks, key=lambda x: x["created_at"], reverse=True)
total = len(tasks)
tasks = tasks[:limit]
return {
"total": total,
"tasks": tasks
}
# ==================== 启动服务 ====================
if __name__ == "__main__":
print("="*60)
print("皮肤分析 API 服务")
print("="*60)
print("监听地址:http://0.0.0.0:8071")
print("提交接口:POST /api/v1/analysis/submit")
print("状态查询:GET /api/v1/analysis/status/{task_id}")
print("健康检查:GET /health")
print("="*60)
print()
uvicorn.run(app, host="0.0.0.0", port=8071)