"""Pure total integer expressions; full-scan e-graph rebuild and tree extraction.
No file output. Costs charge every tree occurrence, not each shared DAG node.
"""
from dataclasses import dataclass
import json
SIG={'x':0,'y':0,'0':0,'1':0,'2':0,'add':2,'mul':2,'dbl':1}
WEIGHT={'x':1,'y':1,'0':1,'1':1,'2':1,'add':2,'mul':4,'dbl':1}
def need(ok,msg):
 if not ok:raise ValueError(msg)
def syntax(t,pattern=False):
 todo=[t];variables=set();count=0
 while todo:
  a=todo.pop();count+=1
  if pattern and type(a)is str and a.startswith('?') and len(a)>1:variables.add(a);continue
  need(type(a)is tuple and a and type(a[0])is str and a[0]in SIG and len(a)==SIG[a[0]]+1,'expression/pattern syntax');todo.extend(a[1:])
 return variables,count
@dataclass(frozen=True)
class Rule:
 name:str
 lhs:object
 rhs:object
 def validate(self):
  l,_=syntax(self.lhs,True);r,_=syntax(self.rhs,True);need(r<=l,'unbound RHS placeholder');return self
class EGraph:
 def __init__(self):self.nodes=[];self.parent=[];self.size=[];self.clean=True;self.version=0
 def find(self,i):
  need(type(i)is int and 0<=i<len(self.parent),'e-class id')
  while self.parent[i]!=i:i=self.parent[i]
  return i
 def merge(self,a,b):
  a,b=self.find(a),self.find(b)
  if a==b:return a
  if (self.size[a],-a)<(self.size[b],-b):a,b=b,a
  self.parent[b]=a;self.size[a]+=self.size[b];self.clean=False;self.version+=1;return a
 def canon(self,n):return (n[0],tuple(self.find(c)for c in n[1]))
 def add(self,op,children):
  need(op in SIG and len(children)==SIG[op],'enode arity');n=self.canon((op,tuple(children)))
  # A deliberately simple search remains valid even during a dirty write phase.
  for i,old in enumerate(self.nodes):
   if self.canon(old)==n:return self.find(i)
  i=len(self.nodes);self.nodes.append(n);self.parent.append(i);self.size.append(1);self.version+=1;return i
 def add_term(self,t,subst=None):
  syntax(t,subst is not None);subst={}if subst is None else subst;todo=[(t,False)];values=[]
  while todo:
   a,ready=todo.pop()
   if type(a)is str:need(a in subst,'missing substitution');values.append(self.find(subst[a]));continue
   if not ready:todo.append((a,True));todo.extend((c,False)for c in reversed(a[1:]));continue
   k=len(a)-1;children=values[-k:]if k else[]
   if k:del values[-k:]
   values.append(self.add(a[0],children))
  return values[0]
 def rebuild(self):
  passes=0;merges=0
  while True:
   passes+=1;table={};changed=False
   for i,n in enumerate(self.nodes):
    key=self.canon(n);owner=self.find(i)
    if key in table:
     other=self.find(table[key])
     if owner!=other:self.merge(owner,other);merges+=1;changed=True
    else:table[key]=owner
   if not changed:break
  self.clean=True;return dict(passes=passes,merges=merges)
 def snapshot(self):
  need(self.clean,'rebuild before query');groups={}
  for i,n in enumerate(self.nodes):groups.setdefault(self.find(i),{})[self.canon(n)]=None
  return {c:tuple(ns)for c,ns in groups.items()}
 def frozen(self,root,weights=WEIGHT):
  groups=self.snapshot();ids={c:i for i,c in enumerate(groups)};out=[]
  for c,ns in groups.items():out.append(tuple((op,tuple(ids[k]for k in kids),weights[op])for op,kids in ns))
  return tuple(out),ids[self.find(root)]
def matches(pattern,root,groups):
 # Pattern work decreases in each branch, even if the e-graph has cycles.
 todo=[([(pattern,root)],{})];out=[];steps=0
 while todo:
  tasks,env=todo.pop();steps+=1
  if not tasks:out.append(env);continue
  p,c=tasks[-1];rest=tasks[:-1]
  if type(p)is str:
   if p not in env or env[p]==c:new=env.copy();new[p]=c;todo.append((rest,new))
   continue
  for op,kids in groups[c]:
   steps+=1
   if op==p[0]:todo.append((rest+list(zip(p[1:],kids)),env.copy()))
 # Different e-node witnesses may induce the same placeholder assignment.
 seen=set();unique=[]
 for env in out:
  key=tuple(sorted(env.items()))
  if key not in seen:seen.add(key);unique.append(env)
 return unique,steps
