CoolFace
Apppublic

Monike123/Document_API_End_points

sourceHugging Facemitupdated 11mo agoView on Hugging Face
0likes
app.py179 linesDownload Raw Back to root
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)))