#!/usr/bin/env python3
"""Dependency-free, deterministic RAG/fine-tuning selection reference."""

from __future__ import annotations

import argparse
import json
import math
from pathlib import Path


REQUIRED_FACTS = (
    "dataRightsVerified",
    "evaluationContractReady",
    "baselineMeetsGate",
    "freshKnowledgeRequired",
    "sourceTraceabilityRequired",
    "permissionAwareRetrievalRequired",
    "stableBehaviorGap",
    "sufficientLabeledData",
    "ragIsolatedGainVerified",
    "fineTuningIsolatedGainVerified",
    "ragAloneMeetsGate",
    "fineTuningAloneMeetsGate",
)

COMPARISON_CONTEXT = (
    "bothCandidatesPassedLockedGate",
    "sameLockedEvaluationSet",
    "sameProductionLoadProfile",
    "sameCostHorizon",
    "sameOperationalAndGovernanceRubrics",
)

COMPARISON_METRICS = {
    "qualityMargin": "higher",
    "p95LatencyMs": "lower",
    "monthlyCostCny": "lower",
    "operationalComplexityScore": "lower",
    "governanceRiskScore": "lower",
}

COMPARISON_CANDIDATES = ("rag", "fineTuning")


def validate_facts(facts: dict) -> None:
    missing = [name for name in REQUIRED_FACTS if name not in facts]
    unknown = sorted(set(facts) - set(REQUIRED_FACTS))
    non_boolean = [name for name in REQUIRED_FACTS if name in facts and type(facts[name]) is not bool]
    contradictions = []
    if not missing and not non_boolean:
        if facts["ragAloneMeetsGate"] and not facts["ragIsolatedGainVerified"]:
            contradictions.append("ragAloneMeetsGate_requires_ragIsolatedGainVerified")
        if facts["fineTuningAloneMeetsGate"] and not facts["fineTuningIsolatedGainVerified"]:
            contradictions.append(
                "fineTuningAloneMeetsGate_requires_fineTuningIsolatedGainVerified"
            )
    if missing or unknown or non_boolean or contradictions:
        raise ValueError(
            json.dumps(
                {
                    "missing": missing,
                    "unknown": unknown,
                    "nonBoolean": non_boolean,
                    "contradictions": contradictions,
                },
                ensure_ascii=False,
                sort_keys=True,
            )
        )


def recommend(facts: dict) -> dict:
    validate_facts(facts)
    knowledge_gap = any(
        facts[name]
        for name in (
            "freshKnowledgeRequired",
            "sourceTraceabilityRequired",
            "permissionAwareRetrievalRequired",
        )
    )
    behavior_gap = facts["stableBehaviorGap"]

    if not facts["dataRightsVerified"] or not facts["evaluationContractReady"]:
        rule_id = "R01_PREPARE_DATA_AND_EVAL"
        action = "prepare_data_before_selection"
    elif facts["baselineMeetsGate"]:
        rule_id = "R02_KEEP_PASSING_BASELINE"
        action = "continue_baseline"
    elif (
        knowledge_gap
        and behavior_gap
        and facts["ragIsolatedGainVerified"]
        and facts["fineTuningIsolatedGainVerified"]
        and not facts["ragAloneMeetsGate"]
        and not facts["fineTuningAloneMeetsGate"]
    ):
        rule_id = "R03_VALIDATE_HYBRID_AFTER_ABLATION"
        action = "validate_hybrid_after_ablation"
    elif knowledge_gap and not facts["ragIsolatedGainVerified"]:
        rule_id = "R04_VALIDATE_RAG_FIRST"
        action = "validate_rag"
    elif behavior_gap and not facts["sufficientLabeledData"]:
        rule_id = "R05_PREPARE_LABELED_DATA"
        action = "prepare_data_before_selection"
    elif behavior_gap and not facts["fineTuningIsolatedGainVerified"]:
        rule_id = "R06_VALIDATE_FINE_TUNING"
        action = "validate_fine_tuning"
    elif (
        knowledge_gap
        and behavior_gap
        and facts["ragAloneMeetsGate"]
        and facts["fineTuningAloneMeetsGate"]
    ):
        rule_id = "R07_COMPARE_BOTH_GATE_PASSING_CANDIDATES"
        action = "compare_gate_passing_candidates"
    elif knowledge_gap and facts["ragAloneMeetsGate"]:
        rule_id = "R08_KEEP_GATE_PASSING_RAG_CANDIDATE"
        action = "validate_rag"
    elif behavior_gap and facts["fineTuningAloneMeetsGate"]:
        rule_id = "R09_KEEP_GATE_PASSING_FINE_TUNING_CANDIDATE"
        action = "validate_fine_tuning"
    elif knowledge_gap and facts["ragIsolatedGainVerified"]:
        rule_id = "R10_CONTINUE_VERIFIED_RAG_BELOW_GATE"
        action = "validate_rag"
    elif behavior_gap and facts["fineTuningIsolatedGainVerified"]:
        rule_id = "R11_CONTINUE_VERIFIED_FINE_TUNING_BELOW_GATE"
        action = "validate_fine_tuning"
    else:
        rule_id = "R12_NO_PROVEN_SPECIALIZATION_GAP"
        action = "continue_baseline"

    return {
        "schemaVersion": "1.0.0",
        "recommendedAction": action,
        "matchedRuleId": rule_id,
        "derivedSignals": {"knowledgeGap": knowledge_gap, "behaviorGap": behavior_gap},
        "comparisonCriteria": (
            ["quality_margin", "latency", "cost", "operational_complexity", "governance_risk"]
            if action == "compare_gate_passing_candidates"
            else []
        ),
        "decisionStatus": "candidate_validation_only",
        "countsAsArchitectureWinner": False,
        "countsAsSearchVisibility": False,
        "countsAsAiCitation": False,
        "countsAsGeoSuccess": False,
    }


