#!/usr/bin/env python3
"""Deterministic local transactional memory patch/repair journal.

This tool proves local byte-state repair and rollback behavior only. It does not
provide signing, authorship, timestamp authority, semantic truth, public witness
consensus, production disaster recovery, deployment, domain control, or
real-world authorization.
"""
from __future__ import annotations
from pathlib import Path, PurePosixPath
import argparse
import hashlib
import json
import os
import shutil
import tempfile
from typing import Any

PROTECTED = {'.uai/totem.uai', '.uai/taboo.uai', '.uai/talisman.uai'}
BUNDLE_SCHEMA = 'memory-patch-bundle/1'
JOURNAL_SCHEMA = 'memory-repair-journal/1'

class RepairTransactionError(RuntimeError):
    def __init__(self, code: str, journal: dict[str, Any] | None = None):
        super().__init__(code)
        self.code = code
        self.journal = journal


def sha_bytes(data: bytes) -> str:
    return hashlib.sha256(data).hexdigest()


def sha_file(path: Path) -> str:
    h = hashlib.sha256()
    with path.open('rb') as f:
        for chunk in iter(lambda: f.read(1024 * 1024), b''):
            h.update(chunk)
    return h.hexdigest()


def load_json(path: Path) -> dict[str, Any]:
    return json.loads(path.read_text(encoding='utf-8'))


def safe_rel(value: str) -> str:
    if not isinstance(value, str) or not value or '\\' in value or '\x00' in value:
        raise RepairTransactionError('unsafe_path')
    p = PurePosixPath(value)
    if p.is_absolute() or any(part in ('', '.', '..') for part in p.parts):
        raise RepairTransactionError('unsafe_path')
    norm = p.as_posix()
    if norm != value:
        raise RepairTransactionError('noncanonical_path')
    return norm


def safe_target(root: Path, rel: str) -> Path:
    rel = safe_rel(rel)
    root = root.resolve()
    p = (root / rel).resolve()
    try:
        p.relative_to(root)
    except ValueError as exc:
        raise RepairTransactionError('unsafe_path') from exc
    return p


def snapshot_index(snapshot: dict[str, Any]) -> dict[str, dict[str, Any]]:
    if snapshot.get('schema') != 'memory-state-snapshot/1':
        raise RepairTransactionError('snapshot_schema')
    files = snapshot.get('files') or []
    if files != sorted(files, key=lambda x: x['path']):
        raise RepairTransactionError('snapshot_order')
    idx: dict[str, dict[str, Any]] = {}
    for entry in files:
        rel = safe_rel(entry['path'])
        if rel in idx:
            raise RepairTransactionError('snapshot_duplicate_path')
        idx[rel] = entry
    return idx


def bundle_material(bundle: dict[str, Any]) -> bytes:
    lines = [
        str(bundle.get('schema', '')),
        str(bundle.get('base_snapshot_id', '')),
        str(bundle.get('target_snapshot_id', '')),
        str(bundle.get('base_release', '')),
        str(bundle.get('target_release', '')),
        str(bundle.get('source_delta_digest_sha256', '')),
    ]
    for a in bundle.get('actions', []):
        lines.append('\0'.join([
            str(a.get('path', '')),
            str(a.get('classification', '')),
            str(a.get('role', '')),
            str(a.get('before_sha256') or '-'),
            str(a.get('after_sha256') or '-'),
            str(a.get('after_size_bytes', '')),
            str(a.get('payload_path') or '-'),
            str(a.get('payload_sha256') or '-'),
            '1' if a.get('automatic') else '0',
            '1' if a.get('protected_anchor') else '0',
        ]))
    dep = bundle.get('dependency_closure', {})
    lines.append('\0'.join([
        str(dep.get('target_closure_file_count', '')),
        str(dep.get('payload_file_count', '')),
        str(dep.get('omitted_unchanged_file_count', '')),
        str(dep.get('removed_review_only_count', '')),
    ]))
    return ('\n'.join(lines) + '\n').encode('utf-8')


def compute_bundle_digest(bundle: dict[str, Any]) -> str:
    return sha_bytes(bundle_material(bundle))


