Ram0894's picture
Update app.py
d85bf0a verified
Raw
History Blame Contribute Delete
4.04 kB
import os
from fastapi import FastAPI, UploadFile, File, Form, Depends, HTTPException
from fastapi.middleware.cors import CORSMiddleware
from fastapi.staticfiles import StaticFiles
from fastapi.responses import FileResponse
from sqlalchemy.orm import Session
from database import init_db, get_db, PredictionRecord
from model_loader import preload_all_models
from predict import predict_image
app = FastAPI(title="AI-Assisted Diagnostic Support API", version="1.0.0")
# Enable CORS for frontend integration
app.add_middleware(
CORSMiddleware,
allow_origins=["*"], # For local development simplicity
allow_credentials=True,
allow_methods=["*"],
allow_headers=["*"],
)
# Startup event: Initialize database and preload models
@app.on_event("startup")
def startup_event():
print("Starting up diagnostic API...")
print("Initializing SQLite Database...")
init_db()
# Preload models to avoid lag on first request
try:
preload_all_models()
print("All models preloaded successfully and ready.")
except Exception as e:
print(f"Warning: Failed to preload models during startup: {str(e)}")
@app.get("/health")
def health_check():
return {
"status": "healthy",
"models_loaded": True,
"supported_modules": ["blood_cell", "skin_lesion"]
}
@app.post("/predict")
async def run_predict(
file: UploadFile = File(...),
module_type: str = Form(...), # 'blood_cell' or 'skin_lesion'
db: Session = Depends(get_db)
):
if module_type not in ["blood_cell", "skin_lesion"]:
raise HTTPException(status_code=400, detail="Invalid module type. Choose 'blood_cell' or 'skin_lesion'.")
try:
# Read uploaded image bytes
image_bytes = await file.read()
# Run inference pipeline
label, confidence, class_probs, heatmap_b64, inf_time = predict_image(
image_bytes=image_bytes,
filename=file.filename,
module_type=module_type
)
# Save record in the SQLite database
record = PredictionRecord(
module=module_type,
filename=file.filename,
predicted_class=label,
confidence=confidence,
notes=f"Inference took {inf_time:.4f} seconds."
)
db.add(record)
db.commit()
db.refresh(record)
return {
"id": record.id,
"timestamp": record.timestamp,
"filename": record.filename,
"prediction": label,
"confidence": confidence,
"class_probabilities": class_probs,
"heatmap_base64": heatmap_b64,
"inference_time_sec": inf_time
}
except Exception as e:
import traceback
traceback.print_exc()
raise HTTPException(status_code=500, detail=f"Inference execution failed: {str(e)}")
@app.get("/history")
def get_prediction_history(db: Session = Depends(get_db)):
records = db.query(PredictionRecord).order_by(PredictionRecord.timestamp.desc()).all()
# Format database records into JSON-friendly structures
history_list = []
for r in records:
history_list.append({
"id": r.id,
"timestamp": r.timestamp.isoformat() + "Z" if r.timestamp else None,
"module": "Blood Cell" if r.module == "blood_cell" else "Skin Lesion",
"filename": r.filename,
"prediction": r.predicted_class,
"confidence": r.confidence,
"notes": r.notes
})
return history_list
# Serve static files and index.html
static_dir = os.path.join(os.path.dirname(__file__), "static")
if os.path.exists(static_dir):
app.mount("/static", StaticFiles(directory=static_dir), name="static")
@app.get("/")
def read_index():
return FileResponse(os.path.join(static_dir, "index.html"))
if __name__ == "__main__":
import uvicorn
uvicorn.run("backend.app:app", host="0.0.0.0", port=8000, reload=True)