FeatureX-Clinical-AI / flask_app_reference.py
Lnadeem's picture
Upload 11 files
53b6de2 verified
Raw History Blame Contribute Delete
11.2 kB
import base64
import binascii
import json
import logging
import os
import re
import traceback
import unicodedata
import uuid
from typing import Any, Dict
import gradio as gr
import spaces
import database
from inference import classes_info, health_info, predict_bytes, root_info
logging.basicConfig(
level=logging.INFO,
format="%(asctime)s | %(levelname)s | %(name)s | %(message)s",
)
logger = logging.getLogger("featurex-zerogpu-server")
SERVER_VERSION = "3.0.0-hf-zerogpu-gradio"
MAX_UPLOAD = 15 * 1024 * 1024
def _json(payload: Any) -> str:
return json.dumps(payload, ensure_ascii=False, separators=(",", ":"))
def _decode_image(value: str) -> bytes:
if not value:
raise ValueError("Missing image_base64.")
text = str(value).strip()
if text.startswith("data:"):
try:
text = text.split(",", 1)[1]
except IndexError:
raise ValueError("Invalid data URL.")
text = re.sub(r"\s+", "", text)
try:
raw = base64.b64decode(text, validate=True)
except (binascii.Error, ValueError) as exc:
raise ValueError("Invalid base64 image payload.") from exc
if not raw:
raise ValueError("Decoded image is empty.")
if len(raw) > MAX_UPLOAD:
raise ValueError("Image exceeds 15MB.")
return raw
def _metadata(
patient_name: str = "",
birth_date: str = "",
gender: str = "",
age: str = "",
user_id: str = "",
username: str = "",
department: str = "",
role: str = "doctor",
supervisor_id: str = "",
) -> tuple[Dict[str, Any], Dict[str, Any]]:
def s(value: Any, default: str = "") -> str:
if value is None:
value = default
return unicodedata.normalize("NFC", str(value))
age_value = None
age_text = s(age)
if age_text:
try:
age_value = int(float(age_text))
except Exception:
age_value = None
patient = {
"patient_name": s(patient_name),
"birth_date": s(birth_date),
"gender": s(gender),
"age": age_value,
}
user = {
"user_id": s(user_id),
"username": s(username),
"department": s(department),
"role": s(role, "doctor") or "doctor",
"supervisor_id": s(supervisor_id),
}
return patient, user
def _build_unified(result: Dict[str, Any]) -> Dict[str, Any]:
reports_ar = result.get("medical_report_ar") or {}
reports_en = result.get("medical_report_en") or {}
selected = reports_en if result.get("language") == "en" else reports_ar
pred = selected.get("prediction") or {}
quality = result.get("quality_metrics") or {}
reliability = selected.get("reliability") or {}
confidence_flags = selected.get("confidence_flags") or []
confidence = float(pred.get("confidence") or 0.0)
safety_flags = []
if confidence < 0.60:
safety_flags.append("low_confidence")
if not quality.get("is_acceptable", True):
safety_flags.append("image_quality")
if not reliability.get("is_reliable", True):
safety_flags.append("prediction_not_reliable")
return {
"success": True,
"language": result.get("language"),
"schema_version": "featurex-xray-v3-zerogpu",
"status": "success",
"analysis_id": result.get("analysis_id"),
"model": {
"classifier": "EfficientNet-B4",
"report_model": (result.get("model_info") or {}).get(
"report_model", "ZrH42/Janus-Pro-CXR-Final"
),
"architecture": (result.get("model_info") or {}).get("architecture"),
"version": (result.get("model_info") or {}).get("version"),
},
"classification": {
"predicted_class": pred.get("primary_diagnosis"),
"confidence": confidence,
"probabilities": {
x.get("disease"): float(x.get("probability", 0.0))
for x in result.get("predictions", [])
if x.get("disease")
},
"top_predictions": result.get("predictions", []),
},
"predictions": result.get("predictions", []),
"medical_report": selected,
"medical_report_ar": reports_ar,
"medical_report_en": reports_en,
"quality_metrics": quality,
"heatmap_stats": result.get("heatmap_stats", {}),
"model_info": result.get("model_info", {}),
"visualizations": result.get("visualizations", {}),
"report": {
"language": result.get("language"),
"source": selected.get("report_source", "EFFICIENTNET_TEMPLATE_FALLBACK"),
"model": selected.get("report_model", "ZrH42/Janus-Pro-CXR-Final"),
"janus_raw_text": selected.get("janus_raw_text", ""),
"findings": [selected.get("findings")] if selected.get("findings") else [],
"impression": [selected.get("impression")] if selected.get("impression") else [],
"recommendations": selected.get("recommendations", []),
"raw_text": "\n".join(
str(x)
for x in [selected.get("title"), selected.get("findings"), selected.get("impression")]
if x
),
"structured": selected,
"report_ar": reports_ar,
"report_en": reports_en,
},
"quality": quality,
"gradcam": {
"enabled": True,
"heatmap_stats": result.get("heatmap_stats", {}),
"visualizations": result.get("visualizations", {}),
"visualization_keys": list((result.get("visualizations") or {}).keys()),
},
"discrepancies": {"flags": [], "count": 0},
"safety": {
"wording_flags": confidence_flags,
"manual_review_required": bool(safety_flags),
"reasons": safety_flags,
},
"database": result.get("database", {}),
"storage": {
"original_image_url": (result.get("database") or {}).get("original_image_url"),
"visualization_urls": (result.get("database") or {}).get("visualization_urls", {}),
},
"janus": result.get("janus", {}),
"report_integrity": result.get("report_integrity", {}),
"legacy_result": result,
}
@spaces.GPU(duration=45)
def analyze(
image_base64: str,
language: str = "ar",
patient_name: str = "",
birth_date: str = "",
gender: str = "",
age: str = "",
user_id: str = "",
username: str = "",
department: str = "",
role: str = "doctor",
supervisor_id: str = "",
) -> str:
request_id = str(uuid.uuid4())
try:
raw = _decode_image(image_base64)
filename = "xray.png"
language = (language or "ar").strip().lower()
if language not in {"ar", "en"}:
language = "ar"
patient, user = _metadata(
patient_name,
birth_date,
gender,
age,
user_id,
username,
department,
role,
supervisor_id,
)
logger.info(
"ANALYZE START request_id=%s bytes=%d user_id=%s patient_name_has_arabic=%s",
request_id,
len(raw),
user.get("user_id"),
bool(re.search(r"[\u0600-\u06FF]", patient.get("patient_name", ""))),
)
result = predict_bytes(raw, filename, language)
if not isinstance(result, dict) or result.get("status") != "success":
raise RuntimeError("Inference did not return success.")
result["analysis_id"] = str(uuid.uuid4())
result["request_id"] = request_id
result["server_version"] = SERVER_VERSION
db_info = {"saved": False, "analysis_id": result["analysis_id"]}
try:
saved = database.save_analysis(
result=result,
original_bytes=raw,
original_filename=filename,
patient=patient,
user=user,
)
db_info.update(
{
"saved": bool(saved.get("saved", False)),
"analysis_id": saved.get("analysis_id", result["analysis_id"]),
"original_image_url": saved.get("original_image_url"),
"visualization_urls": saved.get("visualization_urls", {}),
"storage_saved": bool(
saved.get("storage_saved", saved.get("original_image_url"))
),
}
)
if saved.get("reason"):
db_info["reason"] = saved["reason"]
if saved.get("error"):
db_info["error"] = saved["error"]
db_info["error_type"] = saved.get("error_type")
except Exception as exc:
logger.exception("SUPABASE SAVE FAILED request_id=%s", request_id)
db_info.update(
{
"saved": False,
"storage_saved": False,
"error": str(exc),
"error_type": type(exc).__name__,
}
)
result["database"] = db_info
return _json(_build_unified(result))
except ValueError as exc:
logger.exception("ANALYZE REJECTED request_id=%s", request_id)
return _json(
{
"success": False,
"status": "rejected",
"request_id": request_id,
"error": "Image validation failed",
"message": str(exc),
}
)
except Exception as exc:
logger.exception("ANALYZE ERROR request_id=%s", request_id)
return _json(
{
"success": False,
"status": "error",
"request_id": request_id,
"error": "Internal server error",
"error_type": type(exc).__name__,
"message": str(exc),
"traceback": traceback.format_exc(limit=8),
}
)
@spaces.GPU(duration=30)
def gpu_health() -> str:
info = health_info()
info.update(
{
"server_version": SERVER_VERSION,
"transport": "Gradio Server + ZeroGPU",
"database": database.config_info(),
"api": {"analyze": True, "health": True, "classes": True},
}
)
return _json(info)
def health() -> str:
info = root_info()
info.update(
{
"server_version": SERVER_VERSION,
"transport": "Gradio Server + ZeroGPU",
"zero_gpu": True,
"database": database.config_info(),
}
)
return _json(info)
def classes() -> str:
return _json(classes_info())
# Gradio Server mode provides the HTTP API while keeping the Space a genuine
# Gradio Space, which is required for free ZeroGPU hosting.
server = gr.Server()
server.api(name="health", show_api=True)(health)
server.api(name="gpu_health", show_api=True)(gpu_health)
server.api(name="classes", show_api=True)(classes)
server.api(name="analyze", concurrency_limit=1, show_api=True)(analyze)
if __name__ == "__main__":
server.launch()