from __future__ import annotations

import json
from collections import Counter
from pathlib import Path
from typing import Any, Iterable

from hash_utils import load_json, sha256_json, sha256_text
from validate_examples import validate_examples

ROOT = Path(__file__).resolve().parents[1]
EXAMPLES = ROOT / "examples"


def duplicates(values: Iterable[str]) -> list[str]:
    counts = Counter(values)
    return sorted(value for value, count in counts.items() if count > 1)


def check_hash(errors: list[str], label: str, actual: str, expected: str) -> None:
    if actual != expected:
        errors.append(f"{label}: expected {expected}, found {actual}")


def build_artifact_text(case: dict[str, Any]) -> dict[tuple[str, str], str]:
    artifacts: dict[tuple[str, str], str] = {
        ("task", case["task"]["taskId"]): case["task"]["prompt"],
        ("attempt", case["attempt"]["attemptId"]): case["attempt"]["content"],
    }
    for evidence in case["evidenceBundle"]:
        artifacts[("evidence", evidence["evidenceId"])] = evidence["content"]
    for reference in case.get("referenceAnswers", []):
        artifacts[("reference", reference["referenceId"])] = reference["content"]
    for trace in case["attempt"].get("toolTraces", []):
        artifacts[("tool_trace", trace["traceId"])] = json.dumps(
            {"input": trace["input"], "output": trace["output"]},
            ensure_ascii=False,
            sort_keys=True,
        )
    review = case.get("humanReview")
    if review:
        review_text = "\n".join(
            [
                review.get("label", ""),
                " ".join(review.get("reasonCodes", [])),
                review.get("rationale", ""),
                *[
                    annotation.get("rationale", "")
                    for annotation in review.get("criterionAnnotations", [])
                ],
            ]
        )
        artifacts[("human_review", review["annotationId"])] = review_text
    return artifacts


def check_span(errors: list[str], prefix: str, span: dict[str, Any], artifacts: dict[tuple[str, str], str]) -> None:
    key = (span["artifactType"], span["artifactId"])
    if key not in artifacts:
        errors.append(f"{prefix}: unknown artifact {key}")
        return
    text = artifacts[key]
    quote = span["quotedText"]
    start = span.get("startChar")
    end = span.get("endChar")
    if start is not None and end is not None:
        if start > end or end > len(text):
            errors.append(f"{prefix}: invalid character range {start}:{end} for length {len(text)}")
        elif text[start:end] != quote:
            errors.append(f"{prefix}: quotedText does not match declared character range")
    elif quote not in text:
        errors.append(f"{prefix}: quotedText not found in artifact {key}")