def build_bundle(base_snapshot: dict[str, Any], target_snapshot: dict[str, Any], target_root: Path,
                 payload_root: Path, source_delta: dict[str, Any] | None = None) -> dict[str, Any]:
    base = snapshot_index(base_snapshot)
    target = snapshot_index(target_snapshot)
    actions: list[dict[str, Any]] = []
    unchanged = 0
    removed = 0
    payload_root.mkdir(parents=True, exist_ok=True)
    for rel in sorted(set(base) | set(target)):
        before = base.get(rel)
        after = target.get(rel)
        if before == after:
            unchanged += 1
            continue
        protected = rel in PROTECTED
        if after is None:
            removed += 1
            actions.append({
                'path': rel, 'classification': 'removed', 'role': before['role'],
                'before_sha256': before['sha256'], 'after_sha256': None,
                'after_size_bytes': 0, 'payload_path': None, 'payload_sha256': None,
                'automatic': False, 'protected_anchor': protected,
                'action': 'review_only_no_automatic_delete',
            })
            continue
        src = safe_target(target_root, rel)
        if not src.is_file() or sha_file(src) != after['sha256'] or src.stat().st_size != after['size_bytes']:
            raise RepairTransactionError('target_source_hash:' + rel)
        payload_rel = 'payload/' + rel
        dst = safe_target(payload_root, rel)
        dst.parent.mkdir(parents=True, exist_ok=True)
        shutil.copy2(src, dst)
        classification = 'added' if before is None else 'changed'
        actions.append({
            'path': rel, 'classification': classification, 'role': after['role'],
            'before_sha256': None if before is None else before['sha256'],
            'after_sha256': after['sha256'], 'after_size_bytes': after['size_bytes'],
            'payload_path': payload_rel, 'payload_sha256': after['sha256'],
            'automatic': not protected, 'protected_anchor': protected,
            'action': 'manual_hold_protected' if protected else ('copy_verified_from_target_snapshot' if classification == 'added' else 'replace_verified_from_target_snapshot'),
        })
    bundle: dict[str, Any] = {
        'schema': BUNDLE_SCHEMA,
        'base_snapshot_id': base_snapshot['snapshot_id'],
        'target_snapshot_id': target_snapshot['snapshot_id'],
        'base_release': base_snapshot['release'],
        'target_release': target_snapshot['release'],
        'digest_algorithm': 'sha256',
        'source_delta_digest_sha256': '' if source_delta is None else source_delta.get('delta_digest_sha256',''),
        'source_delta_path': None if source_delta is None else 'assets/data/deltas/memory-delta-v16.4-to-v16.5.json',
        'actions': actions,
        'dependency_closure': {
            'target_closure_file_count': len(target),
            'payload_file_count': sum(1 for a in actions if a['payload_path']),
            'omitted_unchanged_file_count': unchanged,
            'removed_review_only_count': removed,
            'rule': 'payload contains only added/changed target bytes; unchanged closure members are omitted because their base and target entries are byte-identical; removed paths are never automatically deleted',
        },
        'protected_anchors': sorted(PROTECTED),
        'truth_boundary': 'Local content-addressed repair input only; not authorization to modify protected anchors, not signing/authorship/timestamp authority, and not production deployment or disaster-recovery proof.',
    }
    if source_delta is not None:
        if source_delta.get('from_snapshot_id') != base_snapshot.get('snapshot_id') or source_delta.get('to_snapshot_id') != target_snapshot.get('snapshot_id'):
            raise RepairTransactionError('source_delta_snapshot_mismatch')
        delta_paths = sorted(x['path'] for k in ('added','removed','changed') for x in source_delta.get('classifications',{}).get(k,[]))
        if delta_paths != sorted(a['path'] for a in actions):
            raise RepairTransactionError('source_delta_path_mismatch')
    dg = compute_bundle_digest(bundle)
    bundle['bundle_digest_sha256'] = dg
    bundle['bundle_id'] = 'patch-' + dg[:20]
    return bundle


