#!/usr/bin/env python3
"""Fail-closed validator for claims bound to the RAG/fine-tuning source map."""

from __future__ import annotations

import argparse
import json
from pathlib import Path


SOURCE_SCOPES = {
    "microsoft-rag-fine-tuning-choice-20260720": {
        "architecture_mechanism",
        "selection_dimensions",
    },
    "aws-rag-vs-finetuning-20260720": {
        "architecture_mechanism",
        "selection_dimensions",
        "citation_traceability_dimension",
    },
    "google-cloud-rag-retrieval-evaluation-20260720": {
        "retrieval_evaluation_method",
    },
    "nist-ai-rmf-core-20260720": {
        "governance_process",
    },
    "baidu-qianfan-rag-knowledge-base-20260813": {
        "architecture_mechanism",
        "retrieval_evaluation_method",
        "china_rag_workflow",
    },
    "tencent-cloud-rag-guide-20260813": {
        "architecture_mechanism",
        "china_rag_workflow",
    },
}

FORBIDDEN_OUTCOME_SCOPES = {
    "current_project_architecture_winner",
    "provider_quality_ranking",
    "search_visibility",
    "ai_citation",
    "geo_success",
}

ALLOWED_SCOPES = set().union(*SOURCE_SCOPES.values())
ALL_SCOPES = ALLOWED_SCOPES | FORBIDDEN_OUTCOME_SCOPES
TOP_LEVEL_KEYS = {"claims"}
CLAIM_KEYS = {"claimId", "statement", "scope", "evidenceIds"}


def _exact_keys(value: object, expected: set[str], path: str) -> list[dict]:
    if type(value) is not dict:
        return [{"path": path, "code": "must_be_object"}]
    missing = sorted(expected - set(value))
    unknown = sorted(set(value) - expected)
    if missing or unknown:
        return [{"path": path, "code": "invalid_keys", "missing": missing, "unknown": unknown}]
    return []


def _load_source_ids(source_map: dict) -> set[str]:
    mappings = source_map.get("sourceMappings")
    if type(mappings) is not list:
        raise ValueError("sourceMappings must be an array")
    source_ids = [row.get("evidenceId") for row in mappings if type(row) is dict]
    if any(type(value) is not str or not value for value in source_ids):
        raise ValueError("every source mapping must have a non-empty evidenceId")
    if len(source_ids) != len(set(source_ids)):
        raise ValueError("source map contains duplicate evidenceId values")
    expected = set(SOURCE_SCOPES)
    actual = set(source_ids)
    if actual != expected:
        raise ValueError(
            json.dumps(
                {"sourceMapMismatch": {"missing": sorted(expected - actual), "unknown": sorted(actual - expected)}},
                ensure_ascii=False,
                sort_keys=True,
            )
        )
    return actual