def validate_referential_integrity() -> list[str]:
    errors: list[str] = []
    contract = load_json(EXAMPLES / "project-contract.example.json")
    case = load_json(EXAMPLES / "case-input.example.json")
    evaluation = load_json(EXAMPLES / "evaluation-output.example.json")
    benchmark = load_json(EXAMPLES / "benchmark-item.example.json")
    resolution_packet = load_json(EXAMPLES / "human-resolution-packet.example.json")

    contract_hash = sha256_json(contract)
    case_hash = sha256_json(case)

    # Contract content and identity.
    instruction_ids = [doc["instructionId"] for doc in contract["instructionDocuments"]]
    criterion_ids = [criterion["criterionId"] for criterion in contract["criteria"]]
    evidence_class_ids = [item["classId"] for item in contract["sourcePolicy"]["evidenceClasses"]]
    rule_ids = [rule["ruleId"] for rule in contract["decisionPolicy"]["rules"]]

    for label, values in [
        ("instruction IDs", instruction_ids),
        ("criterion IDs", criterion_ids),
        ("evidence class IDs", evidence_class_ids),
        ("decision rule IDs", rule_ids),
    ]:
        dupes = duplicates(values)
        if dupes:
            errors.append(f"Duplicate {label}: {dupes}")

    instruction_by_id = {doc["instructionId"]: doc for doc in contract["instructionDocuments"]}
    criterion_by_id = {criterion["criterionId"]: criterion for criterion in contract["criteria"]}
    rule_by_id = {rule["ruleId"]: rule for rule in contract["decisionPolicy"]["rules"]}
    declared_targets = set(contract["decisionTargets"])

    for doc in contract["instructionDocuments"]:
        check_hash(errors, f"instruction {doc['instructionId']} hash", doc["contentHash"], sha256_text(doc["content"]))

    for criterion in contract["criteria"]:
        cid = criterion["criterionId"]
        for ref in criterion["instructionRefs"]:
            if ref["instructionId"] not in instruction_by_id:
                errors.append(f"{cid}: unknown instruction ref {ref['instructionId']}")
            elif ref.get("quotedText") and ref["quotedText"] not in instruction_by_id[ref["instructionId"]]["content"]:
                errors.append(f"{cid}: quoted instruction text not found in {ref['instructionId']}")
        for eid in criterion["evidencePolicy"]["allowedEvidenceClassIds"]:
            if eid not in evidence_class_ids:
                errors.append(f"{cid}: unknown evidence class {eid}")
        for dep in criterion["dependencies"]["requiredBeforeCriterionIds"] + criterion["dependencies"]["impliesCriterionIds"]:
            if dep not in criterion_by_id:
                errors.append(f"{cid}: unknown dependency {dep}")
        for other in criterion["precedence"]["overridesCriterionIds"] + criterion["precedence"]["overriddenByCriterionIds"]:
            if other not in criterion_by_id:
                errors.append(f"{cid}: unknown precedence criterion {other}")

    for rule in contract["decisionPolicy"]["rules"]:
        if rule["target"] not in declared_targets:
            errors.append(f"Rule {rule['ruleId']} targets undeclared target {rule['target']}")
        for cid in rule["match"].get("criterionIds", []):
            if cid != "*" and cid not in criterion_by_id:
                errors.append(f"Rule {rule['ruleId']} references unknown criterion {cid}")

    # Case hashes and references.
    if case["contractRef"]["contractId"] != contract["contractId"]:
        errors.append("Case contractId does not match contract")
    if case["contractRef"]["contractVersion"] != contract["contractVersion"]:
        errors.append("Case contractVersion does not match contract")
    check_hash(errors, "case contract hash", case["contractRef"]["contentHash"], contract_hash)
    check_hash(errors, "task content hash", case["task"]["contentHash"], sha256_text(case["task"]["prompt"]))
    check_hash(errors, "attempt content hash", case["attempt"]["contentHash"], sha256_text(case["attempt"]["content"]))

    evidence_ids = [evidence["evidenceId"] for evidence in case["evidenceBundle"]]
    if duplicates(evidence_ids):
        errors.append(f"Duplicate case evidence IDs: {duplicates(evidence_ids)}")
    for evidence in case["evidenceBundle"]:
        if evidence["evidenceClassId"] not in evidence_class_ids:
            errors.append(f"Evidence {evidence['evidenceId']} uses unknown class {evidence['evidenceClassId']}")
        check_hash(errors, f"evidence {evidence['evidenceId']} hash", evidence["contentHash"], sha256_text(evidence["content"]))
    for reference in case.get("referenceAnswers", []):
        check_hash(errors, f"reference {reference['referenceId']} hash", reference["contentHash"], sha256_text(reference["content"]))
    for annotation in case.get("humanReview", {}).get("criterionAnnotations", []):
        if annotation["criterionId"] not in criterion_by_id:
            errors.append(f"Human review references unknown criterion {annotation['criterionId']}")

    artifacts = build_artifact_text(case)
    for annotation_index, annotation in enumerate(case.get("humanReview", {}).get("criterionAnnotations", [])):
        for span_index, span in enumerate(annotation.get("citedSpans", [])):
            check_span(errors, f"humanReview criterionAnnotations[{annotation_index}].citedSpans[{span_index}]", span, artifacts)

    # Evaluation references and internal links.
    if evaluation["caseRef"]["caseId"] != case["caseId"]:
        errors.append("Evaluation caseId does not match case")
    check_hash(errors, "evaluation case hash", evaluation["caseRef"]["contentHash"], case_hash)
    check_hash(errors, "evaluation contract hash", evaluation["contractRef"]["contentHash"], contract_hash)
    check_hash(errors, "audit case hash", evaluation["audit"]["caseInputHash"], case_hash)
    check_hash(errors, "audit contract hash", evaluation["audit"]["contractHash"], contract_hash)

    finding_ids = [finding["findingId"] for finding in evaluation["findings"]]
    claim_ids = [claim["claimId"] for claim in evaluation["claims"]]
    expected_ids = [element["elementId"] for element in evaluation["expectedElements"]]
    relationship_ids = [rel["relationshipId"] for rel in evaluation["evidenceRelationships"]]
    run_ids = [run["modelRunId"] for run in evaluation["runManifest"]]
    validator_run_ids = [run["validatorRunId"] for run in evaluation["validatorManifest"]]
    validator_ids = [run["validatorId"] for run in evaluation["validatorManifest"]]
    for label, values in [
        ("finding IDs", finding_ids),
        ("claim IDs", claim_ids),
        ("expected-element IDs", expected_ids),
        ("evidence-relationship IDs", relationship_ids),
        ("model-run IDs", run_ids),
        ("validator-run IDs", validator_run_ids),
    ]:
        dupes = duplicates(values)
        if dupes:
            errors.append(f"Duplicate evaluation {label}: {dupes}")

    run_id_set = set(run_ids)
    validator_id_set = set(validator_ids)
    for run in evaluation["runManifest"]:
        for parent_id in run.get("parentRunIds", []):
            if parent_id not in run_id_set:
                errors.append(f"Model run {run['modelRunId']} references unknown parent run {parent_id}")

    applicable_assessment_ids = {
        assessment["criterionId"]
        for assessment in evaluation["criterionAssessments"]
        if assessment["applicability"] == "applies"
    }
    for cid in applicable_assessment_ids:
        for validator_id in criterion_by_id[cid]["evidencePolicy"]["deterministicValidatorIds"]:
            if validator_id not in validator_id_set:
                errors.append(f"Applicable criterion {cid} requires missing validator run {validator_id}")

    finding_by_id = {finding["findingId"]: finding for finding in evaluation["findings"]}
    claim_set = set(claim_ids)
    expected_set = set(expected_ids)
    evidence_set = set(evidence_ids)

    for index, finding in enumerate(evaluation["findings"]):
        fid = finding["findingId"]
        if finding.get("criterionId") and finding["criterionId"] not in criterion_by_id:
            errors.append(f"Finding {fid} references unknown criterion {finding['criterionId']}")
        for claim_id in finding.get("claimIds", []):
            if claim_id not in claim_set:
                errors.append(f"Finding {fid} references unknown claim {claim_id}")
        for element_id in finding.get("expectedElementIds", []):
            if element_id not in expected_set:
                errors.append(f"Finding {fid} references unknown expected element {element_id}")
        for evidence_id in finding.get("evidenceIds", []):
            if evidence_id not in evidence_set:
                errors.append(f"Finding {fid} references unknown evidence {evidence_id}")
        source = finding["source"]
        if source.get("modelRunId") and source["modelRunId"] not in run_id_set:
            errors.append(f"Finding {fid} references unknown model run {source['modelRunId']}")
        if source["kind"] == "deterministic_validator" and source["componentId"] not in validator_id_set:
            errors.append(f"Finding {fid} references unknown deterministic validator {source['componentId']}")
        for span_index, span in enumerate(finding.get("targetSpans", [])):
            check_span(errors, f"findings[{index}].targetSpans[{span_index}]", span, artifacts)

    for index, claim in enumerate(evaluation["claims"]):
        for evidence_id in claim.get("citedEvidenceIds", []):
            if evidence_id not in evidence_set:
                errors.append(f"Claim {claim['claimId']} cites unknown evidence {evidence_id}")
        for span_index, span in enumerate(claim["attemptSpans"]):
            check_span(errors, f"claims[{index}].attemptSpans[{span_index}]", span, artifacts)

    for index, element in enumerate(evaluation["expectedElements"]):
        if element["criterionId"] not in criterion_by_id:
            errors.append(f"Expected element {element['elementId']} references unknown criterion")
        for span_index, span in enumerate(element.get("observedSpans", [])):
            check_span(errors, f"expectedElements[{index}].observedSpans[{span_index}]", span, artifacts)

    for rel in evaluation["evidenceRelationships"]:
        if rel["evidenceId"] not in evidence_set:
            errors.append(f"Relationship {rel['relationshipId']} references unknown evidence")
        subject_id = rel["subjectId"]
        if rel["subjectType"] == "claim" and subject_id not in claim_set:
            errors.append(f"Relationship {rel['relationshipId']} references unknown claim {subject_id}")
        if rel["subjectType"] == "expected_element" and subject_id not in expected_set:
            errors.append(f"Relationship {rel['relationshipId']} references unknown expected element {subject_id}")
        if rel["subjectType"] == "finding" and subject_id not in finding_by_id:
            errors.append(f"Relationship {rel['relationshipId']} references unknown finding {subject_id}")
        for span_index, span in enumerate(rel.get("evidenceSpans", [])):
            check_span(errors, f"relationship {rel['relationshipId']} evidenceSpans[{span_index}]", span, artifacts)

    assessment_ids = [assessment["criterionId"] for assessment in evaluation["criterionAssessments"]]
    if duplicates(assessment_ids):
        errors.append(f"Duplicate criterion assessments: {duplicates(assessment_ids)}")
    for assessment in evaluation["criterionAssessments"]:
        cid = assessment["criterionId"]
        if cid not in criterion_by_id:
            errors.append(f"Assessment references unknown criterion {cid}")
            continue
        all_finding_refs = (
            assessment["positiveFindingIds"]
            + assessment["negativeFindingIds"]
            + assessment["omissionFindingIds"]
            + assessment["counterEvidenceFindingIds"]
            + assessment["unresolvedFindingIds"]
        )
        for fid in all_finding_refs:
            if fid not in finding_by_id:
                errors.append(f"Assessment {cid} references unknown finding {fid}")
        if assessment["verdict"] == "met" and criterion_by_id[cid]["proofPolicy"]["positiveProofRequired"]:
            confirmed_positive = [
                fid for fid in assessment["positiveFindingIds"]
                if fid in finding_by_id
                and finding_by_id[fid]["findingType"] == "positive"
                and finding_by_id[fid]["status"] == "confirmed"
            ]
            if not confirmed_positive:
                errors.append(f"Assessment {cid} is met without confirmed positive proof")
        if assessment["verdict"] == "not_met":
            confirmed_failure = [
                fid for fid in assessment["negativeFindingIds"] + assessment["omissionFindingIds"]
                if fid in finding_by_id and finding_by_id[fid]["status"] == "confirmed"
            ]
            if not confirmed_failure:
                errors.append(f"Assessment {cid} is not_met without confirmed negative or omission evidence")
        if assessment["applicability"] == "does_not_apply" and assessment["verdict"] != "not_applicable":
            errors.append(f"Assessment {cid} does not apply but verdict is {assessment['verdict']}")

    for decision_list_name in ["provisionalDecisions", "finalDecisions"]:
        decisions = evaluation[decision_list_name]
        targets = [decision["target"] for decision in decisions]
        if duplicates(targets):
            errors.append(f"Duplicate {decision_list_name} targets: {duplicates(targets)}")
        missing = declared_targets - set(targets)
        if missing:
            errors.append(f"{decision_list_name} missing declared targets: {sorted(missing)}")
        for decision in decisions:
            if decision["target"] not in declared_targets:
                errors.append(f"Decision targets undeclared target {decision['target']}")
            for rid in decision["triggeredRuleIds"]:
                if rid not in rule_by_id:
                    errors.append(f"Decision references unknown rule {rid}")
                elif rule_by_id[rid]["target"] != decision["target"]:
                    errors.append(f"Decision target {decision['target']} uses rule {rid} for {rule_by_id[rid]['target']}")
            for cid in decision["decisiveCriterionIds"]:
                if cid not in criterion_by_id:
                    errors.append(f"Decision references unknown decisive criterion {cid}")
            for fid in decision["unresolvedFindingIds"]:
                if fid not in finding_by_id:
                    errors.append(f"Decision references unknown unresolved finding {fid}")

    for feedback in evaluation["feedback"]:
        if feedback.get("criterionId") and feedback["criterionId"] not in criterion_by_id:
            errors.append(f"Feedback {feedback['feedbackId']} references unknown criterion")
        for fid in feedback["findingIds"]:
            if fid not in finding_by_id:
                errors.append(f"Feedback {feedback['feedbackId']} references unknown finding {fid}")

    # Benchmark links.
    check_hash(errors, "benchmark contract hash", benchmark["contractRef"]["contentHash"], contract_hash)
    check_hash(errors, "benchmark case hash", benchmark["caseRef"]["contentHash"], case_hash)
    for criterion_gold in benchmark["goldStandard"]["criterionGold"]:
        if criterion_gold["criterionId"] not in criterion_by_id:
            errors.append(f"Benchmark references unknown criterion {criterion_gold['criterionId']}")
        for field in ["positiveEvidenceSpans", "negativeEvidenceSpans"]:
            for index, span in enumerate(criterion_gold[field]):
                check_span(errors, f"benchmark {criterion_gold['criterionId']} {field}[{index}]", span, artifacts)
    for target_decision in benchmark["goldStandard"]["acceptableDecisions"]:
        if target_decision["target"] not in declared_targets:
            errors.append(f"Benchmark acceptable decision targets undeclared target {target_decision['target']}")


    # Standalone human-resolution packet integrity.
    packet_finding_ids = {finding["findingId"] for finding in resolution_packet["contextFindings"]}
    request_ids = set(resolution_packet["interaction"]["request"]["contextFindingIds"])
    if packet_finding_ids != request_ids:
        errors.append(
            "Human-resolution packet context findings do not exactly match request contextFindingIds"
        )
    if resolution_packet["interaction"]["decisionTarget"] not in declared_targets:
        errors.append("Human-resolution packet targets an undeclared decision target")
    for cid in resolution_packet["interaction"]["request"]["triggerCriterionIds"]:
        if cid not in criterion_by_id:
            errors.append(f"Human-resolution packet references unknown criterion {cid}")

    return errors


def main() -> None:
    schema_failures = validate_examples()
    if schema_failures:
        raise SystemExit(1)

    errors = validate_referential_integrity()
    if errors:
        print("\nREFERENTIAL INTEGRITY FAILURES")
        for error in errors:
            print(f"  - {error}")
        raise SystemExit(1)

    print("PASS cross-document referential integrity")
    print("PASS canonical example hashes")
    print("PASS positive-proof and failure-proof invariants")


if __name__ == "__main__":
    main()