def verify_bundle(bundle: dict[str, Any], payload_root: Path, base_snapshot: dict[str, Any] | None = None,
                  target_snapshot: dict[str, Any] | None = None, source_delta: dict[str, Any] | None = None) -> dict[str, Any]:
    if bundle.get('schema') != BUNDLE_SCHEMA:
        raise RepairTransactionError('bundle_schema')
    actions = bundle.get('actions') or []
    paths: set[str] = set()
    for a in actions:
        rel = safe_rel(a.get('path', ''))
        if rel in paths:
            raise RepairTransactionError('duplicate_path')
        paths.add(rel)
        protected = rel in PROTECTED
        if bool(a.get('protected_anchor')) != protected:
            raise RepairTransactionError('protected_anchor_flag')
        if protected and a.get('automatic'):
            raise RepairTransactionError('protected_anchor_automatic_mutation')
        cls = a.get('classification')
        if cls not in ('added', 'changed', 'removed'):
            raise RepairTransactionError('classification')
        if cls == 'removed':
            if a.get('automatic') or a.get('payload_path') or a.get('after_sha256'):
                raise RepairTransactionError('automatic_delete_forbidden')
            continue
        payload_path = a.get('payload_path')
        if not isinstance(payload_path, str) or not payload_path.startswith('payload/'):
            raise RepairTransactionError('payload_path')
        payload_rel = safe_rel(payload_path[len('payload/'):])
        if payload_rel != rel:
            raise RepairTransactionError('payload_path_mismatch')
        p = safe_target(payload_root, payload_rel)
        if not p.is_file():
            raise RepairTransactionError('payload_missing:' + rel)
        if p.stat().st_size != a.get('after_size_bytes'):
            raise RepairTransactionError('payload_size:' + rel)
        dg = sha_file(p)
        if dg != a.get('payload_sha256') or dg != a.get('after_sha256'):
            raise RepairTransactionError('payload_hash:' + rel)
    if actions != sorted(actions, key=lambda x: x['path']):
        raise RepairTransactionError('action_order')
    calc = compute_bundle_digest(bundle)
    if calc != bundle.get('bundle_digest_sha256') or bundle.get('bundle_id') != 'patch-' + calc[:20]:
        raise RepairTransactionError('bundle_digest')
    dep = bundle.get('dependency_closure') or {}
    if dep.get('payload_file_count') != sum(1 for a in actions if a.get('payload_path')):
        raise RepairTransactionError('dependency_payload_count')
    if base_snapshot is not None:
        bi = snapshot_index(base_snapshot)
        if bundle.get('base_snapshot_id') != base_snapshot.get('snapshot_id'):
            raise RepairTransactionError('wrong_base_snapshot')
    else:
        bi = None
    if target_snapshot is not None:
        ti = snapshot_index(target_snapshot)
        if bundle.get('target_snapshot_id') != target_snapshot.get('snapshot_id'):
            raise RepairTransactionError('stale_target_snapshot')
    else:
        ti = None
    if source_delta is not None:
        if bundle.get('source_delta_digest_sha256') != source_delta.get('delta_digest_sha256'):
            raise RepairTransactionError('source_delta_digest_mismatch')
        if source_delta.get('from_snapshot_id') != bundle.get('base_snapshot_id') or source_delta.get('to_snapshot_id') != bundle.get('target_snapshot_id'):
            raise RepairTransactionError('source_delta_snapshot_mismatch')
    if bi is not None and ti is not None:
        changed_expected = []
        unchanged = 0
        removed = 0
        for rel in sorted(set(bi) | set(ti)):
            if bi.get(rel) == ti.get(rel):
                unchanged += 1
            else:
                changed_expected.append(rel)
                if rel not in ti:
                    removed += 1
        if sorted(paths) != changed_expected:
            raise RepairTransactionError('bundle_delta_mismatch')
        if dep.get('target_closure_file_count') != len(ti) or dep.get('omitted_unchanged_file_count') != unchanged or dep.get('removed_review_only_count') != removed:
            raise RepairTransactionError('dependency_closure_mismatch')
        for a in actions:
            rel = a['path']; before = bi.get(rel); after = ti.get(rel)
            if a['before_sha256'] != (None if before is None else before['sha256']):
                raise RepairTransactionError('before_hash_mismatch:' + rel)
            if a['after_sha256'] != (None if after is None else after['sha256']):
                raise RepairTransactionError('after_hash_mismatch:' + rel)
            if after is not None and (a['role'] != after['role'] or a['after_size_bytes'] != after['size_bytes']):
                raise RepairTransactionError('target_metadata_mismatch:' + rel)
    return {'pass': True, 'bundle_id': bundle['bundle_id'], 'bundle_digest_sha256': calc, 'action_count': len(actions), 'payload_file_count': dep.get('payload_file_count')}


def verify_directory(root: Path, snapshot: dict[str, Any], scope: set[str] | None = None) -> None:
    idx = snapshot_index(snapshot)
    items = idx.items() if scope is None else ((p, idx[p]) for p in sorted(scope))
    for rel, e in items:
        p = safe_target(root, rel)
        if not p.is_file():
            raise RepairTransactionError('snapshot_missing:' + rel)
        if p.stat().st_size != e['size_bytes'] or sha_file(p) != e['sha256']:
            raise RepairTransactionError('snapshot_hash:' + rel)


