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()