apvclassifier / main.py
PrathameshRaut's picture
Update main.py
2e49730 verified
Raw
History Blame Contribute Delete
5.86 kB
import os
import json
import tempfile
import numpy as np
import tensorflow as tf
from tensorflow.keras.preprocessing import image
from pdf2image import convert_from_path
from fastapi import FastAPI, UploadFile, Form, HTTPException, Header, Depends
from fastapi.responses import JSONResponse
# -----------------------------
# CONFIG
# -----------------------------
IMG_SIZE = (224, 224)
MODEL_PATH = "./final_model.keras"
class_names = ['Other', 'Aadhaar Card', 'Pan Card', 'Voter Id']
type_mapping = {
"aadhaar_card": "Aadhaar Card",
"pan_card": "Pan Card",
"voter_id": "Voter Id"
}
# -----------------------------
# AUTH CONFIG (HF SECRETS)
# -----------------------------
MASTER_SECRET_KEY = os.getenv("MASTER_SECRET_KEY")
PROJECT_KEYS_JSON = os.getenv("PROJECT_KEYS_JSON", "{}")
try:
PROJECT_KEYS = json.loads(PROJECT_KEYS_JSON)
except Exception:
PROJECT_KEYS = {}
# -----------------------------
# AUTH VALIDATION
# -----------------------------
def verify_keys(
x_secret_key: str = Header(None),
x_project_id: str = Header(None),
x_project_key: str = Header(None),
):
# Server configuration check
if not MASTER_SECRET_KEY:
raise HTTPException(status_code=500, detail="MASTER_SECRET_KEY not configured in Space secrets")
# 1) Validate master key
if not x_secret_key or x_secret_key != MASTER_SECRET_KEY:
raise HTTPException(status_code=401, detail="Invalid Secret Key")
# 2) Validate project key
if not x_project_id or not x_project_key:
raise HTTPException(status_code=401, detail="Project Id and Project Key are required")
if x_project_id not in PROJECT_KEYS:
raise HTTPException(status_code=401, detail=f"Unknown project: {x_project_id}")
if PROJECT_KEYS.get(x_project_id) != x_project_key:
raise HTTPException(status_code=401, detail="Invalid Project Key")
return True
# -----------------------------
# LOAD MODEL
# -----------------------------
model = tf.keras.models.load_model(MODEL_PATH)
# -----------------------------
# Predict a single image array (H,W,C)
# -----------------------------
def predict_array(img_array):
img_array = tf.cast(img_array, tf.float32)
img_array = tf.image.resize(img_array, IMG_SIZE)
img_array = tf.expand_dims(img_array, axis=0)
pred = model.predict(img_array, verbose=0)[0]
class_id = int(np.argmax(pred))
conf = float(np.max(pred))
return class_names[class_id], conf, pred
# -----------------------------
# Predict from image path
# -----------------------------
def predict_image(img_path):
img = image.load_img(img_path, target_size=IMG_SIZE, color_mode="rgb")
img_array = image.img_to_array(img)
label, conf, _ = predict_array(img_array)
return {
"type": label,
"confidence": round(conf, 6)
}
# -----------------------------
# Predict from PDF (max 2 pages only)
# -----------------------------
def predict_pdf(pdf_path, dpi=200, max_pages=2):
pages = convert_from_path(pdf_path, dpi=dpi)
if len(pages) > max_pages:
return {
"error": "Maximum page limit reached",
"max_pages": max_pages,
"found_pages": len(pages)
}
results = []
labels = []
for page in pages:
page_np = np.array(page)
# RGBA -> RGB
if page_np.shape[-1] == 4:
page_np = page_np[..., :3]
label, conf, _ = predict_array(page_np)
results.append((label, conf))
labels.append(label)
# 2 pages validation
if len(labels) == 2:
if labels[0] != labels[1]:
return {
"error": "Invalid input",
"reason": "Pages belong to different document types",
"page1_type": labels[0],
"page2_type": labels[1]
}
final_label = labels[0]
final_conf = float(max(results[0][1], results[1][1]))
return {
"type": final_label,
"confidence": round(final_conf, 6)
}
# 1 page output
return {
"type": labels[0],
"confidence": round(float(results[0][1]), 6)
}
# -----------------------------
# Main function
# -----------------------------
def predict_file(path):
if not os.path.exists(path):
return {"error": "File not found", "path": path}
ext = os.path.splitext(path)[1].lower()
if ext in [".png", ".jpg", ".jpeg", ".bmp", ".webp"]:
return predict_image(path)
if ext == ".pdf":
return predict_pdf(path)
return {"error": "Unsupported file type", "ext": ext}
# -----------------------------
# FastAPI App
# -----------------------------
app = FastAPI(
title="ID Document Validator",
description="Predict document type and authenticate based on expected type."
)
@app.post("/predict")
async def predict(
file: UploadFile,
expected_type: str = Form(None),
auth: bool = Depends(verify_keys) # ✅ Authentication required
):
if not file.filename:
raise HTTPException(status_code=400, detail="No file provided")
with tempfile.NamedTemporaryFile(delete=False, suffix=os.path.splitext(file.filename)[1]) as temp_file:
temp_path = temp_file.name
temp_file.write(await file.read())
try:
result = predict_file(temp_path)
# authentication based on expected_type
if expected_type and expected_type in type_mapping and "type" in result:
if result["type"] == type_mapping[expected_type]:
result["Authentication"] = "Valid"
else:
result["Authentication"] = "Not valid"
return JSONResponse(content=result)
finally:
if os.path.exists(temp_path):
os.unlink(temp_path)
@app.get("/")
def read_root():
return {"message": "ID Validator API is running"}