from typing import List, Dict, Any, Optional
from dataclasses import dataclass, field
from core.model_pool import Verdict, GateDecision, GateType, get_config


@dataclass
class ConsensusResult:
    verdict: Verdict
    agreement_ratio: float
    vote_distribution: Dict[str, int]


class Verifier:
    """
    Multi-model Verification Layer.
    Implements 4-way (parallel) and Duplex (Cross-check) patterns.

    Gates:
    - tool_call: Allow/deny tool execution
    - memory_write: Allow/deny write to semantic memory/whiteboard
    - escalation: Allow/deny Hybrid → Powerhouse escalation
    """

    def __init__(self, roles: List[str]):
        self.roles = roles
        self.configs = [get_config(r) for r in roles]

    def verify_duplex(
        self, action: Dict[str, Any], context: str, gate: GateType
    ) -> GateDecision:
        """
        Duplex Verification: Cross-check between two perspectives.
        D1 = Logic/Syntax check
        D2 = Safety/Invariant check
        """
        v1 = self._check_perspective(
            action, context, "logic", self.roles[0] if self.roles else "duplex_d1"
        )
        v2 = self._check_perspective(
            action,
            context,
            "safety",
            self.roles[1] if len(self.roles) > 1 else "duplex_d2",
        )

        allowed = v1.decision == "approve" and v2.decision == "approve"

        if not allowed:
            rationale = v1.rationale_tags + v2.rationale_tags
            return GateDecision(
                gate=gate,
                allowed=False,
                verdict=Verdict(
                    decision="reject",
                    confidence=min(v1.confidence, v2.confidence),
                    rationale_tags=rationale,
                    source_role="duplex",
                ),
            )

        return GateDecision(
            gate=gate,
            allowed=True,
            verdict=Verdict(
                decision="approve",
                confidence=(v1.confidence + v2.confidence) / 2,
                rationale_tags=[],
                source_role="duplex",
            ),
        )

    def consensus_4way(
        self, action: Dict[str, Any], context: str, gate: GateType
    ) -> GateDecision:
        """
        4-way parallel consensus.
        Outputs: {approve|reject|revise} + confidence + rationale_tags
        Threshold: 3/4 agreement required.
        """
        verdicts: List[Verdict] = []

        for i, role in enumerate(self.roles[:4]):
            v = self._check_perspective(action, context, f"perspective_{i}", role)
            verdicts.append(v)

        vote_counts: Dict[str, int] = {}
        for v in verdicts:
            vote_counts[v.decision] = vote_counts.get(v.decision, 0) + 1

        leader = max(vote_counts, key=lambda k: vote_counts[k])
        agreement = vote_counts[leader] / len(verdicts) if verdicts else 0

        winning_verdicts = [v for v in verdicts if v.decision == leader]
        avg_confidence = (
            sum(v.confidence for v in winning_verdicts) / len(winning_verdicts)
            if winning_verdicts
            else 0
        )

        combined_tags: List[str] = []
        for v in winning_verdicts:
            combined_tags.extend(v.rationale_tags)

        if agreement >= 0.75 and leader == "approve":
            return GateDecision(
                gate=gate,
                allowed=True,
                verdict=Verdict(
                    decision="approve",
                    confidence=avg_confidence,
                    rationale_tags=list(set(combined_tags)),
                    source_role="consensus_4way",
                ),
            )

        if leader == "revise" and agreement >= 0.5:
            return GateDecision(
                gate=gate,
                allowed=False,
                verdict=Verdict(
                    decision="revise",
                    confidence=avg_confidence,
                    rationale_tags=list(set(combined_tags)),
                    source_role="consensus_4way",
                ),
                override_reason="Revision requested by consensus",
            )

        return GateDecision(
            gate=gate,
            allowed=False,
            verdict=Verdict(
                decision="reject",
                confidence=avg_confidence,
                rationale_tags=list(set(combined_tags)),
                source_role="consensus_4way",
            ),
        )

    def _check_perspective(
        self, action: Dict[str, Any], context: str, check_type: str, role: str
    ) -> Verdict:
        """
        Single-perspective check.
        check_type: logic | safety | perspective_N
        """
        logic_valid = self._validate_schema(action)
        safety_valid = self._check_invariants(action, context)

        if check_type == "logic":
            if logic_valid:
                return Verdict("approve", 0.9, ["schema_valid"], role)
            return Verdict("reject", 0.95, ["schema_invalid"], role)

        if check_type == "safety":
            if safety_valid:
                return Verdict("approve", 0.85, ["invariants_pass"], role)
            return Verdict("reject", 0.9, ["safety_violation"], role)

        if logic_valid and safety_valid:
            return Verdict("approve", 0.8, [], role)
        if not logic_valid:
            return Verdict("reject", 0.85, ["schema_invalid"], role)
        return Verdict("revise", 0.7, ["needs_review"], role)

    def _validate_schema(self, action: Dict) -> bool:
        if "type" not in action:
            return False
        if action["type"] == "tool_call":
            return "tool" in action and "args" in action
        if action["type"] == "memory_write":
            return "key" in action and "value" in action
        return True

    def _check_invariants(self, action: Dict, context: str) -> bool:
        """
        Hard invariants:
        - Never fabricate tool outputs
        - Never write to memory without provenance
        - Stop after N failed retries
        """
        if action.get("type") == "tool_call":
            if action.get("fabricated", False):
                return False

        if action.get("type") == "memory_write":
            if "provenance" not in action:
                return False

        return True