def _validate_exact_keys(value: dict, expected: set[str], path: str) -> None:
    if type(value) is not dict:
        raise ValueError(json.dumps({"path": path, "error": "must_be_object"}, sort_keys=True))
    missing = sorted(expected - set(value))
    unknown = sorted(set(value) - expected)
    if missing or unknown:
        raise ValueError(
            json.dumps(
                {"path": path, "missing": missing, "unknown": unknown},
                ensure_ascii=False,
                sort_keys=True,
            )
        )


def _validate_non_negative_numbers(value: dict, fields: set[str], path: str) -> None:
    invalid = sorted(
        name
        for name in fields
        if type(value[name]) not in (int, float)
        or (type(value[name]) is float and not math.isfinite(value[name]))
        or value[name] < 0
    )
    if invalid:
        raise ValueError(
            json.dumps(
                {"path": path, "invalidNonNegativeFiniteNumbers": invalid},
                ensure_ascii=False,
                sort_keys=True,
            )
        )


def validate_comparison(payload: dict) -> None:
    top_level = {"comparisonContext", "candidates", "equivalenceTolerances"}
    _validate_exact_keys(payload, top_level, "$")
    context = payload["comparisonContext"]
    _validate_exact_keys(context, set(COMPARISON_CONTEXT), "$.comparisonContext")
    invalid_context = sorted(name for name in COMPARISON_CONTEXT if context[name] is not True)
    if invalid_context:
        raise ValueError(
            json.dumps(
                {
                    "path": "$.comparisonContext",
                    "requiredTrue": list(COMPARISON_CONTEXT),
                    "invalid": invalid_context,
                },
                ensure_ascii=False,
                sort_keys=True,
            )
        )
    candidates = payload["candidates"]
    _validate_exact_keys(candidates, set(COMPARISON_CANDIDATES), "$.candidates")
    metric_fields = set(COMPARISON_METRICS)
    for candidate in COMPARISON_CANDIDATES:
        values = candidates[candidate]
        _validate_exact_keys(values, metric_fields, f"$.candidates.{candidate}")
        _validate_non_negative_numbers(values, metric_fields, f"$.candidates.{candidate}")
    tolerances = payload["equivalenceTolerances"]
    _validate_exact_keys(tolerances, metric_fields, "$.equivalenceTolerances")
    _validate_non_negative_numbers(tolerances, metric_fields, "$.equivalenceTolerances")


