#!/usr/bin/env python3
"""Validate and retrieve a small, file-backed memory store."""

from __future__ import annotations

import argparse
import hashlib
import json
import sys
from datetime import date
from pathlib import Path


def load_json(path: Path) -> dict:
    data = json.loads(path.read_text(encoding="utf-8"))
    if not isinstance(data, dict):
        raise ValueError(f"{path}: expected a JSON object")
    return data


def safe_path(root: Path, relative: str) -> Path:
    candidate = (root / relative).resolve()
    if not candidate.is_relative_to(root.resolve()):
        raise ValueError(f"path escapes the memory root: {relative}")
    if not candidate.is_file():
        raise ValueError(f"file does not exist: {relative}")
    return candidate


def parse_day(value: object, field: str, record_id: str) -> date:
    if not isinstance(value, str):
        raise ValueError(f"{record_id}: {field} must be YYYY-MM-DD")
    try:
        return date.fromisoformat(value)
    except ValueError as exc:
        raise ValueError(f"{record_id}: {field} must be YYYY-MM-DD") from exc


def instruction_from_source(path: Path) -> str:
    for line in path.read_text(encoding="utf-8").splitlines():
        candidate = line.strip()
        if candidate and not candidate.startswith("#"):
            return candidate
    raise ValueError(f"{path}: durable source contains no instruction")


