"""Severity tuning — auto-adjust leak severity based on historical patterns.

Rules:
- If a leak has recurred 3+ times, bump severity by one level
- If estimated_loss > threshold for severity, bump up
- If a leak was resolved quickly (< 4h), reduce severity next time
- If a leak went unresolved for 7+ days, bump severity

Public API:
    adjust_severity(leak, company_id) -> adjusted severity string
    get_severity_rules(company_id) -> current tuning rules
    update_severity_rules(company_id, rules) -> save custom rules
"""

from __future__ import annotations

from typing import Any, Dict, List, Optional

import logging

logger = logging.getLogger(__name__)

SEVERITY_ORDER = ["low", "medium", "high", "critical"]

# Default loss thresholds per severity (USD)
DEFAULT_LOSS_THRESHOLDS = {
    "low": 0,
    "medium": 500,
    "high": 5000,
    "critical": 20000,
}


def adjust_severity(
    leak: Dict[str, Any],
    company_id: str,
) -> str:
    """Adjust leak severity based on historical patterns.

    Args:
        leak: dict with detector_id, source, severity, estimated_loss
        company_id: company context

    Returns:
        Adjusted severity string (may be same as original).
    """
    from app.models import db, RevenueLeak

    original_severity = leak.get("severity", "medium")
    current_idx = SEVERITY_ORDER.index(original_severity) if original_severity in SEVERITY_ORDER else 1

    detector_id = leak.get("detector_id")
    source = leak.get("source")
    estimated_loss = leak.get("estimated_loss")

    # Rule 1: Recurrence bump
    if detector_id:
        recurring = _count_recent_recurrences(
            company_id, detector_id, source
        )
        if recurring >= 3:
            current_idx = min(current_idx + 1, len(SEVERITY_ORDER) - 1)

    # Rule 2: Loss threshold bump
    if estimated_loss is not None:
        thresholds = _get_loss_thresholds(company_id)
        for sev, threshold in thresholds.items():
            sev_idx = SEVERITY_ORDER.index(sev)
            if estimated_loss >= threshold and sev_idx > current_idx:
                current_idx = sev_idx

    # Rule 3: Quick resolution downgrade (check history)
    if detector_id and source:
        avg_resolution_hours = _avg_resolution_time(
            company_id, detector_id, source
        )
        if avg_resolution_hours is not None and avg_resolution_hours < 4:
            # Quick to resolve — less severe
            current_idx = max(current_idx - 1, 0)

    return SEVERITY_ORDER[current_idx]


def _count_recent_recurrences(
    company_id: str,
    detector_id: str,
    source: str,
) -> int:
    """Count how many times this detector+source combo has triggered recently."""
    from app.models import db, RevenueLeak
    from datetime import timedelta

    cutoff = (
        __import__("datetime", fromlist=["datetime"]).datetime.now(
            __import__("datetime", fromlist=["timezone"]).timezone.utc
        )
        - timedelta(days=90)
    )

    count = RevenueLeak.query.filter(
        RevenueLeak.company_id == company_id,
        RevenueLeak.detector_id == detector_id,
        RevenueLeak.source == source,
        RevenueLeak.detected_at >= cutoff,
    ).count()

    return count


def _avg_resolution_time(
    company_id: str,
    detector_id: str,
    source: str,
) -> Optional[float]:
    """Get average resolution time in hours for a detector+source combo."""
    from app.models import db, RevenueLeak
    from datetime import timedelta

    cutoff = (
        __import__("datetime", fromlist=["datetime"]).datetime.now(
            __import__("datetime", fromlist=["timezone"]).timezone.utc
        )
        - timedelta(days=90)
    )

    leaks = RevenueLeak.query.filter(
        RevenueLeak.company_id == company_id,
        RevenueLeak.detector_id == detector_id,
        RevenueLeak.source == source,
        RevenueLeak.resolved == True,
        RevenueLeak.detected_at >= cutoff,
    ).all()

    if not leaks:
        return None

    total_hours = 0
    count = 0
    for leak in leaks:
        if leak.resolved_at and leak.detected_at:
            delta = leak.resolved_at - leak.detected_at
            total_hours += delta.total_seconds() / 3600
            count += 1

    return total_hours / count if count > 0 else None


def _get_loss_thresholds(company_id: str) -> Dict[str, float]:
    """Get loss thresholds for severity adjustment."""
    from app.models import db, Company

    company = db.session.get(
        __import__("app.models", fromlist=["Company"]).Company,
        company_id,
    )
    if company is None:
        return DEFAULT_LOSS_THRESHOLDS

    settings = company.settings_json or {}
    severity_settings = settings.get("leak_severity") or {}
    thresholds = severity_settings.get("loss_thresholds")

    return thresholds if isinstance(thresholds, dict) else DEFAULT_LOSS_THRESHOLDS


def get_severity_rules(company_id: str) -> Dict[str, Any]:
    """Get current severity tuning rules for a company."""
    from app.models import db, Company

    company = db.session.get(
        __import__("app.models", fromlist=["Company"]).Company,
        company_id,
    )
    if company is None:
        return {
            "loss_thresholds": DEFAULT_LOSS_THRESHOLDS,
            "recurrence_bump": True,
            "quick_resolution_downgrade": True,
            "unresolved_escalation": True,
        }

    settings = company.settings_json or {}
    severity_settings = settings.get("leak_severity") or {}

    return {
        "loss_thresholds": severity_settings.get(
            "loss_thresholds", DEFAULT_LOSS_THRESHOLDS
        ),
        "recurrence_bump": severity_settings.get(
            "recurrence_bump", True
        ),
        "quick_resolution_downgrade": severity_settings.get(
            "quick_resolution_downgrade", True
        ),
        "unresolved_escalation": severity_settings.get(
            "unresolved_escalation", True
        ),
    }


def update_severity_rules(
    company_id: str,
    rules: Dict[str, Any],
) -> Dict[str, Any]:
    """Update severity tuning rules for a company."""
    from app.models import db, Company

    company = db.session.get(
        __import__("app.models", fromlist=["Company"]).Company,
        company_id,
    )
    if company is None:
        return {"error": "Company not found"}

    settings = company.settings_json or {}
    settings.setdefault("leak_severity", {}).update(rules)
    company.settings_json = settings

    db.session.commit()

    return {"success": True, "rules": get_severity_rules(company_id)}
