Monike123/Document_API_End_points
0
1# app.py2import os3import io4from pathlib import Path5from dotenv import load_dotenv6 7from flask import Flask, request, jsonify8from flask_cors import CORS9from flask_limiter import Limiter10from flask_limiter.util import get_remote_address11from ultralytics import YOLO12from PIL import Image13 14# Load .env when running locally (do not commit .env)15load_dotenv()16 17# -----------------------18# Config from environment19# -----------------------20FLASK_ENV = os.getenv("FLASK_ENV", "production")21SECRET_KEY = os.getenv("SECRET_KEY") # must be set in host secrets22CORS_ORIGINS = os.getenv("CORS_ORIGINS", "") # comma separated list of allowed origins23 24# Per-model API keys (set these in the host as secrets)25API_KEY_AADHAAR = os.getenv("HF_AADHAAR_API_KEY") # token for Aadhaar endpoint26API_KEY_PAN = os.getenv("HF_PAN_API_KEY") # token for PAN endpoint27API_KEY_DL = os.getenv("HF_DRIVING_LICENSE_API_KEY") # token for DL endpoint28 29# Model paths (relative to repo; put .pt files under models/)30AADHAAR_MODEL_PATH = os.getenv("AADHAAR_MODEL_PATH", "./models/best.pt")31PAN_MODEL_PATH = os.getenv("PAN_MODEL_PATH", "./models/Pan_best.pt")32DL_MODEL_PATH = os.getenv("DL_MODEL_PATH", "./models/DL_best.pt")33 34# Rate limiting - adjust as needed35DEFAULT_RATE = os.getenv("DEFAULT_RATE", "30/minute") # 30 requests per minute per IP36 37# -----------------------38# App creation & security39# -----------------------40app = Flask(__name__)41app.secret_key = SECRET_KEY or os.urandom(24) # fallback for local dev only42 43# Configure CORS44allowed_origins = [o.strip() for o in CORS_ORIGINS.split(",") if o.strip()]45if not allowed_origins:46 # safer default: no origins allowed unless explicitly set47 cors = CORS(app, resources={r"/*": {"origins": []}})48else:49 cors = CORS(app, resources={r"/*": {"origins": allowed_origins}}, supports_credentials=True)50 51# Rate limiter (prevents simple abuse)52limiter = Limiter(53 key_func=get_remote_address,54 default_limits=[DEFAULT_RATE],55 storage_uri="memory://", # default in-memory store (ok for small apps)56)57limiter.init_app(app)58 59# -----------------------60# Load models at startup61# -----------------------62models = {}63model_map = {64 "aadhaar": Path(AADHAAR_MODEL_PATH),65 "pan": Path(PAN_MODEL_PATH),66 "dl": Path(DL_MODEL_PATH),67}68 69def safe_load_model(key, path: Path):70 if not path.exists():71 app.logger.warning(f"Model file not found for {key}: {path}")72 return None73 try:74 app.logger.info(f"Loading model for {key} from {path.name}")75 model = YOLO(str(path))76 return model77 except Exception as e:78 app.logger.error(f"Failed to load model {key}: {e}")79 return None80 81for k,p in model_map.items():82 models[k] = safe_load_model(k, p)83 84# -----------------------85# Utility: inference86# -----------------------87def run_inference(model, pil_image):88 """Run ultralytics YOLO inference on a PIL image and return simple JSON."""89 if model is None:90 return []91 results = model(pil_image) # ultralytics lets you pass PIL Image92 r = results[0] # first (and usually only) result93 boxes = getattr(r, "boxes", None)94 names = getattr(r, "names", {})95 detections = []96 if boxes is None:97 return detections98 99 for box in boxes:100 # ultralytics Box object fields depend on package version; defensive access:101 try:102 xyxy = box.xyxy.cpu().numpy().tolist()103 if isinstance(xyxy[0], list):104 xyxy = xyxy[0]105 106 except Exception:107 xyxy = getattr(box, "xyxy", None)108 if xyxy is None:109 continue110 try:111 conf = float(box.conf.cpu().numpy()) if hasattr(box, "conf") else None112 except Exception:113 conf = None114 try:115 cls = int(box.cls.cpu().numpy()) if hasattr(box, "cls") else None116 except Exception:117 cls = None118 label = names.get(cls, str(cls)) if cls is not None else None119 detections.append({120 "box": [float(x) for x in xyxy],121 "confidence": conf,122 "class": label123 })124 return detections125 126# -----------------------127# Auth helper128# -----------------------129def check_api_key(model_key):130 """Return True if request has a valid x-api-key header for the given model_key."""131 incoming = request.headers.get("x-api-key", "")132 if model_key == "aadhaar":133 return bool(incoming and API_KEY_AADHAAR and incoming == API_KEY_AADHAAR)134 if model_key == "pan":135 return bool(incoming and API_KEY_PAN and incoming == API_KEY_PAN)136 if model_key == "dl":137 return bool(incoming and API_KEY_DL and incoming == API_KEY_DL)138 return False139 140# -----------------------141# Routes142# -----------------------143@app.route("/", methods=["GET"])144def health():145 loaded = [k for k,v in models.items() if v is not None]146 return jsonify({"status": "ok", "models_loaded": loaded})147 148@app.route("/predict/<model_key>", methods=["POST"])149@limiter.limit(DEFAULT_RATE)150def predict(model_key):151 model_key = model_key.lower()152 if model_key not in models:153 return jsonify({"error": "unknown model key"}), 404154 155 # Authentication: must provide x-api-key header matching the env secret156 if not check_api_key(model_key):157 return jsonify({"error": "unauthorized"}), 401158 159 if "image" not in request.files:160 return jsonify({"error": "no image provided; use multipart form field 'image'"}), 400161 file = request.files["image"]162 try:163 img = Image.open(io.BytesIO(file.read())).convert("RGB")164 except Exception as e:165 return jsonify({"error": "invalid image", "detail": str(e)}), 400166 167 try:168 detections = run_inference(models[model_key], img)169 return jsonify({"predictions": detections})170 except Exception as e:171 app.logger.exception("Inference failed")172 return jsonify({"error": "inference failed", "detail": str(e)}), 500173 174# -----------------------175# Run (for local dev)176# -----------------------177if __name__ == "__main__":178 # Port 7860 is common on HF Spaces; change if needed179 app.run(host="0.0.0.0", port=int(os.getenv("PORT", 7860)))