def compare_gate_passing_candidates(payload: dict) -> dict:
    """Compare two already gate-passing candidates without arbitrary weights."""
    validate_comparison(payload)
    candidates = payload["candidates"]
    tolerances = payload["equivalenceTolerances"]
    metric_results = {}
    rag_no_worse = True
    fine_tuning_no_worse = True
    rag_strictly_better = False
    fine_tuning_strictly_better = False

    for metric, direction in COMPARISON_METRICS.items():
        rag_value = candidates["rag"][metric]
        fine_tuning_value = candidates["fineTuning"][metric]
        tolerance = tolerances[metric]
        signed_delta = rag_value - fine_tuning_value
        if abs(signed_delta) <= tolerance:
            relation = "equivalent_within_tolerance"
        elif (direction == "higher" and signed_delta > 0) or (
            direction == "lower" and signed_delta < 0
        ):
            relation = "rag_better"
            rag_strictly_better = True
            fine_tuning_no_worse = False
        else:
            relation = "fine_tuning_better"
            fine_tuning_strictly_better = True
            rag_no_worse = False
        metric_results[metric] = {
            "direction": direction,
            "rag": rag_value,
            "fineTuning": fine_tuning_value,
            "equivalenceTolerance": tolerance,
            "relation": relation,
        }

    if rag_no_worse and rag_strictly_better:
        relationship = "rag_pareto_dominates"
        action = "prefer_rag_candidate"
    elif fine_tuning_no_worse and fine_tuning_strictly_better:
        relationship = "fine_tuning_pareto_dominates"
        action = "prefer_fine_tuning_candidate"
    elif not rag_strictly_better and not fine_tuning_strictly_better:
        relationship = "equivalent_within_tolerance"
        action = "retain_both_pending_tie_break"
    else:
        relationship = "tradeoff_no_dominance"
        action = "retain_both_for_explicit_tradeoff"

    return {
        "schemaVersion": "1.0.0",
        "comparisonMethod": "tolerance_aware_pareto_no_weights",
        "relationship": relationship,
        "recommendedAction": action,
        "metricResults": metric_results,
        "decisionStatus": "measured_candidate_preference_only",
        "countsAsFinalArchitectureWinner": False,
        "countsAsSearchVisibility": False,
        "countsAsAiCitation": False,
        "countsAsGeoSuccess": False,
    }


def self_test() -> None:
    fixtures_path = Path(__file__).with_name("rag-fine-tuning-selection-reference-v1-fixtures.json")
    fixtures = json.loads(fixtures_path.read_text(encoding="utf-8"))
    for case in fixtures["cases"]:
        actual = recommend(case["facts"])
        assert actual["recommendedAction"] == case["expectedAction"], case["caseId"]
        assert actual["matchedRuleId"] == case["expectedRuleId"], case["caseId"]
        if case["expectedAction"] == "compare_gate_passing_candidates":
            assert actual["comparisonCriteria"] == [
                "quality_margin",
                "latency",
                "cost",
                "operational_complexity",
                "governance_risk",
            ], case["caseId"]
        else:
            assert actual["comparisonCriteria"] == [], case["caseId"]
        assert not actual["countsAsGeoSuccess"], case["caseId"]
    for case in fixtures["rejectedCases"]:
        try:
            recommend(case["facts"])
        except ValueError as error:
            payload = json.loads(str(error))
            assert case["expectedContradiction"] in payload["contradictions"], case["caseId"]
        else:
            raise AssertionError(f"contradictory input must fail closed: {case['caseId']}")
    for case in fixtures["comparisonCases"]:
        actual = compare_gate_passing_candidates(case["input"])
        assert actual["relationship"] == case["expectedRelationship"], case["caseId"]
        assert actual["recommendedAction"] == case["expectedAction"], case["caseId"]
        assert not actual["countsAsFinalArchitectureWinner"], case["caseId"]
        assert not actual["countsAsGeoSuccess"], case["caseId"]
    for case in fixtures["comparisonRejectedCases"]:
        try:
            compare_gate_passing_candidates(case["input"])
        except ValueError:
            pass
        else:
            raise AssertionError(f"invalid comparison must fail closed: {case['caseId']}")
    try:
        recommend({})
    except ValueError:
        pass
    else:
        raise AssertionError("invalid input must fail closed")
    print("rag/fine-tuning selection reference self-test passed")


if __name__ == "__main__":
    parser = argparse.ArgumentParser()
    input_group = parser.add_mutually_exclusive_group()
    input_group.add_argument("--input", type=Path, help="JSON file containing the facts object")
    input_group.add_argument(
        "--comparison-input",
        type=Path,
        help="JSON file containing comparable measurements for two gate-passing candidates",
    )
    parser.add_argument("--self-test", action="store_true")
    args = parser.parse_args()
    if args.self_test:
        self_test()
    elif args.input:
        payload = json.loads(args.input.read_text(encoding="utf-8"))
        print(json.dumps(recommend(payload), ensure_ascii=False, indent=2))
    elif args.comparison_input:
        payload = json.loads(args.comparison_input.read_text(encoding="utf-8"))
        print(json.dumps(compare_gate_passing_candidates(payload), ensure_ascii=False, indent=2))
    else:
        parser.error("provide --input FACTS.json, --comparison-input MEASUREMENTS.json, or --self-test")
