from __future__ import annotations import argparse import logging import os import uvicorn from fastapi import FastAPI, Request from fastapi.concurrency import run_in_threadpool from fastapi.exceptions import RequestValidationError from fastapi.responses import JSONResponse from .contracts import InvokeRequest from .errors import OCRServiceError from .service import OCRService logger = logging.getLogger(__name__) service = OCRService() def create_app(ocr_service: OCRService | None = None) -> FastAPI: active_service = ocr_service or service app = FastAPI(title="OCR Services", version="0.1.0") @app.exception_handler(OCRServiceError) async def handle_service_error( request: Request, exc: OCRServiceError ) -> JSONResponse: del request return JSONResponse( status_code=exc.status_code, content={"ok": False, "result": None, "error": exc.as_dict()}, ) @app.exception_handler(RequestValidationError) async def handle_validation_error( request: Request, exc: RequestValidationError ) -> JSONResponse: del request return JSONResponse( status_code=422, content={ "ok": False, "result": None, "error": { "code": "INVALID_REQUEST", "message": "请求结构无效", "details": {"errors": exc.errors()}, }, }, ) @app.middleware("http") async def log_unexpected_errors(request: Request, call_next): try: return await call_next(request) except Exception: logger.exception( "unexpected error while handling %s %s", request.method, request.url.path, ) raise @app.get("/health") async def health() -> dict[str, object]: return {"ok": True, "service": "ocr-services", "version": "0.1.0"} @app.post("/api/invoke") async def invoke(request: InvokeRequest) -> dict[str, object]: result = await run_in_threadpool(active_service.invoke, request) return {"ok": True, "result": result, "error": None} return app app = create_app() def main() -> None: parser = argparse.ArgumentParser(description="启动统一 OCR 服务") parser.add_argument("--host", default=os.getenv("OCR_SERVICE_HOST", "127.0.0.1")) parser.add_argument( "--port", type=int, default=int(os.getenv("OCR_SERVICE_PORT", "5020")) ) parser.add_argument("--log-level", default="info") args = parser.parse_args() logging.basicConfig( level=getattr(logging, args.log_level.upper()), format="%(asctime)s %(levelname)s %(name)s %(message)s", ) uvicorn.run(app, host=args.host, port=args.port, log_level=args.log_level) if __name__ == "__main__": main()