def validate(payload: object, source_map: dict) -> dict:
    source_ids = _load_source_ids(source_map)
    errors = _exact_keys(payload, TOP_LEVEL_KEYS, "$")
    claims = payload.get("claims") if type(payload) is dict else None
    if type(claims) is not list or not claims:
        errors.append({"path": "$.claims", "code": "must_be_non_empty_array"})
        claims = []

    claim_ids: set[str] = set()
    results = []
    for index, claim in enumerate(claims):
        path = f"$.claims[{index}]"
        claim_errors = _exact_keys(claim, CLAIM_KEYS, path)
        if claim_errors:
            errors.extend(claim_errors)
            continue

        claim_id = claim["claimId"]
        statement = claim["statement"]
        scope = claim["scope"]
        evidence_ids = claim["evidenceIds"]

        if type(claim_id) is not str or not claim_id.strip():
            claim_errors.append({"path": f"{path}.claimId", "code": "must_be_non_empty_string"})
        elif claim_id in claim_ids:
            claim_errors.append({"path": f"{path}.claimId", "code": "duplicate_claim_id"})
        else:
            claim_ids.add(claim_id)
        if type(statement) is not str or not statement.strip():
            claim_errors.append({"path": f"{path}.statement", "code": "must_be_non_empty_string"})
        if scope not in ALL_SCOPES:
            claim_errors.append({"path": f"{path}.scope", "code": "unknown_scope", "value": scope})
        if (
            type(evidence_ids) is not list
            or not evidence_ids
            or any(type(value) is not str or not value for value in evidence_ids)
        ):
            claim_errors.append({"path": f"{path}.evidenceIds", "code": "must_be_non_empty_string_array"})
            evidence_ids = []
        elif len(evidence_ids) != len(set(evidence_ids)):
            claim_errors.append({"path": f"{path}.evidenceIds", "code": "duplicate_evidence_id"})

        unknown_ids = sorted(set(evidence_ids) - source_ids)
        if unknown_ids:
            claim_errors.append(
                {"path": f"{path}.evidenceIds", "code": "unknown_evidence_id", "values": unknown_ids}
            )

        if scope in FORBIDDEN_OUTCOME_SCOPES:
            claim_errors.append(
                {
                    "path": f"{path}.scope",
                    "code": "source_guidance_cannot_prove_outcome",
                    "scope": scope,
                }
            )
        elif scope in ALLOWED_SCOPES and not unknown_ids:
            unsupported = sorted(
                evidence_id
                for evidence_id in evidence_ids
                if scope not in SOURCE_SCOPES[evidence_id]
            )
            if unsupported:
                claim_errors.append(
                    {
                        "path": f"{path}.evidenceIds",
                        "code": "evidence_outside_declared_scope",
                        "scope": scope,
                        "values": unsupported,
                    }
                )

        if claim_errors:
            errors.extend(claim_errors)
            status = "rejected"
        else:
            status = "supported_within_source_scope"
        results.append(
            {
                "claimId": claim_id,
                "scope": scope,
                "status": status,
                "evidenceIds": evidence_ids,
                "countsAsCurrentProjectOutcome": False,
                "countsAsSearchVisibility": False,
                "countsAsAiCitation": False,
                "countsAsGeoSuccess": False,
            }
        )

    return {
        "schemaVersion": "1.0.0",
        "passed": not errors,
        "claimCount": len(claims),
        "results": results,
        "errors": errors,
        "boundary": {
            "sourceScopeValidationOnly": True,
            "requiresProjectExperimentForArchitectureWinner": True,
            "requiresIndependentObservationForSearchOrCitationOutcome": True,
            "countsAsSearchVisibility": False,
            "countsAsAiCitation": False,
            "countsAsGeoSuccess": False,
        },
    }


def self_test(source_map: dict, fixtures: dict) -> None:
    for case in fixtures["acceptedCases"]:
        result = validate(case["input"], source_map)
        assert result["passed"], (case["caseId"], result)
    for case in fixtures["rejectedCases"]:
        result = validate(case["input"], source_map)
        assert not result["passed"], case["caseId"]
        codes = {error["code"] for error in result["errors"]}
        assert case["expectedCode"] in codes, (case["caseId"], codes)
    print(
        json.dumps(
            {
                "status": "self_test_passed",
                "acceptedCases": len(fixtures["acceptedCases"]),
                "rejectedCases": len(fixtures["rejectedCases"]),
                "countsAsGeoSuccess": False,
            },
            ensure_ascii=False,
            sort_keys=True,
        )
    )


def main() -> None:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--input", type=Path)
    parser.add_argument(
        "--source-map",
        type=Path,
        default=Path(__file__).with_name("rag-fine-tuning-reliable-source-evidence-map-v1.json"),
    )
    parser.add_argument(
        "--fixtures",
        type=Path,
        default=Path(__file__).with_name("rag-fine-tuning-source-claim-validator-v1-fixtures.json"),
    )
    parser.add_argument("--self-test", action="store_true")
    args = parser.parse_args()
    source_map = json.loads(args.source_map.read_text(encoding="utf-8"))
    if args.self_test:
        fixtures = json.loads(args.fixtures.read_text(encoding="utf-8"))
        self_test(source_map, fixtures)
        return
    if args.input is None:
        parser.error("--input is required unless --self-test is used")
    payload = json.loads(args.input.read_text(encoding="utf-8"))
    result = validate(payload, source_map)
    print(json.dumps(result, ensure_ascii=False, indent=2, sort_keys=True))
    raise SystemExit(0 if result["passed"] else 2)


if __name__ == "__main__":
    main()