def validate(root: Path, as_of: date) -> tuple[dict, dict[str, dict], int]:
    index = load_json(safe_path(root, "index.json"))
    entries = index.get("entries")
    if index.get("schema_version") != "1.0" or not isinstance(entries, list) or not entries:
        raise ValueError("index.json: schema_version 1.0 and a non-empty entries list are required")

    records: dict[str, dict] = {}
    for path in sorted((root / "records").glob("*.json")):
        record = load_json(path)
        record_id = record.get("id")
        if not isinstance(record_id, str) or not record_id:
            raise ValueError(f"{path}: id is required")
        if record_id in records:
            raise ValueError(f"duplicate record id: {record_id}")
        record["_path"] = path
        records[record_id] = record

    if not records:
        raise ValueError("records: at least one record is required")

    source_count = 0
    for record_id, record in records.items():
        status = record.get("status")
        if status not in {"current", "superseded"}:
            raise ValueError(f"{record_id}: status must be current or superseded")

        source_ref = record.get("durable_source")
        expected_hash = record.get("source_sha256")
        if not isinstance(source_ref, str) or not isinstance(expected_hash, str):
            raise ValueError(f"{record_id}: durable_source and source_sha256 are required")
        source = safe_path(root, source_ref)
        actual_hash = hashlib.sha256(source.read_bytes()).hexdigest()
        if actual_hash != expected_hash:
            raise ValueError(f"{record_id}: durable source hash mismatch")
        source_count += 1

        parse_day(record.get("recorded_on"), "recorded_on", record_id)
        review_on = parse_day(record.get("review_on"), "review_on", record_id)
        expires_on = parse_day(record.get("expires_on"), "expires_on", record_id)
        if review_on >= expires_on:
            raise ValueError(f"{record_id}: review_on must precede expires_on")

        replacement = record.get("superseded_by")
        predecessor = record.get("supersedes")
        if status == "superseded":
            if not isinstance(replacement, str) or not replacement:
                raise ValueError(f"{record_id}: superseded_by is required")
            if replacement == record_id:
                raise ValueError(f"{record_id}: superseded_by cannot be a self-link")
            if replacement not in records:
                raise ValueError(f"{record_id}: superseded_by does not resolve")
        elif replacement is not None:
            raise ValueError(f"{record_id}: a current record cannot declare superseded_by")

        if predecessor is not None:
            if not isinstance(predecessor, str) or not predecessor:
                raise ValueError(f"{record_id}: supersedes must be a non-empty record id")
            if predecessor == record_id:
                raise ValueError(f"{record_id}: supersedes cannot be a self-link")
            if predecessor not in records:
                raise ValueError(f"{record_id}: supersedes does not resolve")

    # Reject a cycle before accepting any reciprocal edge. A cycle can be
    # internally consistent in both directions and still have no current end.
    for start_id in records:
        seen: set[str] = set()
        cursor = start_id
        while records[cursor].get("status") == "superseded":
            if cursor in seen:
                raise ValueError(f"{start_id}: supersession cycle detected")
            seen.add(cursor)
            cursor = records[cursor]["superseded_by"]

    supersession_count = 0
    for record_id, record in records.items():
        if record.get("status") == "superseded":
            replacement_id = record["superseded_by"]
            replacement = records[replacement_id]
            if replacement.get("supersedes") != record_id:
                raise ValueError(
                    f"{record_id} -> {replacement_id}: reverse supersedes link does not agree"
                )
            supersession_count += 1

        predecessor_id = record.get("supersedes")
        if predecessor_id is not None:
            predecessor = records[predecessor_id]
            if predecessor.get("status") != "superseded":
                raise ValueError(f"{record_id}: supersedes must identify a superseded record")
            if predecessor.get("superseded_by") != record_id:
                raise ValueError(
                    f"{record_id} -> {predecessor_id}: forward superseded_by link does not agree"
                )

    seen_keys: set[str] = set()
    for entry in entries:
        if not isinstance(entry, dict):
            raise ValueError("index.json: each entry must be an object")
        key = entry.get("key")
        trigger = entry.get("read_when")
        record_id = entry.get("record_id")
        target = entry.get("target")
        record_hash = entry.get("record_sha256")
        if not all(isinstance(value, str) and value for value in (key, trigger, record_id, target, record_hash)):
            raise ValueError("index.json: key, read_when, record_id, target, and record_sha256 are required")
        if key in seen_keys:
            raise ValueError(f"index.json: duplicate key {key}")
        seen_keys.add(key)
        target_path = safe_path(root, target)
        if hashlib.sha256(target_path.read_bytes()).hexdigest() != record_hash:
            raise ValueError(f"{key}: record hash mismatch")
        record = records.get(record_id)
        if record is None:
            raise ValueError(f"{key}: record_id does not resolve")
        if target_path != record["_path"].resolve():
            raise ValueError(f"{key}: target and record_id point to different records")
        if record.get("status") != "current":
            raise ValueError(f"{key}: index points to a superseded record")
        review_on = parse_day(record.get("review_on"), "review_on", record_id)
        expires_on = parse_day(record.get("expires_on"), "expires_on", record_id)
        if as_of >= expires_on:
            raise ValueError(f"{key}: record expired on {expires_on.isoformat()}")
        if as_of >= review_on:
            raise ValueError(f"{key}: review due on {review_on.isoformat()}")

    print(f"PASS index: {len(entries)} current pointer")
    print(f"PASS records: {len(records)} checked")
    print(
        f"PASS supersession: {supersession_count} reciprocal link; "
        "no self-links or cycles"
    )
    print(f"PASS source hashes: {source_count} matched")
    print(f"PASS dates: {len(entries)} current record valid as of {as_of.isoformat()}")
    return index, records, supersession_count


def main() -> int:
    parser = argparse.ArgumentParser()
    parser.add_argument("--root", type=Path, default=Path(__file__).parent / "demo-store")
    parser.add_argument("--as-of", type=date.fromisoformat, default=date.today())
    subparsers = parser.add_subparsers(dest="command", required=True)
    subparsers.add_parser("validate")
    retrieve = subparsers.add_parser("retrieve")
    retrieve.add_argument("key")
    args = parser.parse_args()

    try:
        index, records, _ = validate(args.root.resolve(), args.as_of)
        if args.command == "retrieve":
            entry = next((item for item in index["entries"] if item["key"] == args.key), None)
            if entry is None:
                raise ValueError(f"no pointer matches key: {args.key}")
            record = records[entry["record_id"]]
            print(f"key={entry['key']}")
            print(f"source={record['durable_source']}")
            source = safe_path(args.root.resolve(), record["durable_source"])
            print(f"instruction={instruction_from_source(source)}")
            print(f"verified={record['last_verified']}")
            print(f"review_on={record['review_on']}")
            print(f"expires_on={record['expires_on']}")
    except (OSError, ValueError, json.JSONDecodeError) as exc:
        print(f"FAIL: {exc}", file=sys.stderr)
        return 2
    return 0


if __name__ == "__main__":
    raise SystemExit(main())
