imcc / app.py
mwauranjorogekelvin's picture
Update app.py
f1f87f1 verified
Raw
History Blame Contribute Delete
6.33 kB
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
)