"""FastAPI 路由""" from __future__ import annotations import platform from pathlib import Path from typing import List from fastapi import APIRouter, File, HTTPException, UploadFile from app import __version__ from app.config import settings from app.core import get_engine from app.models import HealthResponse, InvoiceResult, PathRecognizeRequest from app.services.recognize_service import recognize_file, recognize_path router = APIRouter() # ---------- 路径白名单 ---------- def _parse_allowed_dirs() -> List[Path]: """解析 ALLOWED_DIRS 配置为绝对路径列表""" if not settings.allowed_dirs.strip(): return [] sep = ";" if platform.system() == "Windows" else ":" roots: List[Path] = [] for raw in settings.allowed_dirs.split(sep): raw = raw.strip().strip('"').strip("'") if not raw: continue try: p = Path(raw).resolve() if p.is_dir(): roots.append(p) except Exception: pass return roots def _check_path_allowed(file_path: Path) -> None: """校验路径在白名单内(路径遍历攻击防护) resolve 后必须是某个 allowed_dir 的子路径。 """ roots = _parse_allowed_dirs() if not roots: raise HTTPException( status_code=403, detail="路径接口未启用:在 .env 配置 ALLOWED_DIRS 后重启服务", ) try: abs_path = file_path.resolve() except Exception as e: raise HTTPException(status_code=400, detail=f"路径无效: {e}") for root in roots: try: abs_path.relative_to(root) return except ValueError: continue raise HTTPException( status_code=403, detail=f"路径不在白名单内(允许: {', '.join(str(r) for r in roots)})", ) # ---------- 路由 ---------- @router.get("/health", response_model=HealthResponse, summary="健康检查") def health(): engine_ok = True try: get_engine() except Exception: engine_ok = False return HealthResponse( status="ok" if engine_ok else "degraded", version=__version__, engine_ready=engine_ok, ) @router.post("/recognize/invoice", response_model=InvoiceResult, summary="识别发票(上传文件)") async def recognize_invoice(file: UploadFile = File(..., description="发票图片或 PDF")): """识别发票并返回结构化字段 支持:PNG/JPG/JPEG/BMP/WEBP/TIFF/PDF """ content = await file.read() max_bytes = settings.max_upload_mb * 1024 * 1024 if len(content) > max_bytes: raise HTTPException(status_code=413, detail=f"文件超过 {settings.max_upload_mb}MB 限制") if not content: raise HTTPException(status_code=400, detail="文件为空") return recognize_file(file.filename or "unknown", content) @router.post("/recognize/invoice/by-path", response_model=InvoiceResult, summary="识别发票(服务器本地路径)") def recognize_invoice_by_path(req: PathRecognizeRequest): """传入服务器本地路径识别发票(避免重复上传大文件) **安全**:路径必须在 .env 的 ALLOWED_DIRS 白名单内才会被执行。 防止任意文件读取 / 路径遍历攻击。 """ p = Path(req.file_path) _check_path_allowed(p) if not p.exists(): raise HTTPException(status_code=404, detail=f"文件不存在: {p}") if not p.is_file(): raise HTTPException(status_code=400, detail=f"不是文件: {p}") return recognize_path(p, delete_after=False) @router.post("/recognize/text", summary="仅做字段抽取(不上传文件)") def recognize_text(raw_text: str): """对已有的 OCR 文本做字段抽取(便于接入其他 OCR 引擎)""" from app.services import extract_invoice fields = extract_invoice(raw_text) return {"fields": fields}