import asyncio from motor.motor_asyncio import AsyncIOMotorClient from app.config import config from typing import Optional, List, Dict, Any from datetime import datetime class DatabaseManager: def __init__(self): self.client = None self.database = None self.collection = None self.is_connected = False async def initialize(self): """Initialize MongoDB connection if available""" try: if config.ENABLE_MONGODB and config.MONGODB_URL: self.client = AsyncIOMotorClient(config.MONGODB_URL) self.database = self.client[config.DATABASE_NAME] self.collection = self.database[config.COLLECTION_NAME] self.is_connected = True print("MongoDB configured successfully") else: print("MongoDB not configured - predictions will not be persisted") except asyncio.TimeoutError: print("MongoDB connection timeout - running without database") self.is_connected = False except Exception as e: print(f"MongoDB configuration failed: {e}") print("Running without database - predictions will not be persisted") self.is_connected = False async def save_prediction(self, prediction_data: Dict[str, Any]) -> bool: """Save prediction to database if available""" if not self.is_connected: return False try: await self.collection.insert_one({ **prediction_data, "_id": prediction_data["id"] }) return True except Exception as e: print(f"Failed to save prediction: {e}") return False async def get_prediction(self, prediction_id: str) -> Optional[Dict[str, Any]]: """Get prediction by ID""" if not self.is_connected: return None try: doc = await self.collection.find_one({"_id": prediction_id}) if doc: doc["id"] = doc.pop("_id") return doc except Exception as e: print(f"Failed to get prediction: {e}") return None async def get_predictions_paginated( self, page: int = 1, limit: int = 10, filters: Optional[Dict[str, Any]] = None, sort_field: str = "created_at", sort_direction: int = -1 ) -> tuple: """Get paginated predictions with filtering""" if not self.is_connected: return [], 0 try: filter_query = filters or {} skip = (page - 1) * limit # Get total count total_count = await self.collection.count_documents(filter_query) # Get paginated results cursor = self.collection.find(filter_query).sort(sort_field, sort_direction).skip(skip).limit(limit) predictions = [] async for doc in cursor: doc_dict = dict(doc) doc_dict["id"] = doc_dict.pop("_id") predictions.append(doc_dict) return predictions, total_count except Exception as e: print(f"Failed to get predictions: {e}") return [], 0 async def delete_prediction(self, prediction_id: str) -> bool: """Delete prediction by ID""" if not self.is_connected: return False try: result = await self.collection.delete_one({"_id": prediction_id}) return result.deleted_count > 0 except Exception as e: print(f"Failed to delete prediction: {e}") return False async def get_statistics(self, filter_query: Optional[Dict[str, Any]] = None, include_enhanced: bool = True) -> Dict[str, Any]: """ Get prediction statistics with optional enhanced metrics Args: filter_query: MongoDB filter for date ranges, etc. include_enhanced: Whether to include enhanced metrics (trends, analytics, performance, health) Returns: Dictionary with statistics data """ if not self.is_connected: return { "total_predictions": 0, "stroke_detected": 0, "normal_cases": 0, "stroke_detection_rate": 0, "average_confidence": 0 } try: filter_query = filter_query or {} # Basic statistics total_predictions = await self.collection.count_documents(filter_query) stroke_detected = await self.collection.count_documents({**filter_query, "has_stroke": True}) # Get average confidence pipeline = [ {"$match": filter_query}, {"$group": {"_id": None, "avg_confidence": {"$avg": "$confidence_score"}}} ] avg_confidence_result = await self.collection.aggregate(pipeline).to_list(1) avg_confidence = avg_confidence_result[0]["avg_confidence"] if avg_confidence_result else 0 basic_stats = { "total_predictions": total_predictions, "stroke_detected": stroke_detected, "normal_cases": total_predictions - stroke_detected, "stroke_detection_rate": (stroke_detected / total_predictions * 100) if total_predictions > 0 else 0, "average_confidence": avg_confidence if avg_confidence else 0 } if not include_enhanced: return basic_stats # Enhanced statistics trends = await self._calculate_trends(filter_query, basic_stats) detailed_analytics = await self._calculate_detailed_analytics(filter_query) performance_metrics = await self._calculate_performance_metrics(filter_query) stroke_type_breakdown = await self._calculate_stroke_type_breakdown(filter_query) stroke_type_accuracy = await self._calculate_stroke_type_accuracy(filter_query) return { **basic_stats, "stroke_type_breakdown": stroke_type_breakdown, "stroke_type_accuracy": stroke_type_accuracy, "trends": trends, "detailed_analytics": detailed_analytics, "performance_metrics": performance_metrics } except Exception as e: print(f"Failed to get statistics: {e}") return { "total_predictions": 0, "stroke_detected": 0, "normal_cases": 0, "stroke_detection_rate": 0, "average_confidence": 0 } async def _calculate_trends(self, filter_query: Dict[str, Any], current_stats: Dict[str, Any]) -> Dict[str, float]: """Calculate trends by comparing current period to previous period""" try: # Extract date range from filter if "created_at" not in filter_query: # No date filter, can't calculate trends reliably return { "total_predictions": 0.0, "stroke_detected": 0.0, "normal_cases": 0.0, "detection_rate": 0.0, "average_confidence": 0.0 } current_date_filter = filter_query["created_at"] if "$gte" not in current_date_filter: return { "total_predictions": 0.0, "stroke_detected": 0.0, "normal_cases": 0.0, "detection_rate": 0.0, "average_confidence": 0.0 } # Calculate previous period start_date = current_date_filter["$gte"] end_date = current_date_filter.get("$lte", datetime.now()) period_length = end_date - start_date previous_start = start_date - period_length previous_end = start_date # Get previous period stats previous_filter = { **{k: v for k, v in filter_query.items() if k != "created_at"}, "created_at": { "$gte": previous_start, "$lte": previous_end } } prev_total = await self.collection.count_documents(previous_filter) prev_stroke = await self.collection.count_documents({**previous_filter, "has_stroke": True}) pipeline = [ {"$match": previous_filter}, {"$group": {"_id": None, "avg_confidence": {"$avg": "$confidence_score"}}} ] prev_conf_result = await self.collection.aggregate(pipeline).to_list(1) prev_confidence = prev_conf_result[0]["avg_confidence"] if prev_conf_result else 0 # Calculate percentage changes def calc_change(current, previous): if previous == 0: return 100.0 if current > 0 else 0.0 return ((current - previous) / previous) * 100 return { "total_predictions": round(calc_change(current_stats["total_predictions"], prev_total), 1), "stroke_detected": round(calc_change(current_stats["stroke_detected"], prev_stroke), 1), "normal_cases": round(calc_change(current_stats["normal_cases"], prev_total - prev_stroke), 1), "detection_rate": round(calc_change(current_stats["stroke_detection_rate"], (prev_stroke / prev_total * 100) if prev_total > 0 else 0), 1), "average_confidence": round(calc_change(current_stats["average_confidence"], prev_confidence if prev_confidence else 0), 1) } except Exception as e: print(f"Failed to calculate trends: {e}") return { "total_predictions": 0.0, "stroke_detected": 0.0, "normal_cases": 0.0, "detection_rate": 0.0, "average_confidence": 0.0 } async def _calculate_detailed_analytics(self, filter_query: Dict[str, Any]) -> Dict[str, float]: """Calculate detailed analytics metrics using real ground truth when available""" try: # Get processing time statistics pipeline = [ {"$match": filter_query}, {"$group": { "_id": None, "avg_processing_time": {"$avg": "$processing_time"}, "total": {"$sum": 1}, "stroke_count": { "$sum": {"$cond": ["$has_stroke", 1, 0]} } }} ] result = await self.collection.aggregate(pipeline).to_list(1) if not result: return self._default_detailed_analytics() data = result[0] avg_processing_time = data.get("avg_processing_time", 0) # Check if we have ground truth data ground_truth_filter = {**filter_query, "doctor_expectation": {"$ne": None}} ground_truth_count = await self.collection.count_documents(ground_truth_filter) if ground_truth_count > 0: # Calculate real FP/FN rates using doctor expectations confusion_pipeline = [ {"$match": ground_truth_filter}, {"$group": { "_id": { "predicted": "$has_stroke", "actual": "$doctor_expectation" }, "count": {"$sum": 1} }} ] confusion_result = await self.collection.aggregate(confusion_pipeline).to_list(10) tp = tn = fp = fn = 0 for item in confusion_result: predicted = item["_id"]["predicted"] actual = item["_id"]["actual"] count = item["count"] if actual == True and predicted == True: tp = count elif actual == False and predicted == False: tn = count elif actual == False and predicted == True: fp = count # False Positive elif actual == True and predicted == False: fn = count # False Negative total_with_gt = tp + tn + fp + fn fp_rate = (fp / total_with_gt * 100) if total_with_gt > 0 else 0 fn_rate = (fn / total_with_gt * 100) if total_with_gt > 0 else 0 detection_accuracy = ((tp + tn) / total_with_gt * 100) if total_with_gt > 0 else 0 else: # Estimate FP/FN rates based on confidence distribution low_confidence_threshold = 0.6 low_conf_pipeline = [ {"$match": {**filter_query, "confidence_score": {"$lt": low_confidence_threshold}}}, {"$count": "count"} ] low_conf_result = await self.collection.aggregate(low_conf_pipeline).to_list(1) low_conf_count = low_conf_result[0]["count"] if low_conf_result else 0 total = data.get("total", 1) # Estimate: low confidence predictions are potential false positives/negatives fp_rate = (low_conf_count / total * 100) if total > 0 else 0 fn_rate = fp_rate * 0.8 # Assume FN is slightly lower than FP detection_accuracy = (data.get("stroke_count", 0) / total * 100) if total > 0 else 0 return { "detection_accuracy": round(min(detection_accuracy, 100.0), 1), "detection_accuracy_trend": 0.0, # Would need historical data "average_processing_time": round(avg_processing_time, 1), "processing_time_trend": 0.0, # Would need historical data "false_positive_rate": round(min(fp_rate, 100.0), 1), "false_positive_trend": 0.0, "false_negative_rate": round(min(fn_rate, 100.0), 1), "false_negative_trend": 0.0 } except Exception as e: print(f"Failed to calculate detailed analytics: {e}") return self._default_detailed_analytics() def _default_detailed_analytics(self) -> Dict[str, float]: """Return default detailed analytics""" return { "detection_accuracy": 0.0, "detection_accuracy_trend": 0.0, "average_processing_time": 0.0, "processing_time_trend": 0.0, "false_positive_rate": 0.0, "false_positive_trend": 0.0, "false_negative_rate": 0.0, "false_negative_trend": 0.0 } async def _calculate_performance_metrics(self, filter_query: Dict[str, Any]) -> Dict[str, float]: """ Calculate ML model performance metrics using real doctor expectations when available. Falls back to estimates based on confidence scores if no ground truth. """ try: # First, check if we have any cases with doctor expectations ground_truth_filter = {**filter_query, "doctor_expectation": {"$ne": None}} ground_truth_count = await self.collection.count_documents(ground_truth_filter) if ground_truth_count > 0: # Use real ground truth data return await self._calculate_real_performance_metrics(ground_truth_filter) else: # Fall back to estimates return await self._calculate_estimated_performance_metrics(filter_query) except Exception as e: print(f"Failed to calculate performance metrics: {e}") return self._default_performance_metrics() async def _calculate_real_performance_metrics(self, filter_query: Dict[str, Any]) -> Dict[str, float]: """Calculate real performance metrics using doctor expectations as ground truth""" try: pipeline = [ {"$match": filter_query}, {"$facet": { "confusion_matrix": [ {"$group": { "_id": { "predicted": "$has_stroke", "actual": "$doctor_expectation" }, "count": {"$sum": 1} }} ], "avg_confidence": [ {"$group": {"_id": None, "avg": {"$avg": "$confidence_score"}}} ] }} ] result = await self.collection.aggregate(pipeline).to_list(1) if not result or not result[0]["confusion_matrix"]: return self._default_performance_metrics() # Parse confusion matrix tp = tn = fp = fn = 0 for item in result[0]["confusion_matrix"]: predicted = item["_id"]["predicted"] actual = item["_id"]["actual"] count = item["count"] if actual == True and predicted == True: tp = count # True Positive elif actual == False and predicted == False: tn = count # True Negative elif actual == False and predicted == True: fp = count # False Positive elif actual == True and predicted == False: fn = count # False Negative # Calculate metrics sensitivity = (tp / (tp + fn)) if (tp + fn) > 0 else 0 # True Positive Rate specificity = (tn / (tn + fp)) if (tn + fp) > 0 else 0 # True Negative Rate ppv = (tp / (tp + fp)) if (tp + fp) > 0 else 0 # Positive Predictive Value / Precision npv = (tn / (tn + fn)) if (tn + fn) > 0 else 0 # Negative Predictive Value # F1 Score precision = ppv recall = sensitivity f1 = (2 * precision * recall / (precision + recall)) if (precision + recall) > 0 else 0 # AUC-ROC estimate (would need ROC curve calculation for exact value) avg_conf = result[0]["avg_confidence"][0]["avg"] if result[0]["avg_confidence"] else 0.5 auc_roc = min(0.5 + (avg_conf * 0.5), 0.99) return { "sensitivity": round(sensitivity * 100, 1), "specificity": round(specificity * 100, 1), "positive_predictive_value": round(ppv * 100, 1), "negative_predictive_value": round(npv * 100, 1), "f1_score": round(f1 * 100, 1), "auc_roc": round(auc_roc, 3) } except Exception as e: print(f"Failed to calculate real performance metrics: {e}") return self._default_performance_metrics() async def _calculate_estimated_performance_metrics(self, filter_query: Dict[str, Any]) -> Dict[str, float]: """Calculate estimated performance metrics based on confidence scores (fallback method)""" try: pipeline = [ {"$match": filter_query}, {"$facet": { "total": [{"$count": "count"}], "stroke_predictions": [ {"$match": {"has_stroke": True}}, {"$count": "count"} ], "high_confidence_stroke": [ {"$match": {"has_stroke": True, "confidence_score": {"$gte": 0.85}}}, {"$count": "count"} ], "high_confidence_normal": [ {"$match": {"has_stroke": False, "confidence_score": {"$gte": 0.85}}}, {"$count": "count"} ], "avg_confidence": [ {"$group": {"_id": None, "avg": {"$avg": "$confidence_score"}}} ] }} ] result = await self.collection.aggregate(pipeline).to_list(1) if not result: return self._default_performance_metrics() data = result[0] total = data["total"][0]["count"] if data["total"] else 0 if total == 0: return self._default_performance_metrics() stroke_pred = data["stroke_predictions"][0]["count"] if data["stroke_predictions"] else 0 normal_pred = total - stroke_pred high_conf_stroke = data["high_confidence_stroke"][0]["count"] if data["high_confidence_stroke"] else 0 high_conf_normal = data["high_confidence_normal"][0]["count"] if data["high_confidence_normal"] else 0 avg_conf = data["avg_confidence"][0]["avg"] if data["avg_confidence"] else 0.5 # Estimate performance metrics estimated_tp = high_conf_stroke estimated_tn = high_conf_normal estimated_fp = stroke_pred - high_conf_stroke estimated_fn = normal_pred - high_conf_normal # Ensure non-negative values estimated_fp = max(0, estimated_fp) estimated_fn = max(0, estimated_fn) # Calculate metrics sensitivity = (estimated_tp / (estimated_tp + estimated_fn)) if (estimated_tp + estimated_fn) > 0 else 0 specificity = (estimated_tn / (estimated_tn + estimated_fp)) if (estimated_tn + estimated_fp) > 0 else 0 ppv = (estimated_tp / (estimated_tp + estimated_fp)) if (estimated_tp + estimated_fp) > 0 else 0 npv = (estimated_tn / (estimated_tn + estimated_fn)) if (estimated_tn + estimated_fn) > 0 else 0 # F1 Score precision = ppv recall = sensitivity f1 = (2 * precision * recall / (precision + recall)) if (precision + recall) > 0 else 0 # AUC-ROC estimate auc_roc = min(0.5 + (avg_conf * 0.5), 0.99) return { "sensitivity": round(min(sensitivity * 100, 100.0), 1), "specificity": round(min(specificity * 100, 100.0), 1), "positive_predictive_value": round(min(ppv * 100, 100.0), 1), "negative_predictive_value": round(min(npv * 100, 100.0), 1), "f1_score": round(min(f1 * 100, 100.0), 1), "auc_roc": round(auc_roc, 3) } except Exception as e: print(f"Failed to calculate estimated performance metrics: {e}") return self._default_performance_metrics() async def _calculate_stroke_type_breakdown(self, filter_query: Dict[str, Any]) -> Dict[str, Any]: """Calculate breakdown of cases by doctor-classified stroke type""" try: pipeline = [ {"$match": filter_query}, {"$group": { "_id": "$stroke_type", "count": {"$sum": 1} }} ] results = await self.collection.aggregate(pipeline).to_list(10) normal_count = 0 ischemia_count = 0 bleeding_count = 0 unclassified_count = 0 for item in results: stroke_type = item["_id"] count = item["count"] if stroke_type == "normal": normal_count = count elif stroke_type == "ischemia": ischemia_count = count elif stroke_type == "bleeding": bleeding_count = count else: unclassified_count += count classified_total = normal_count + ischemia_count + bleeding_count return { "normal_count": normal_count, "ischemia_count": ischemia_count, "bleeding_count": bleeding_count, "unclassified_count": unclassified_count, "normal_percentage": round((normal_count / classified_total * 100) if classified_total > 0 else 0, 1), "ischemia_percentage": round((ischemia_count / classified_total * 100) if classified_total > 0 else 0, 1), "bleeding_percentage": round((bleeding_count / classified_total * 100) if classified_total > 0 else 0, 1), } except Exception as e: print(f"Failed to calculate stroke type breakdown: {e}") return { "normal_count": 0, "ischemia_count": 0, "bleeding_count": 0, "unclassified_count": 0, "normal_percentage": 0.0, "ischemia_percentage": 0.0, "bleeding_percentage": 0.0 } async def _calculate_stroke_type_accuracy(self, filter_query: Dict[str, Any]) -> Dict[str, float]: """Calculate AI accuracy for each stroke type category using doctor classification as ground truth""" try: # Only consider cases where doctor provided a stroke_type typed_filter = {**filter_query, "stroke_type": {"$ne": None}} typed_count = await self.collection.count_documents(typed_filter) if typed_count == 0: return {"normal_accuracy": 0.0, "ischemia_accuracy": 0.0, "bleeding_accuracy": 0.0} # For each type, check how often AI agreed with the doctor accuracies = {} for stype in ["normal", "ischemia", "bleeding"]: type_filter = {**filter_query, "stroke_type": stype} type_total = await self.collection.count_documents(type_filter) if type_total == 0: accuracies[f"{stype}_accuracy"] = 0.0 continue if stype == "normal": # Doctor says normal → AI should predict no stroke (has_stroke=False) correct_count = await self.collection.count_documents({**type_filter, "has_stroke": False}) else: # Doctor says ischemia/bleeding → AI should predict stroke (has_stroke=True) correct_count = await self.collection.count_documents({**type_filter, "has_stroke": True}) accuracies[f"{stype}_accuracy"] = round((correct_count / type_total * 100), 1) return accuracies except Exception as e: print(f"Failed to calculate stroke type accuracy: {e}") return {"normal_accuracy": 0.0, "ischemia_accuracy": 0.0, "bleeding_accuracy": 0.0} def _default_performance_metrics(self) -> Dict[str, float]: """Return default performance metrics""" return { "sensitivity": 0.0, "specificity": 0.0, "positive_predictive_value": 0.0, "negative_predictive_value": 0.0, "f1_score": 0.0, "auc_roc": 0.5 } def is_available(self) -> bool: """Check if database is available""" return self.is_connected # Global database manager instance db_manager = DatabaseManager()