| 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__) |
|
|
| |
| 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.set_grad_enabled(False) |
| torch.set_num_threads(1) |
| torch.set_num_interop_threads(1) |
|
|
| |
| 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 |
|
|
| |
| 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") |
|
|
| |
| 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, :] |
|
|
|
|
| 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] |
|
|
|
|
| |
| 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") |
|
|
|
|
| |
| @asynccontextmanager |
| async def lifespan(app: FastAPI): |
| warmup() |
| yield |
|
|
|
|
| |
| class SolveRequest(BaseModel): |
| image_base64: str |
|
|
| class SolveResponse(BaseModel): |
| success: bool |
| text: str = "" |
| processing_time: float = 0.0 |
| error: str = "" |
|
|
|
|
| |
| app = FastAPI(title="CAPTCHA Solver API", lifespan=lifespan) |
|
|
| app.add_middleware( |
| CORSMiddleware, |
| allow_origins=["*"], |
| allow_methods=["*"], |
| allow_headers=["*"], |
| ) |
|
|
|
|
| |
| @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 |
| ) |