import os import io import time import base64 import secrets import logging import numpy as np from contextlib import asynccontextmanager from fastapi import FastAPI, HTTPException, Security, Depends from fastapi.security.api_key import APIKeyHeader from fastapi.middleware.cors import CORSMiddleware from pydantic import BaseModel from PIL import Image import torch import onnxruntime as rt from utils.tokenizer_base import Tokenizer logging.basicConfig(level=logging.INFO) logger = logging.getLogger(__name__) # ── Config ──────────────────────────────────────────────────────────────────── API_KEY = os.environ.get("API_KEY", "changeme") IMG_SIZE = (128, 32) VOCAB = r"0123456789abcdefghijklmnopqrstuvwxyzABCDEFGHIJKLMNOPQRSTUVWXYZ!\"#$%&'()*+,-./:;<=>?@[\\]^_`{|}~" MODEL_PATH = "./model/model.onnx" ORT_INTRA_THREADS = int(os.getenv("ORT_INTRA_THREADS", "1")) ORT_INTER_THREADS = int(os.getenv("ORT_INTER_THREADS", "1")) # ── Torch global optimizations ──────────────────────────────────────────────── torch.set_grad_enabled(False) # no autograd overhead on tensor ops torch.set_num_threads(1) # don't compete with ORT threads torch.set_num_interop_threads(1) # no inter-op parallelism from torch # ── API Key Auth ────────────────────────────────────────────────────────────── api_key_header = APIKeyHeader(name="X-API-Key", auto_error=False) def verify_key(key: str = Security(api_key_header)): if not key or not secrets.compare_digest(key, API_KEY): raise HTTPException(status_code=401, detail="Invalid or missing API key") return key # ── Load model ──────────────────────────────────────────────────────────────── logger.info("Loading ONNX model...") tokenizer = Tokenizer(VOCAB) opts = rt.SessionOptions() opts.intra_op_num_threads = ORT_INTRA_THREADS opts.inter_op_num_threads = ORT_INTER_THREADS opts.execution_mode = rt.ExecutionMode.ORT_SEQUENTIAL opts.graph_optimization_level = rt.GraphOptimizationLevel.ORT_ENABLE_ALL opts.optimized_model_filepath = MODEL_PATH + ".opt" opts.enable_mem_pattern = True opts.enable_cpu_mem_arena = True session = rt.InferenceSession( MODEL_PATH, sess_options=opts, providers=["CPUExecutionProvider"] ) input_name = session.get_inputs()[0].name logger.info("✅ ONNX model ready") # ── Inference ───────────────────────────────────────────────────────────────── def preprocess(image: Image.Image) -> np.ndarray: image = image.convert("RGB").resize(IMG_SIZE, Image.BILINEAR) x = np.ascontiguousarray(image, dtype=np.float32) x = (x / 255.0 - 0.5) / 0.5 return x.transpose(2, 0, 1)[np.newaxis, :] # [1, 3, H, W] def solve_image(image: Image.Image) -> str: x = preprocess(image) logits = session.run(None, {input_name: x})[0] probs = torch.tensor(logits).softmax(-1) preds, _ = tokenizer.decode(probs) return preds[0] # ── Warmup ──────────────────────────────────────────────────────────────────── def warmup(): logger.info("Warming up model...") dummy = Image.new("RGB", IMG_SIZE, color=(128, 128, 128)) for _ in range(3): solve_image(dummy) logger.info("✅ Warmup complete — model is hot") # ── Lifespan ────────────────────────────────────────────────────────────────── @asynccontextmanager async def lifespan(app: FastAPI): warmup() yield # ── Schemas ─────────────────────────────────────────────────────────────────── class SolveRequest(BaseModel): image_base64: str class SolveResponse(BaseModel): success: bool text: str = "" processing_time: float = 0.0 error: str = "" # ── FastAPI ─────────────────────────────────────────────────────────────────── app = FastAPI(title="CAPTCHA Solver API", lifespan=lifespan) app.add_middleware( CORSMiddleware, allow_origins=["*"], allow_methods=["*"], allow_headers=["*"], ) # ── Endpoints ───────────────────────────────────────────────────────────────── @app.get("/health") def health(): return { "status": "ok", "device": "cpu", "model": "ONNX INT8", "quantized": True, "workers": os.getenv("WEB_CONCURRENCY", "1"), "intra_threads": ORT_INTRA_THREADS, } @app.post("/solve-captcha-base64", response_model=SolveResponse) def solve(req: SolveRequest, _: str = Depends(verify_key)): start = time.time() try: raw = req.image_base64 if "," in raw: raw = raw.split(",", 1)[1] image = Image.open(io.BytesIO(base64.b64decode(raw))) text = solve_image(image).strip()[:5] elapsed = time.time() - start logger.info(f"✅ Solved: '{text}' in {elapsed:.3f}s") return SolveResponse(success=True, text=text, processing_time=elapsed) except Exception as e: logger.error(f"Error: {e}") return SolveResponse( success=False, error=str(e), processing_time=time.time() - start )