def event_material(txid: str, event: dict[str, Any]) -> bytes:
    return ('\0'.join([
        txid,
        str(event['sequence']),
        event['state'],
        event.get('reason', ''),
        event.get('previous_event_digest_sha256', ''),
        event.get('base_snapshot_id', ''),
        event.get('target_snapshot_id', ''),
        event.get('bundle_digest_sha256', ''),
    ]) + '\n').encode('utf-8')


def add_event(journal: dict[str, Any], state: str, reason: str = '') -> dict[str, Any]:
    events = journal['events']
    event = {
        'sequence': len(events) + 1,
        'state': state,
        'reason': reason,
        'previous_event_digest_sha256': events[-1]['event_digest_sha256'] if events else '0' * 64,
        'base_snapshot_id': journal['base_snapshot_id'],
        'target_snapshot_id': journal['target_snapshot_id'],
        'bundle_digest_sha256': journal['bundle_digest_sha256'],
    }
    event['event_digest_sha256'] = sha_bytes(event_material(journal['transaction_id'], event))
    events.append(event)
    return event


def finalize_journal(journal: dict[str, Any]) -> dict[str, Any]:
    material = ''.join(e['event_digest_sha256'] + '\n' for e in journal['events']).encode('ascii')
    journal['journal_digest_sha256'] = sha_bytes(material)
    journal['event_count'] = len(journal['events'])
    journal['terminal_state'] = journal['events'][-1]['state'] if journal['events'] else 'none'
    return journal


def verify_journal(journal: dict[str, Any]) -> dict[str, Any]:
    if journal.get('schema') != JOURNAL_SCHEMA:
        raise RepairTransactionError('journal_schema')
    prev = '0' * 64
    seen_seq = set()
    for i, e in enumerate(journal.get('events') or [], 1):
        if e.get('sequence') != i or i in seen_seq:
            raise RepairTransactionError('journal_sequence')
        seen_seq.add(i)
        if e.get('previous_event_digest_sha256') != prev:
            raise RepairTransactionError('journal_link')
        if e.get('base_snapshot_id') != journal.get('base_snapshot_id') or e.get('target_snapshot_id') != journal.get('target_snapshot_id') or e.get('bundle_digest_sha256') != journal.get('bundle_digest_sha256'):
            raise RepairTransactionError('journal_context')
        calc = sha_bytes(event_material(journal['transaction_id'], e))
        if calc != e.get('event_digest_sha256'):
            raise RepairTransactionError('journal_event_digest')
        prev = calc
    calc_chain = sha_bytes(''.join(e['event_digest_sha256'] + '\n' for e in journal.get('events') or []).encode('ascii'))
    if calc_chain != journal.get('journal_digest_sha256'):
        raise RepairTransactionError('journal_digest')
    if journal.get('event_count') != len(journal.get('events') or []):
        raise RepairTransactionError('journal_count')
    if journal.get('terminal_state') != (journal['events'][-1]['state'] if journal.get('events') else 'none'):
        raise RepairTransactionError('journal_terminal')
    return {'pass': True, 'journal_digest_sha256': calc_chain, 'terminal_state': journal.get('terminal_state')}


def _remove_tree(path: Path) -> None:
    if path.exists():
        if path.is_dir(): shutil.rmtree(path)
        else: path.unlink()