def saturate(t,rules,round_limit):
 need(type(round_limit)is int and round_limit>=0,'nonnegative round limit')
 for r in rules:r.validate()
 eg=EGraph();root=eg.add_term(t);eg.rebuild();history=[]
 for iteration in range(round_limit):
  groups=eg.snapshot();batch=[];work=0
  for rule in rules:
   for c in groups:
    ms,n=matches(rule.lhs,c,groups);work+=n;batch.extend((rule,c,s)for s in ms)
  before=eg.version;created=len(eg.nodes);applied=[]
  for rule,c,s in batch:
   rhs=eg.add_term(rule.rhs,s);eg.merge(c,rhs);applied.append(dict(rule=rule.name,root=c,subst=s))
  restoration=eg.rebuild();history.append(dict(iteration=iteration+1,matches=len(batch),match_steps=work,new_nodes=len(eg.nodes)-created,classes=len(eg.snapshot()),rebuild=restoration,applications=applied))
  if eg.version==before:return eg,root,dict(status='saturated',history=history)
 return eg,root,dict(status='iteration-limit',history=history)
def graph_valid(g):
 need(type(g)is tuple and g,'nonempty finite e-graph');seen={};arities={}
 for c,nodes in enumerate(g):
  need(type(nodes)is tuple and nodes,'nonempty e-class')
  for n in nodes:
   need(type(n)is tuple and len(n)==3,'cost enode');op,kids,w=n
   need(type(op)is str and op and type(kids)is tuple and type(w)is int and w>0,'positive integral node cost')
   need(all(type(k)is int and 0<=k<len(g)for k in kids),'child class')
   need(op not in arities or arities[op]==len(kids),'fixed symbol arity');arities[op]=len(kids)
   key=(op,kids);need(key not in seen,'deduplicated canonical enodes');seen[key]=c
 return g
def candidate(n,dist):
 _,children,w=n
 return None if any(dist[c]is None for c in children)else w+sum(dist[c]for c in children)
def extract(g,root):
 graph_valid(g);need(type(root)is int and 0<=root<len(g),'extraction root');C=len(g);dist=[None]*C;trace=[]
 for height in range(1,C+1):
  nxt=[]
  for ns in g:
   costs=[v for n in ns if (v:=candidate(n,dist))is not None];nxt.append(min(costs)if costs else None)
  trace.append(nxt.copy());dist=nxt
 chosen=[]
 for c,ns in enumerate(g):
  if dist[c]is None:chosen.append(None)
  else:chosen.append(next(i for i,n in enumerate(ns)if candidate(n,dist)==dist[c]))
 proof=dict(dist=dist,chosen=chosen);verify(g,proof)
 return dict(status='NO_FINITE_TERM'if dist[root]is None else'OK',root=root,cost=dist[root],proof=proof,height_rows=trace)
def verify(g,proof):
 graph_valid(g);dist=proof['dist'];chosen=proof['chosen'];need(len(dist)==len(g)==len(chosen),'certificate shape')
 need(all(d is None or(type(d)is int and d>0)for d in dist),'distance')
 for c,ns in enumerate(g):
  d=dist[c]
  for n in ns:
   value=candidate(n,dist)
   if value is not None:need(d is not None and d<=value,'lower-bound inequality')
  if d is None:need(chosen[c]is None,'unproductive choice');continue
  i=chosen[c];need(type(i)is int and 0<=i<len(ns),'chosen enode');need(candidate(ns[i],dist)==d,'realizing choice')
  need(all(dist[k]<d for k in ns[i][1]),'strict descent')
 return True
def unfold(g,result):
 need(result['status']=='OK','no finite term');verify(g,result['proof']);todo=[(result['root'],False)];values=[];choice=result['proof']['chosen']
 while todo:
  c,ready=todo.pop();op,kids,_=g[c][choice[c]]
  if not ready:todo.append((c,True));todo.extend((k,False)for k in reversed(kids));continue
  k=len(kids);children=values[-k:]if k else[]
  if k:del values[-k:]
  values.append((op,*children))
 return values[0]
def evaluate(t,env):
 syntax(t);todo=[(t,False)];values=[]
 while todo:
  a,ready=todo.pop();op=a[0]
  if not ready:todo.append((a,True));todo.extend((c,False)for c in reversed(a[1:]));continue
  k=len(a)-1;v=values[-k:]if k else[]
  if k:del values[-k:]
  if op in ('x','y'):z=env[op]
  elif op in ('0','1','2'):z=int(op)
  elif op=='add':z=v[0]+v[1]
  elif op=='mul':z=v[0]*v[1]
  else:z=2*v[0]
  values.append(z)
 return values[0]
def tree_cost(t,weights=WEIGHT):
 todo=[t];total=0
 while todo:a=todo.pop();total+=weights[a[0]];todo.extend(a[1:])
 return total
