﻿"""CTCAE Grading Engine for HAD Digital MVP.

Loads grading rules from data/ctcae_rules.json.
Maps structured symptom inputs to CTCAE grades 1-5.
Grades >= 2 are flagged provisional until clinician confirmed.
"""

import json
from pathlib import Path
from database import get_db


class CTCAEEngine:
    """CTCAE v5.0 grading engine with configurable rules."""

    def __init__(self, rules_path: str | None = None):
        if rules_path is None:
            from config_manager import config
            rules_path = config.get("ctcae", "rules_path", default="data/ctcae_rules.json")
        self.rules_path = Path(rules_path)
        self.rules = {}
        self.load_rules()

    def load_rules(self):
        """Load CTCAE rules from JSON file."""
        if not self.rules_path.exists():
            raise FileNotFoundError(f"CTCAE rules file not found: {self.rules_path}")
        with open(self.rules_path, "r", encoding="utf-8") as f:
            data = json.load(f)
        self.rules = data.get("categories", {})
        self.version = data.get("version", "unknown")

    def get_symptom_rule(self, symptom_id: str) -> dict | None:
        """Find a symptom rule by ID across all categories."""
        for category_name, symptoms in self.rules.items():
            if symptom_id in symptoms:
                return symptoms[symptom_id]
        return None

    def list_symptoms(self) -> list[dict]:
        """List all available symptoms with their categories."""
        result = []
        for category_name, symptoms in self.rules.items():
            for symptom_id, rule in symptoms.items():
                result.append({
                    "symptom_id": symptom_id,
                    "display_name": rule.get("display_name", symptom_id),
                    "category": category_name,
                })
        return result

    def grade_symptom(self, symptom_id: str, inputs: dict) -> dict:
        """Grade a symptom based on structured inputs.
        
        Args:
            symptom_id: The symptom identifier (e.g., 'nausea', 'neutropenia')
            inputs: Dict of measurable values (e.g., severity_score, episodes_per_24h,
                    anc_mm3, hemoglobin_gdl, temp_celsius, bsa_percent, etc.)
        
        Returns:
            Dict with grade, criteria, intervention, provisional flag, and matched input.
            Returns None values for unmatched grades.
        """
        rule = self.get_symptom_rule(symptom_id)
        if rule is None:
            return {
                "symptom_id": symptom_id,
                "grade": None,
                "criteria": f"Unknown symptom: {symptom_id}",
                "intervention": None,
                "provisional": True,
                "error": f"No rule found for symptom '{symptom_id}'",
            }

        grades = rule.get("grades", {})
        best_grade = None
        best_grade_num = 0
        best_threshold_key = None
        best_threshold_value = None

        # Find the highest matching grade
        for grade_str in sorted(grades.keys(), key=int, reverse=True):
            grade_num = int(grade_str)
            grade_info = grades[grade_str]
            thresholds = grade_info.get("thresholds", {})

            # Check if any threshold matches
            matched = False
            matched_key = None
            matched_value = None

            for threshold_key, (low, high) in thresholds.items():
                input_value = inputs.get(threshold_key)
                if input_value is not None:
                    try:
                        val = float(input_value)
                        if low <= val <= high:
                            matched = True
                            matched_key = threshold_key
                            matched_value = val
                            break
                    except (ValueError, TypeError):
                        continue

            if matched and grade_num > best_grade_num:
                best_grade = grade_info
                best_grade_num = grade_num
                best_threshold_key = matched_key
                best_threshold_value = matched_value

        if best_grade is None:
            return {
                "symptom_id": symptom_id,
                "display_name": rule.get("display_name", symptom_id),
                "category": rule.get("category", "unknown"),
                "grade": None,
                "criteria": "No matching grade found for provided inputs",
                "intervention": None,
                "provisional": True,
                "error": "Inputs did not match any grade threshold",
            }

        # Grade >= 2 is provisional until clinician confirmed
        provisional = best_grade_num >= 2

        return {
            "symptom_id": symptom_id,
            "display_name": rule.get("display_name", symptom_id),
            "category": rule.get("category", "unknown"),
            "grade": best_grade_num,
            "criteria": best_grade.get("criteria", ""),
            "intervention": best_grade.get("intervention", ""),
            "provisional": provisional,
            "matched_threshold": best_threshold_key,
            "matched_value": best_threshold_value,
            "engine_version": self.version,
        }

    def grade_and_save(self, report_id: int, patient_id: int, symptom_id: str, inputs: dict) -> dict:
        """Grade a symptom and save the result to the database.
        
        Args:
            report_id: The toxicity report ID
            patient_id: The patient ID
            symptom_id: The symptom identifier
            inputs: The measurable inputs
        
        Returns:
            The grading result dict with the database grade_id.
        """
        result = self.grade_symptom(symptom_id, inputs)

        if result.get("grade") is None:
            return result

        db = get_db()
        cursor = db.execute(
            """INSERT INTO toxicity_grades 
               (report_id, patient_id, symptom_id, grade, criteria, intervention, 
                provisional, grading_engine_version)
               VALUES (?, ?, ?, ?, ?, ?, ?, ?)""",
            (
                report_id,
                patient_id,
                symptom_id,
                result["grade"],
                result["criteria"],
                result.get("intervention", ""),
                1 if result["provisional"] else 0,
                result.get("engine_version", "unknown"),
            ),
        )
        db.commit()
        result["grade_id"] = cursor.lastrowid

        # Update report status
        db.execute(
            "UPDATE toxicity_reports SET status = 'graded', updated_at = datetime('now') WHERE id = ?",
            (report_id,),
        )
        db.commit()

        return result

    def confirm_grade(self, grade_id: int, confirmed_by: int) -> bool:
        """Clinician confirms a provisional grade."""
        db = get_db()
        db.execute(
            """UPDATE toxicity_grades 
               SET provisional = 0, confirmed_by = ?, confirmed_at = datetime('now'),
                   updated_at = datetime('now')
               WHERE id = ?""",
            (confirmed_by, grade_id),
        )
        db.commit()
        return True


# Global singleton
_engine = None


def get_ctcae_engine(rules_path: str | None = None) -> CTCAEEngine:
    """Get the global CTCAE engine instance."""
    global _engine
    if _engine is None:
        _engine = CTCAEEngine(rules_path)
    return _engine