def run_transaction(base_dir: Path, bundle: dict[str, Any], payload_root: Path,
                    base_snapshot: dict[str, Any], target_snapshot: dict[str, Any],
                    failpoint: str | None = None) -> dict[str, Any]:
    base_dir = base_dir.resolve()
    verify_bundle(bundle, payload_root, base_snapshot, target_snapshot)
    verify_directory(base_dir, base_snapshot)
    txid = 'txn-' + bundle['bundle_digest_sha256'][:20]
    journal: dict[str, Any] = {
        'schema': JOURNAL_SCHEMA,
        'transaction_id': txid,
        'bundle_id': bundle['bundle_id'],
        'bundle_digest_sha256': bundle['bundle_digest_sha256'],
        'base_snapshot_id': bundle['base_snapshot_id'],
        'target_snapshot_id': bundle['target_snapshot_id'],
        'events': [],
        'truth_boundary': 'Local transactional byte-state evidence only; not signing/authorship/timestamp authority, semantic truth, live deployment, production disaster recovery, domain control, or real-world authorization.',
    }
    parent = base_dir.parent
    stage = parent / (base_dir.name + '.' + txid + '.stage')
    backup = parent / (base_dir.name + '.' + txid + '.backup')
    _remove_tree(stage); _remove_tree(backup)
    renamed_base = False
    installed_stage = False
    add_event(journal, 'prepare')
    try:
        shutil.copytree(base_dir, stage)
        for a in bundle['actions']:
            if not a['automatic']:
                if a['protected_anchor']:
                    raise RepairTransactionError('protected_anchor_requires_human_authorization')
                continue
            if a['classification'] == 'removed':
                raise RepairTransactionError('automatic_delete_forbidden')
            src = safe_target(payload_root, a['path'])
            dst = safe_target(stage, a['path'])
            dst.parent.mkdir(parents=True, exist_ok=True)
            shutil.copy2(src, dst)
        add_event(journal, 'apply')
        if failpoint == 'after_apply':
            raise RepairTransactionError('simulated_interrupt_after_apply')
        verify_directory(stage, target_snapshot)
        add_event(journal, 'verify')
        if failpoint == 'after_verify':
            raise RepairTransactionError('simulated_interrupt_after_verify')
        os.replace(base_dir, backup); renamed_base = True
        if failpoint in ('after_base_rename', 'rollback_failure'):
            raise RepairTransactionError('simulated_interrupt_after_base_rename')
        os.replace(stage, base_dir); installed_stage = True
        if failpoint == 'after_stage_rename':
            raise RepairTransactionError('simulated_interrupt_after_stage_rename')
        verify_directory(base_dir, target_snapshot)
        add_event(journal, 'commit')
        _remove_tree(backup)
        result = finalize_journal(journal)
        result['exact_target_parity'] = True
        result['exact_base_rollback_parity'] = False
        verify_journal(result)
        return result
    except Exception as exc:
        reason = exc.code if isinstance(exc, RepairTransactionError) else type(exc).__name__
        rollback_ok = False
        try:
            if failpoint == 'rollback_failure':
                # Make rollback verification deterministically fail without hiding the failure.
                if backup.exists():
                    marker = backup / '.rollback-corruption'
                    marker.write_bytes(b'corrupt')
                    victim_rel = next(iter(snapshot_index(base_snapshot)))
                    victim = safe_target(backup, victim_rel)
                    if victim.is_file(): victim.write_bytes(victim.read_bytes() + b'X')
            if installed_stage and base_dir.exists():
                _remove_tree(base_dir)
            if renamed_base and backup.exists():
                os.replace(backup, base_dir)
            _remove_tree(stage)
            verify_directory(base_dir, base_snapshot)
            rollback_ok = True
            add_event(journal, 'rollback', reason)
        except Exception as rollback_exc:
            add_event(journal, 'rollback_failed', reason + ':' + (rollback_exc.code if isinstance(rollback_exc, RepairTransactionError) else type(rollback_exc).__name__))
        result = finalize_journal(journal)
        result['exact_target_parity'] = False
        result['exact_base_rollback_parity'] = rollback_ok
        verify_journal(result)
        return result


def main() -> None:
    ap = argparse.ArgumentParser()
    ap.add_argument('--bundle', required=True)
    ap.add_argument('--payload-root', required=True)
    ap.add_argument('--base-snapshot', required=True)
    ap.add_argument('--target-snapshot', required=True)
    ap.add_argument('--base-dir')
    ap.add_argument('--verify-only', action='store_true')
    ap.add_argument('--failpoint', choices=['after_apply','after_verify','after_base_rename','after_stage_rename','rollback_failure'])
    ap.add_argument('--journal-out')
    args = ap.parse_args()
    bundle = load_json(Path(args.bundle)); base = load_json(Path(args.base_snapshot)); target = load_json(Path(args.target_snapshot)); payload = Path(args.payload_root)
    if args.verify_only:
        out = verify_bundle(bundle, payload, base, target)
    else:
        if not args.base_dir: raise SystemExit('--base-dir required unless --verify-only')
        out = run_transaction(Path(args.base_dir), bundle, payload, base, target, args.failpoint)
    if args.journal_out:
        Path(args.journal_out).write_text(json.dumps(out, ensure_ascii=False, indent=2) + '\n', encoding='utf-8')
    print(json.dumps(out, ensure_ascii=False))

if __name__ == '__main__':
    main()
