mwauranjorogekelvin commited on
Commit
42eced1
Β·
verified Β·
1 Parent(s): 4d7f40d

Update app.py

Browse files

Add FastAPI endpoint with API key auth

Files changed (1) hide show
  1. app.py +107 -28
app.py CHANGED
@@ -1,36 +1,115 @@
 
 
 
 
 
 
 
 
1
  import gradio as gr
2
- from transformers import VisionEncoderDecoderModel, TrOCRProcessor
3
- import torch
 
 
4
  from PIL import Image
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
5
 
6
- def recognize_captcha(input, mdl):
7
-
8
- # Load model and processor
9
- processor = TrOCRProcessor.from_pretrained(mdl)
10
- model = VisionEncoderDecoderModel.from_pretrained(mdl)
11
-
12
- # Prepare image
13
- pixel_values = processor(input, return_tensors="pt").pixel_values
14
-
15
- # Generate text
16
- generated_ids = model.generate(pixel_values)
17
- generated_text = processor.batch_decode(generated_ids, skip_special_tokens=True)[0]
18
-
19
- return generated_text
20
 
 
21
  iface = gr.Interface(
22
  fn=recognize_captcha,
23
- inputs=[
24
- gr.Image(),
25
- gr.Dropdown(
26
- ['anuashok/ocr-captcha-v3','anuashok/ocr-captcha-v2','anuashok/ocr-captcha-v1','microsoft/trocr-base-printed'], label='Model to use'
27
- )
28
- ],
29
- outputs=['text'],
30
- title = "Character Sequence Recognition From Captcha Image",
31
- description = "Using some TrOCR models found on the HF Hub to test/break tough text captchas. Will you have to train your own?",
32
- article="Created by Neeraj with ❀️ !!!"
33
  )
34
 
35
- iface.queue(max_size=10)
36
- iface.launch()
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import io
3
+ import time
4
+ import base64
5
+ import secrets
6
+ import logging
7
+ from functools import lru_cache
8
+
9
  import gradio as gr
10
+ from fastapi import FastAPI, HTTPException, Security, Depends
11
+ from fastapi.security.api_key import APIKeyHeader
12
+ from fastapi.middleware.cors import CORSMiddleware
13
+ from pydantic import BaseModel
14
  from PIL import Image
15
+ import torch
16
+ from transformers import VisionEncoderDecoderModel, TrOCRProcessor
17
+
18
+ logging.basicConfig(level=logging.INFO)
19
+ logger = logging.getLogger(__name__)
20
+
21
+ # ── Config ────────────────────────────────────────────────────────────────────
22
+ API_KEY = os.environ.get("API_KEY", "changeme")
23
+ DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
24
+ MODELS = [
25
+ 'anuashok/ocr-captcha-v3',
26
+ 'anuashok/ocr-captcha-v2',
27
+ 'anuashok/ocr-captcha-v1',
28
+ 'microsoft/trocr-base-printed'
29
+ ]
30
+ DEFAULT_MODEL = MODELS[0]
31
+
32
+ # ── API Key Auth ──────────────────────────────────────────────────────────────
33
+ api_key_header = APIKeyHeader(name="X-API-Key", auto_error=False)
34
+
35
+ def verify_key(key: str = Security(api_key_header)):
36
+ if not key or not secrets.compare_digest(key, API_KEY):
37
+ raise HTTPException(status_code=401, detail="Invalid or missing API key")
38
+ return key
39
+
40
+ # ── Model Cache ───────────────────────────────────────────────────────────────
41
+ @lru_cache(maxsize=4)
42
+ def load_model(model_id: str):
43
+ logger.info(f"Loading {model_id} on {DEVICE}...")
44
+ processor = TrOCRProcessor.from_pretrained(model_id)
45
+ model = VisionEncoderDecoderModel.from_pretrained(model_id).to(DEVICE)
46
+ model.eval()
47
+ logger.info(f"βœ… {model_id} ready")
48
+ return processor, model
49
 
50
+ # ── Inference ─────────────────────────────────────────────────────────────────
51
+ def recognize_captcha(image, model_id):
52
+ processor, model = load_model(model_id)
53
+ pixel_values = processor(image, return_tensors="pt").pixel_values.to(DEVICE)
54
+ with torch.no_grad():
55
+ generated_ids = model.generate(pixel_values, max_new_tokens=32)
56
+ return processor.batch_decode(generated_ids, skip_special_tokens=True)[0]
 
 
 
 
 
 
 
57
 
58
+ # ── Gradio UI ─────────────────────────────────────────────────────────────────
59
  iface = gr.Interface(
60
  fn=recognize_captcha,
61
+ inputs=[gr.Image(), gr.Dropdown(MODELS, label="Model", value=DEFAULT_MODEL)],
62
+ outputs=["text"],
63
+ title="CAPTCHA Solver",
64
+ description="API available at /solve-captcha-base64"
 
 
 
 
 
 
65
  )
66
 
67
+ # ── FastAPI ───────────────────────────────────────────────────────────────────
68
+ app = gr.mount_gradio_app(
69
+ FastAPI(title="CAPTCHA Solver API"),
70
+ iface,
71
+ path="/"
72
+ )
73
+
74
+ app.add_middleware(
75
+ CORSMiddleware,
76
+ allow_origins=["*"],
77
+ allow_methods=["*"],
78
+ allow_headers=["*"],
79
+ )
80
+
81
+ # ── Schemas ───────────────────────────────────────────────────────────────────
82
+ class SolveRequest(BaseModel):
83
+ image_base64: str
84
+ model: str = DEFAULT_MODEL
85
+
86
+ class SolveResponse(BaseModel):
87
+ success: bool
88
+ text: str = ""
89
+ processing_time: float = 0.0
90
+ model_used: str = ""
91
+ error: str = ""
92
+
93
+ # ── Endpoints ─────────────────────────────────────────────────────────────────
94
+ @app.get("/health")
95
+ def health():
96
+ return {"status": "ok", "device": DEVICE, "quantized": False, "default_model": DEFAULT_MODEL}
97
+
98
+ @app.post("/solve-captcha-base64", response_model=SolveResponse)
99
+ def solve(req: SolveRequest, _: str = Depends(verify_key)):
100
+ start = time.time()
101
+ try:
102
+ raw = req.image_base64
103
+ if "," in raw:
104
+ raw = raw.split(",", 1)[1]
105
+
106
+ image = Image.open(io.BytesIO(base64.b64decode(raw))).convert("RGB")
107
+ text = recognize_captcha(image, req.model)
108
+
109
+ elapsed = time.time() - start
110
+ logger.info(f"Solved '{text}' in {elapsed:.2f}s")
111
+ return SolveResponse(success=True, text=text.strip(), processing_time=elapsed, model_used=req.model)
112
+
113
+ except Exception as e:
114
+ logger.error(f"Error: {e}")
115
+ return SolveResponse(success=False, error=str(e), processing_time=time.time() - start)