"""
services/fraud_service.py — FraudDetectionService
Passive background monitoring: tab switches, copy-paste, camera,
multiple faces, inactivity. Runs asynchronously per-candidate.

Face violation rules:
  1st offence → warning
  2nd offence → warning
  3rd offence → disqualify + log
"""
from __future__ import annotations
import asyncio
from config import settings
from utils import json_db as db


class FraudService:

    def __init__(self) -> None:
        # In-memory counters per candidate (candidate_id → counts)
        self._tab_switches:  dict[str, int] = {}
        self._paste_events:  dict[str, int] = {}
        self._face_violations: dict[str, int] = {}
        self._inactivity:    dict[str, int] = {}
        self._total_violations: dict[str, int] = {}
        self._disqualified:  set[str]       = set()

    # ── Public API ─────────────────────────────────────────────────────────

    async def log_event(
        self,
        candidate_id: str,
        event_type: str,
        message: str,
        timestamp: str,
        round_num: int = 1,
        metadata: dict | None = None,
    ) -> dict:
        """
        Log a fraud event from the browser FraudDetector class.
        Returns action to take: none / warn / disqualify.
        """
        if candidate_id in self._disqualified:
            return self._make_response(
                candidate_id, 0, "disqualify",
                "Candidate has already been disqualified."
            )

        # Dispatch to specific counter
        action = await self._handle_event(candidate_id, event_type)

        # Persist event to fraud_log.json
        db.insert(settings.FRAUD_LOG_FILE, {
            "candidate_id": candidate_id,
            "event_type":   event_type,
            "message":      message,
            "timestamp":    timestamp,
            "round":        round_num,
            "action":       action,
            "metadata":     metadata or {},
        })

        # Total violations for this candidate
        total = self._total_violations.get(candidate_id, 0)

        # Warning level
        if action == "disqualify":
            level = "critical"
        elif total >= 2:
            level = "warning"
        else:
            level = "ok"

        msg_map = {
            "none":        "Activity logged.",
            "warn":        f"Warning {total}: {message}",
            "disqualify":  "⚠ You have been disqualified due to repeated violations.",
        }

        return {
            "logged":          True,
            "violation_count": total,
            "warning_level":   level,
            "action":          action,
            "message":         msg_map.get(action, "Logged."),
        }

    def is_disqualified(self, candidate_id: str) -> bool:
        return candidate_id in self._disqualified

    def get_report(self, candidate_id: str) -> dict:
        return {
            "candidate_id":     candidate_id,
            "tab_switches":     self._tab_switches.get(candidate_id, 0),
            "paste_events":     self._paste_events.get(candidate_id, 0),
            "face_violations":  self._face_violations.get(candidate_id, 0),
            "total_violations": self._total_violations.get(candidate_id, 0),
            "disqualified":     candidate_id in self._disqualified,
            "risk_score":       self._calc_risk(candidate_id),
        }

    def get_all_events(self, candidate_id: str) -> list[dict]:
        return db.find_many(settings.FRAUD_LOG_FILE, candidate_id=candidate_id)

    # ── Private helpers ────────────────────────────────────────────────────

    async def _handle_event(self, cid: str, event_type: str) -> str:
        """Increment counter and determine action."""
        action = "none"

        if event_type == "TAB_SWITCH":
            self._tab_switches[cid] = self._tab_switches.get(cid, 0) + 1
            if self._tab_switches[cid] >= settings.FRAUD_MAX_TAB_SWITCHES:
                action = "disqualify"
                self._disqualify(cid)
            else:
                action = "warn"
                self._total_violations[cid] = self._total_violations.get(cid, 0) + 1

        elif event_type in ("PASTE", "COPY"):
            self._paste_events[cid] = self._paste_events.get(cid, 0) + 1
            if self._paste_events[cid] >= settings.FRAUD_MAX_PASTE_EVENTS:
                action = "disqualify"
                self._disqualify(cid)
            else:
                action = "warn"
                self._total_violations[cid] = self._total_violations.get(cid, 0) + 1

        elif event_type == "MULTIPLE_FACES":
            self._face_violations[cid] = self._face_violations.get(cid, 0) + 1
            count = self._face_violations[cid]
            if count >= settings.FRAUD_MAX_FACE_VIOLATIONS:
                action = "disqualify"
                self._disqualify(cid)
            else:
                # 1st → warn, 2nd → warn, 3rd → disqualify
                action = "warn"
                self._total_violations[cid] = self._total_violations.get(cid, 0) + 1

        elif event_type in ("NO_CAMERA", "CAMERA_LOST", "NO_MIC"):
            action = "warn"
            self._total_violations[cid] = self._total_violations.get(cid, 0) + 1

        elif event_type == "INACTIVITY":
            action = "warn"
            self._total_violations[cid] = self._total_violations.get(cid, 0) + 1

        elif event_type in ("WINDOW_BLUR", "SCREEN_CAPTURE"):
            # Log only — minor
            action = "none"

        else:
            action = "none"

        return action

    def _disqualify(self, cid: str) -> None:
        self._disqualified.add(cid)
        self._total_violations[cid] = self._total_violations.get(cid, 0) + 1
        db.insert(settings.FRAUD_LOG_FILE, {
            "candidate_id": cid,
            "event_type":   "DISQUALIFIED",
            "message":      "Candidate auto-disqualified due to max violations",
            "timestamp":    db.now_iso(),
            "round":        0,
            "action":       "disqualify",
            "metadata":     self.get_report(cid),
        })

    def _calc_risk(self, cid: str) -> str:
        total = self._total_violations.get(cid, 0)
        if cid in self._disqualified:
            return "CRITICAL"
        if total >= 5:
            return "HIGH"
        if total >= 2:
            return "MEDIUM"
        return "LOW"

    def _make_response(self, cid: str, total: int, action: str, msg: str) -> dict:
        return {
            "logged":          True,
            "violation_count": total,
            "warning_level":   "critical" if action == "disqualify" else "warning",
            "action":          action,
            "message":         msg,
        }


fraud_service = FraudService()
