#!/usr/bin/env python3
"""Deterministic Merkle tree and inclusion proofs over a memory-state snapshot."""
from pathlib import Path
import argparse, hashlib, json
DOMAIN_LEAF=b'memory-leaf-v1\0'; DOMAIN_NODE=b'memory-node-v1\0'
PROTECTED={'.uai/totem.uai','.uai/taboo.uai','.uai/talisman.uai'}
class MerkleError(ValueError): pass
def h(b): return hashlib.sha256(b).hexdigest()
def leaf_hash(e):
    mat=(e['path']+'\0'+e['role']+'\0'+str(e['size_bytes'])+'\0'+e['sha256']).encode('utf-8')
    return h(DOMAIN_LEAF+mat)
def node_hash(left,right): return h(DOMAIN_NODE+bytes.fromhex(left)+bytes.fromhex(right))
def validate_entries(entries):
    if entries!=sorted(entries,key=lambda e:e['path']): raise MerkleError('entry_order')
    paths=[e['path'] for e in entries]
    if len(paths)!=len(set(paths)): raise MerkleError('duplicate_leaf')
    for e in entries:
        if not isinstance(e.get('sha256'),str) or len(e['sha256'])!=64: raise MerkleError('bad_source_hash')
    return entries
def build_levels(entries):
    validate_entries(entries)
    if not entries: raise MerkleError('empty_tree')
    levels=[[leaf_hash(e) for e in entries]]
    while len(levels[-1])>1:
        cur=levels[-1]; nxt=[]
        for i in range(0,len(cur),2):
            left=cur[i]; right=cur[i+1] if i+1<len(cur) else cur[i]
            nxt.append(node_hash(left,right))
        levels.append(nxt)
    return levels
def proof_for(entries,levels,index):
    proof=[]; idx=index
    for level_no,level in enumerate(levels[:-1]):
        if idx%2==0:
            sib=idx+1 if idx+1<len(level) else idx; side='right'
        else: sib=idx-1; side='left'
        proof.append({'level':level_no,'side':side,'sha256':level[sib]})
        idx//=2
    return proof
def verify_proof(entry,proof,root):
    cur=leaf_hash(entry)
    for expected_level,p in enumerate(proof):
        if p.get('level')!=expected_level or p.get('side') not in {'left','right'}: raise MerkleError('proof_structure')
        sib=p.get('sha256','')
        if len(sib)!=64: raise MerkleError('proof_hash')
        cur=node_hash(sib,cur) if p['side']=='left' else node_hash(cur,sib)
    return cur==root
def build(snapshot):
    entries=validate_entries(snapshot['files']); levels=build_levels(entries); root=levels[-1][0]
    target_paths=[e['path'] for e in entries if e['role']=='active_memory' or e['path'] in PROTECTED]
    proofs={}
    by={e['path']:e for e in entries}
    indices={e['path']:i for i,e in enumerate(entries)}
    for path in sorted(set(target_paths)):
        proof=proof_for(entries,levels,indices[path])
        if not verify_proof(by[path],proof,root): raise MerkleError('self_verify:'+path)
        proofs[path]={'entry':by[path],'leaf_sha256':leaf_hash(by[path]),'protected_anchor':path in PROTECTED,'proof':proof}
    return {
      'schema':'memory-merkle-tree/1','release':snapshot['release'],'snapshot_id':snapshot['snapshot_id'],
      'snapshot_digest_sha256':snapshot['snapshot_digest_sha256'],'leaf_count':len(entries),'tree_height':len(levels)-1,
      'leaf_domain':'memory-leaf-v1','node_domain':'memory-node-v1','odd_node_rule':'duplicate_last_hash_at_each_level',
      'merkle_root_sha256':root,'proof_count':len(proofs),'protected_anchor_count':sum(1 for p in proofs.values() if p['protected_anchor']),
      'proofs':proofs,
      'truth_boundary':'Local inclusion/integrity proof over the declared snapshot only; not signing, authorship, timestamp authority, public witness consensus, semantic truth, deployment, domain control, or real-world authorization.'
    }
def verify_tree(snapshot,tree):
    rebuilt=build(snapshot)
    if tree['snapshot_id']!=snapshot['snapshot_id'] or tree['leaf_count']!=len(snapshot['files']): raise MerkleError('snapshot_link')
    if tree['merkle_root_sha256']!=rebuilt['merkle_root_sha256']: raise MerkleError('root_mismatch')
    by={e['path']:e for e in snapshot['files']}
    for path,rec in tree.get('proofs',{}).items():
        if path not in by or rec.get('entry')!=by[path]: raise MerkleError('path_substitution:'+path)
        if not verify_proof(by[path],rec.get('proof',[]),tree['merkle_root_sha256']): raise MerkleError('invalid_proof:'+path)
    expected={e['path'] for e in snapshot['files'] if e['role']=='active_memory' or e['path'] in PROTECTED}
    if set(tree.get('proofs',{}))!=expected: raise MerkleError('proof_set')
    return True
def main():
    ap=argparse.ArgumentParser(); ap.add_argument('--root',default=str(Path(__file__).resolve().parents[1])); ap.add_argument('--snapshot',default='assets/data/memory-state-snapshot.json'); ap.add_argument('--out',default='assets/data/memory-merkle-tree.json'); ap.add_argument('--verify',action='store_true'); args=ap.parse_args()
    root=Path(args.root).resolve(); snap=json.loads((root/args.snapshot).read_text(encoding='utf-8')); outp=root/args.out
    if args.verify:
        tree=json.loads(outp.read_text(encoding='utf-8')); verify_tree(snap,tree); print(json.dumps({'pass':True,'root':tree['merkle_root_sha256'],'proofs':len(tree['proofs'])})); return
    tree=build(snap); outp.parent.mkdir(parents=True,exist_ok=True); outp.write_text(json.dumps(tree,ensure_ascii=False,indent=2)+'\n',encoding='utf-8'); print(json.dumps({'root':tree['merkle_root_sha256'],'leaves':tree['leaf_count'],'proofs':tree['proof_count']},ensure_ascii=False))
if __name__=='__main__': main()