X=('x',);Y=('y',);ZERO=('0',);ONE=('1',);TWO=('2',)
RULES=(Rule('add-zero',('add','?a',ZERO),'?a'),Rule('mul-one',('mul','?a',ONE),'?a'),Rule('strength',('mul','?a',TWO),('dbl','?a')),Rule('factor',('add',('mul','?a','?c'),('mul','?b','?c')),('mul',('add','?a','?b'),'?c')))
SOURCE=('add',('mul',('add',X,ZERO),TWO),('mul',Y,TWO))
def regressions():
 rejected=[]
 def reject(label,fn):
  try:fn()
  except ValueError:rejected.append(label);return
  raise RuntimeError('expected rejection '+label)
 e=EGraph();x=e.add_term(X);a=e.add_term(('add',X,ZERO));p=e.add_term(('dbl',('add',X,ZERO)));q=e.add_term(('dbl',X));pp=e.add_term(('mul',('dbl',('add',X,ZERO)),TWO));qq=e.add_term(('mul',('dbl',X),TWO));e.merge(a,x)
 reject('dirty query',e.snapshot);need(e.find(p)!=e.find(q),'parents pending before rebuild');repair=e.rebuild();need(e.find(p)==e.find(q)and e.find(pp)==e.find(qq),'congruence propagation')
 for label,g in [('zero weight',((('z',(),0),),)),('negative weight',((('z',(),-1),),)),('missing child',((('u',(9,),1),),)),('empty class',((),)),('duplicate signature',((('z',(),1),),(('z',(),1),))),('inconsistent arity',((('u',(),1),('u',(0,),1)),))]:reject(label,lambda g=g:extract(g,0))
 fixture=((('x',(),1),('u',(0,),1)),);good=extract(fixture,0);over={'dist':[2],'chosen':[0]};under={'dist':[1],'chosen':[1]};wrong_none={'dist':[None],'chosen':[None]}
 for label,proof in [('overpriced certificate',over),('unrealized lower bound',under),('false nonproductive',wrong_none),('nonnumeric distance',{'dist':['1'],'chosen':[0]})]:reject(label,lambda proof=proof:verify(fixture,proof))
 reject('unbound RHS rule',lambda:Rule('bad',X,'?a').validate())
 groups={0:(('x',()),),1:(('y',()),),2:(('add',(0,1)),),3:(('add',(0,0)),)}
 same=('add','?v','?v');need(not matches(same,2,groups)[0] and matches(same,3,groups)[0]==[{'?v':0}],'repeated placeholder agrees')
 return dict(rejections=rejected,rebuild=repair,repeated_placeholder=True)
def main():
 eg,root,run=saturate(SOURCE,RULES,10);g,r=eg.frozen(root);answer=extract(g,r);term=unfold(g,answer);need(run['status']=='saturated' and answer['cost']==5,'main optimum');need(tree_cost(SOURCE)==17 and tree_cost(term)==5,'tree occurrence charge')
 weak,wr,wrun=saturate(SOURCE,RULES[:-1],10);wg,wri=weak.frozen(wr);wa=extract(wg,wri);need(wa['cost']==6,'without factoring')
 short,sr,srun=saturate(SOURCE,RULES,1);sg,sri=short.frozen(sr);short_cost=extract(sg,sri)['cost'];need(srun['status']=='iteration-limit' and short_cost==6,'one round not saturation')
 empty,er,erun=saturate(SOURCE,RULES,0);eg0,eri=empty.frozen(er);need(erun['status']=='iteration-limit' and extract(eg0,eri)['cost']==17,'zero budget keeps clean initial graph')
 observations=[]
 for x in range(-4,5):
  for y in range(-4,5):need(evaluate(SOURCE,dict(x=x,y=y))==evaluate(term,dict(x=x,y=y)),'numeric crosscheck')
 for x,y in [(-3,5),(0,0),(7,-2)]:observations.append(dict(x=x,y=y,source=evaluate(SOURCE,dict(x=x,y=y)),target=evaluate(term,dict(x=x,y=y))))
 # A self-cycle with a leaf is productive; without a leaf there is no finite term.
 cyclic=((('loop',(0,),1),('leaf',(),3)),);pure=((('loop',(0,),1),),);ca=extract(cyclic,0);pa=extract(pure,0);need(ca['cost']==3 and pa['status']=='NO_FINITE_TERM','cycle productivity')
 # Tree charging versus shared-DAG charging; premises a=g(q), b=h(q).
 dag=((('pair',(1,2),1),),(('a',(),6),('g',(3,),1)),(('b',(),6),('h',(3,),1)),(('q',(),6),));da=extract(dag,0);need(da['cost']==13 and 1+1+1+6==9,'DAG sharing boundary')
 # A semantically valid rule can keep creating new argument classes forever.
 growth=(Rule('zero-product',('mul','?a',ZERO),('mul',('add','?a',ONE),ZERO)),);ge,gr,grun=saturate(('mul',X,ZERO),growth,3);need(grun['status']=='iteration-limit','unbounded growth not saturation')
 print(json.dumps(dict(status='PASS',checks=regressions(),source=SOURCE,source_cost=17,result=term,run=run,graph=g,extraction=answer,without_factoring=dict(run=wrun,cost=wa['cost']),one_round=srun,one_round_cost=short_cost,zero_round=erun,observations=observations,cycles=dict(productive=ca,unproductive=pa),tree_vs_dag=dict(tree=da,shared_dag_cost=9),growth=grun),ensure_ascii=False,indent=2))
if __name__=='__main__':main()
