"""Complete structural irrRAF catalogues via deterministic supplier resolutions.

Catalysts do NOT gate food closure. Catalysis is disjunctive. This is a finite
structural calculation, not a kinetic or thermodynamic reactor simulation.
"""
from __future__ import annotations
import argparse
import csv
from dataclasses import dataclass, asdict
from fractions import Fraction
import hashlib
from itertools import combinations, product
import json
import math
from pathlib import Path
import platform
import time

# USER INPUTS ---------------------------------------------------------------
FOOD = ('f',)                     # species freely available to every reaction
# id, reactants, products, alternative catalysts; paper's k=2 gate example.
REACTIONS = (
    (0, ('f',), ('x1',), ('z2',)),
    (1, ('f',), ('x1',), ('z2',)),
    (2, ('f',), ('x2',), ('z2',)),
    (3, ('f',), ('x2',), ('z2',)),
    (4, ('x1',), ('z1',), ('z2',)),
    (5, ('z1', 'x2'), ('z2',), ('z2',)),
)
DELETION_COSTS = {0: 2, 1: 3, 2: 4, 3: 1, 4: 7, 5: 6}  # positive abstract costs
AVAILABLE_PROBABILITY = '1/2'     # independent static reaction availability
REQUIRED_REACTIONS = (0,)
FORBIDDEN_REACTIONS = (2,)
DELETED_REACTIONS = (0,)
SCALING_SIZES = (1, 2, 3, 4, 5, 6)  # paper families; controlled size sweep
LARGE_INDEPENDENT_MODULES = 64
DETERMINISTIC_CHAIN_LENGTH = 64
RESOLUTION_BUDGET = 4096          # fail explicitly before exceeding this budget
EXHAUSTIVE_REACTION_CAP = 18       # small-instance verification/queries only
PAPER_SHA256 = 'c92ad401d52719b68ab0118296ae5dc686c994cbaf9bea19a5bb8c4c4820523a'
# All inputs are finite identifiers/sets, counts, costs or probabilities;
# there are no concentration units or rate constants in the RAF predicate.
# --------------------------------------------------------------------------


@dataclass(frozen=True)
class Reaction:
    identifier: int
    reactants: frozenset[str]
    products: frozenset[str]
    catalysts: frozenset[str]

    def __post_init__(self):
        if not isinstance(self.identifier,int) or isinstance(self.identifier,bool):
            raise ValueError('Reaction identifiers must be integers.')
        for field in ('reactants','products','catalysts'):
            values=frozenset(getattr(self,field))
            if any(not isinstance(x,str) or not x for x in values):
                raise ValueError('Species names must be nonempty strings.')
            object.__setattr__(self,field,values)


class ReactionSystem:
    def __init__(self,food,reactions):
        self.food=frozenset(food)
        if any(not isinstance(x,str) or not x for x in self.food):
            raise ValueError('Food species names must be nonempty strings.')
        reactions=tuple(reactions)
        self.reactions={r.identifier:r for r in reactions}
        if len(self.reactions)!=len(reactions): raise ValueError('Reaction identifiers must be unique.')
        self.identifiers=frozenset(self.reactions)

    def selected(self,ids=None):
        chosen=self.identifiers if ids is None else frozenset(ids)
        if not chosen<=self.identifiers: raise ValueError('Unknown reaction identifier.')
        return chosen

    def closure_stages(self,ids=None):
        selected=self.selected(ids); stages=[self.food]
        while True:
            current=stages[-1]
            new=current.union(*(self.reactions[r].products for r in selected
                                if self.reactions[r].reactants<=current))
            if new==current: return stages
            stages.append(new)

    def closure(self,ids=None): return self.closure_stages(ids)[-1]

    def is_raf(self,ids):
        ids=self.selected(ids);closure=self.closure(ids)
        return bool(ids) and all(self.reactions[r].reactants<=closure
                                and self.reactions[r].catalysts&closure for r in ids)

    def maximal_raf(self,ids=None):
        current=self.selected(ids)
        while current:
            closure=self.closure(current)
            remaining=frozenset(r for r in current if self.reactions[r].reactants<=closure
                                and self.reactions[r].catalysts&closure)
            if remaining==current: return current
            current=remaining
        return frozenset()

    def restrict(self,ids):
        return ReactionSystem(self.food,[self.reactions[r] for r in sorted(self.selected(ids))])

    def record(self):
        return dict(food=sorted(self.food),reactions=[dict(id=r.identifier,reactants=sorted(r.reactants),
                products=sorted(r.products),catalysts=sorted(r.catalysts)) for r in self.reactions.values()])

    @classmethod
    def from_record(cls,data):
        return cls(data['food'],[Reaction(r['id'],r['reactants'],r['products'],r['catalysts']) for r in data['reactions']])


@dataclass
class SupplierOptions:
    maximal: frozenset[int]
    producers: dict[str,tuple[int,...]]
    catalysts: dict[int,tuple[str,...]]

    @classmethod
    def from_system(cls,system):
        maximal=system.maximal_raf()
        available=system.food.union(*(system.reactions[r].products for r in maximal))
        producers={x:tuple(sorted(r for r in maximal if x in system.reactions[r].products))
                   for x in sorted(available-system.food)}
        catalysts={}
        for r in sorted(maximal):
            food=system.reactions[r].catalysts&system.food
            catalysts[r]=(min(food),) if food else tuple(sorted(system.reactions[r].catalysts&available))
        return cls(maximal,producers,catalysts)

    @property
    def beta(self): return sum(len(v)-1 for v in (*self.producers.values(),*self.catalysts.values()))

    @property
    def count(self): return math.prod(len(v) for v in (*self.producers.values(),*self.catalysts.values()))

    def resolutions(self):
        keys=list(self.producers)+list(self.catalysts)
        options=list(self.producers.values())+list(self.catalysts.values())
        boundary=len(self.producers)
        for values in product(*options):
            yield dict(zip(keys[:boundary],values[:boundary])),dict(zip(keys[boundary:],values[boundary:]))


def resolve(system,options,producer,catalyst):
    return ReactionSystem(system.food,[Reaction(r,system.reactions[r].reactants,
        frozenset(x for x in system.reactions[r].products if x in system.food or producer[x]==r),
        frozenset((catalyst[r],))) for r in sorted(options.maximal)])


def sink_components(graph):
    """Iterative Kosaraju SCCs, avoiding recursion-depth limits on long chains."""
    visited=set();order=[]
    for root in sorted(graph):
        if root in visited: continue
        visited.add(root); stack=[(root,iter(sorted(graph[root])))]
        while stack:
            node,children=stack[-1]
            child=next(children,None)
            if child is None: order.append(node);stack.pop()
            elif child not in visited:
                visited.add(child);stack.append((child,iter(sorted(graph[child]))))
    reverse={v:set() for v in graph}
    for v,edges in graph.items():
        for w in edges: reverse[w].add(v)
    visited=set();components=[]
    for root in reversed(order):
        if root in visited: continue
        component=set();stack=[root];visited.add(root)
        while stack:
            node=stack.pop();component.add(node)
            for child in reverse[node]-visited:
                visited.add(child);stack.append(child)
        components.append(frozenset(component))
    return [c for c in components if all(graph[v]<=c for v in c)]


@dataclass
class Catalogue:
    members: tuple[frozenset[int],...]
    beta: int
    global_resolution_count: int
    resolutions_inspected: int
    candidates: int
    rejected: int
    validation_calls: int
    component_count: int
    largest_local_beta: int

    def filter(self,required=(),forbidden=()):
        required,forbidden=frozenset(required),frozenset(forbidden)
        if required&forbidden: raise ValueError('Required and forbidden identifiers overlap.')
        return tuple(c for c in self.members if required<=c and not forbidden&c)

    def frequencies(self):
        return {r:sum(r in c for c in self.members) for r in sorted(set().union(*self.members))}


class SupplierEnumerator:
    def __init__(self,resolution_budget=RESOLUTION_BUDGET):
        if not isinstance(resolution_budget,int) or resolution_budget<1:
            raise ValueError('Resolution budget must be a positive integer.')
        self.budget=resolution_budget

    def global_catalogue(self,system):
        options=SupplierOptions.from_system(system)
        if options.count>self.budget:
            raise ValueError(f'{options.count} resolutions exceed budget {self.budget}; use components or raise budget explicitly.')
        members=set();candidates=rejected=calls=inspected=0
        if options.maximal:
            for producer,catalyst in options.resolutions():
                inspected+=1
                resolved=resolve(system,options,producer,catalyst)
                kept=resolved.maximal_raf()
                graph={r:{producer[x] for x in (resolved.reactions[r].reactants|resolved.reactions[r].catalysts)-system.food}
                       for r in kept}
                if any(not edges<=kept for edges in graph.values()): raise ArithmeticError('Dependency escaped pruned system.')
                resolution_calls=0
                for component in sink_components(graph):
                    candidates+=1;valid=True
                    for r in sorted(component):
                        calls+=1;resolution_calls+=1
                        if system.maximal_raf(component-{r}): valid=False;break
                    if valid: members.add(component)
                    else: rejected+=1
                if resolution_calls>len(options.maximal): raise ArithmeticError('Validation budget invariant failed.')
        ordered=tuple(sorted(members,key=lambda c:tuple(sorted(c))))
        result=Catalogue(ordered,options.beta,options.count,inspected,candidates,rejected,calls,
                         1 if options.maximal else 0,options.beta)
        if sum(map(len,ordered))>len(options.maximal)*options.count: raise ArithmeticError('Output membership bound failed.')
        return result

    def components(self,system):
        options=SupplierOptions.from_system(system)
        # Every nonfood reactant, product and effective catalyst incidence counts.
        touching={}
        for r in options.maximal:
            reaction=system.reactions[r]
            for x in (reaction.reactants|reaction.products|frozenset(options.catalysts[r]))-system.food:
                touching.setdefault(x,set()).add(r)
        adjacency={r:set() for r in options.maximal}
        for group in touching.values():
            first=min(group)
            for r in group: adjacency[first].add(r);adjacency[r].add(first)
        blocks=[];seen=set()
        for r in sorted(adjacency):
            if r in seen: continue
            block=set();stack=[r];seen.add(r)
            while stack:
                v=stack.pop();block.add(v)
                for w in adjacency[v]-seen: seen.add(w);stack.append(w)
            blocks.append(frozenset(block))
        return blocks

    def catalogue(self,system):
        options=SupplierOptions.from_system(system);blocks=self.components(system)
        local=[system.restrict(b) for b in blocks]
        required=sum(SupplierOptions.from_system(q).count for q in local)
        if required>self.budget: raise ValueError(f'Total local resolutions {required} exceed budget {self.budget}.')
        results=[self.global_catalogue(q) for q in local]
        return Catalogue(tuple(sorted((c for result in results for c in result.members),key=lambda c:tuple(sorted(c)))),
                         options.beta,options.count,sum(r.resolutions_inspected for r in results),
                         sum(r.candidates for r in results),sum(r.rejected for r in results),
                         sum(r.validation_calls for r in results),len(blocks),max((r.beta for r in results),default=0))


def coupled_gates(k):
    if not isinstance(k,int) or k<1: raise ValueError('Gate count must be a positive integer.')
    reactions=[]
    for i in range(k):
        for alternative in range(2):
            reactions.append(Reaction(2*i+alternative,{'f'},{f'x{i+1}'},{f'z{k}'}))
    for i in range(k):
        reactants={f'x{i+1}'}|({f'z{i}'} if i else set())
        reactions.append(Reaction(2*k+i,reactants,{f'z{i+1}'},{f'z{k}'}))
    return ReactionSystem({'f'},reactions)


def independent_modules(count,choices=True):
    if not isinstance(count,int) or count<1: raise ValueError('Module count must be positive integer.')
    reactions=[];per=3 if choices else 2
    for i in range(count):
        for a in range(per-1): reactions.append(Reaction(per*i+a,{'f'},{f'x{i}'},{f'z{i}'}))
        reactions.append(Reaction(per*i+per-1,{f'x{i}'},{f'z{i}'},{f'z{i}'}))
    return ReactionSystem({'f'},reactions)


def deterministic_chain(length):
    if not isinstance(length,int) or length<1: raise ValueError('Chain length must be positive integer.')
    return ReactionSystem({'f'},[Reaction(i,{'f' if i==0 else f'x{i}'},{f'x{i+1}'},{f'x{length}'}) for i in range(length)])


def minimal_cuts(system,catalogue):
    if len(system.identifiers)>EXHAUSTIVE_REACTION_CAP: raise ValueError('Exact cut search exceeds subset budget.')
    found=[]
    for size in range(len(system.identifiers)+1):
        for values in combinations(sorted(system.identifiers),size):
            cut=frozenset(values)
            if not any(c<=cut for c in found) and all(cut&member for member in catalogue.members): found.append(cut)
    return tuple(found)


def greedy_cut(catalogue,costs):
    costs={r:Fraction(str(c)) for r,c in costs.items()}
    if any(c<0 for c in costs.values()): raise ValueError('Deletion costs must be nonnegative.')
    remaining=set(catalogue.members);chosen=[]
    while remaining:
        options=[(cost/sum(r in c for c in remaining),r) for r,cost in costs.items() if any(r in c for c in remaining)]
        if not options: raise ValueError('Eligible reactions cannot hit all irrRAFs.')
        _,r=min(options);chosen.append(r);remaining={c for c in remaining if r not in c}
    return chosen,sum((costs[r] for r in chosen),Fraction(0))


def exact_availability(system,catalogue,p=Fraction(AVAILABLE_PROBABILITY)):
    p=Fraction(p)
    if not 0<=p<=1: raise ValueError('Availability must lie in [0,1].')
    n=len(system.identifiers)
    if n>EXHAUSTIVE_REACTION_CAP: raise ValueError('Exact availability search exceeds subset budget.')
    total=Fraction(0)
    for size in range(n+1):
        for ids in combinations(sorted(system.identifiers),size):
            if any(c<=frozenset(ids) for c in catalogue.members): total+=p**size*(1-p)**(n-size)
    return total


def zero_excess_reliability(catalogue,p=Fraction(AVAILABLE_PROBABILITY)):
    p=Fraction(p)
    if catalogue.beta!=0: raise ValueError('Disjoint zero-excess formula does not apply to this catalogue.')
    if not 0<=p<=1: raise ValueError('Availability must lie in [0,1].')
    return 1-math.prod(1-p**len(c) for c in catalogue.members)


def family_record(result):
    return dict(members=[sorted(c) for c in result.members],beta=result.beta,
                global_resolution_count=result.global_resolution_count,resolutions_inspected=result.resolutions_inspected,
                candidates=result.candidates,rejected=result.rejected,validation_calls=result.validation_calls,
                component_count=result.component_count,largest_local_beta=result.largest_local_beta,
                frequencies=result.frequencies(),total_memberships=sum(map(len,result.members)))


def figures(system,catalogue,scaling,output):
    import matplotlib
    matplotlib.use('Agg')
    import matplotlib.pyplot as plt
    with plt.rc_context({'font.size':11,'axes.spines.top':False,'axes.spines.right':False,
                         'figure.facecolor':'white','savefig.facecolor':'white','svg.fonttype':'none'}):
        if catalogue.members and len(catalogue.members)<=32 and len(system.identifiers)<=30:
            fig,ax=plt.subplots(figsize=(7.2,4.5),layout='constrained')
            ids=sorted(system.identifiers)
            ax.imshow([[int(r in c) for r in ids] for c in catalogue.members],cmap='Blues',vmin=0,vmax=1,aspect='auto',interpolation='nearest')
            ax.set(xticks=range(len(ids)),xticklabels=ids,yticks=range(len(catalogue.members)),
                   yticklabels=[f'I{i+1}' for i in range(len(catalogue.members))],
                   xlabel='Original reaction identifier',ylabel='Irreducible RAF')
            for i,c in enumerate(catalogue.members):
                for j,r in enumerate(ids):
                    ax.text(j,i,'1' if r in c else '0',ha='center',va='center',color='white' if r in c else '#202124')
            fig.savefig(output/'catalogue.png',dpi=220);fig.savefig(output/'catalogue.svg');plt.close(fig)
        fig,ax=plt.subplots(figsize=(7.2,4.5),layout='constrained')
        for family,mode,color,style in [('independent','global','#D55E00','--'),('independent','components','#0072B2','-'),('coupled','components','#7B3294',':')]:
            rows=[r for r in scaling if r['family']==family and r['mode']==mode]
            ax.plot([r['size'] for r in rows],[r['resolutions_inspected'] for r in rows],marker='o',linestyle=style,color=color,label=f'{family.capitalize()}, {mode}')
        ax.set(xlabel='Number of modules or gates',ylabel='Supplier resolutions inspected',yscale='log',xticks=list(SCALING_SIZES))
        ax.legend(fontsize=9);ax.grid(axis='y',color='#e3e6e8')
        fig.savefig(output/'scaling.png',dpi=220);fig.savefig(output/'scaling.svg');plt.close(fig)


def main():
    parser=argparse.ArgumentParser(description=__doc__)
    parser.add_argument('--output',type=Path,default=Path('outputs'))
    parser.add_argument('--input',type=Path,help='Optional JSON food/reactions model; see model.json.')
    args=parser.parse_args();args.output.mkdir(parents=True,exist_ok=True);start=time.perf_counter()
    system=ReactionSystem.from_record(json.loads(args.input.read_text())) if args.input else ReactionSystem(FOOD,[Reaction(*r) for r in REACTIONS])
    inputs=dict(model=system.record(),resolution_budget=RESOLUTION_BUDGET,availability=AVAILABLE_PROBABILITY,
                deletion_costs=DELETION_COSTS,required=REQUIRED_REACTIONS,forbidden=FORBIDDEN_REACTIONS,deleted=DELETED_REACTIONS,
                scaling_sizes=SCALING_SIZES,large_modules=LARGE_INDEPENDENT_MODULES,chain_length=DETERMINISTIC_CHAIN_LENGTH)
    intro='Resolved inputs: '+json.dumps(inputs);print(intro,flush=True)
    engine=SupplierEnumerator();catalogue=engine.catalogue(system)
    summary=dict(inputs=inputs,catalogue=family_record(catalogue))
    if not args.input:
        cuts=minimal_cuts(system,catalogue);chosen,cost=greedy_cut(catalogue,DELETION_COSTS)
        summary['worked_queries']=dict(required=list(REQUIRED_REACTIONS),forbidden=list(FORBIDDEN_REACTIONS),
            matching=[sorted(c) for c in catalogue.filter(REQUIRED_REACTIONS,FORBIDDEN_REACTIONS)],
            after_deletion=[sorted(c) for c in catalogue.filter(forbidden=DELETED_REACTIONS)],
            minimal_cuts=[sorted(c) for c in cuts],minimum_cut_cost=min(sum(DELETION_COSTS[r] for r in c) for c in cuts),
            greedy_cut=chosen,greedy_cost=str(cost),availability_probability=str(exact_availability(system,catalogue)))
    # Paper's catalyst-alternative example: resolved minimum is not original minimum.
    counterexample=ReactionSystem({'f'},[Reaction(0,{'f'},{'x'},{'x','y'}),Reaction(1,{'x'},{'y'},{'x'})])
    summary['necessary_validation']=family_record(engine.global_catalogue(counterexample))
    scaling=[]
    for size in SCALING_SIZES:
        for family,builder in [('independent',independent_modules),('coupled',coupled_gates)]:
            model=builder(size)
            for mode,enumerate_ in [('global',engine.global_catalogue),('components',engine.catalogue)]:
                result=enumerate_(model)
                scaling.append(dict(family=family,mode=mode,size=size,reactions=len(model.identifiers),
                    beta=result.beta,global_resolution_count=result.global_resolution_count,
                    resolutions_inspected=result.resolutions_inspected,outputs=len(result.members),
                    validation_calls=result.validation_calls,total_memberships=sum(map(len,result.members))))
    large=engine.catalogue(independent_modules(LARGE_INDEPENDENT_MODULES))
    summary['large_independent']=family_record(large)
    chain=deterministic_chain(DETERMINISTIC_CHAIN_LENGTH)
    summary['deterministic_chain']=dict(**family_record(engine.global_catalogue(chain)),closure_depth=len(chain.closure_stages())-1)
    zero=engine.catalogue(independent_modules(4,choices=False))
    summary['zero_excess_availability']=str(zero_excess_reliability(zero))
    summary['evidence']='Exact finite set computations and rational availability probabilities. Bounded independent tests, not Lean extraction or compilation. Structural RAF membership does not assert kinetic function.'
    with (args.output/'scaling.csv').open('w',newline='',encoding='utf-8') as f:
        writer=csv.DictWriter(f,fieldnames=list(scaling[0]));writer.writeheader();writer.writerows(scaling)
    with (args.output/'catalogue.csv').open('w',newline='',encoding='utf-8') as f:
        writer=csv.writer(f);writer.writerow(['irrRAF','reaction_ids']);writer.writerows((i+1,' '.join(map(str,sorted(c)))) for i,c in enumerate(catalogue.members))
    (args.output/'model.json').write_text(json.dumps(system.record(),indent=2)+'\n',encoding='utf-8')
    figures(system,catalogue,scaling,args.output)
    transcript=json.dumps(summary,indent=2)
    (args.output/'summary.json').write_text(transcript+'\n',encoding='utf-8')
    (args.output/'console.txt').write_text(intro+'\n'+transcript+'\n',encoding='utf-8')
    import matplotlib
    metadata=dict(paper_sha256=PAPER_SHA256,source_sha256=hashlib.sha256(Path(__file__).read_bytes()).hexdigest(),
        matplotlib=matplotlib.__version__,
        python=platform.python_version(),platform=platform.platform(),processor=platform.processor(),
        elapsed_seconds=time.perf_counter()-start,command='python example.py --output outputs',
        seed_policy='Deterministic enumeration; no stochastic sampling.',
        output_sha256={p.name:hashlib.sha256(p.read_bytes()).hexdigest() for p in sorted(args.output.iterdir()) if p.is_file() and p.name!='run_metadata.json'})
    (args.output/'run_metadata.json').write_text(json.dumps(metadata,indent=2)+'\n',encoding='utf-8')
    print(transcript)


if __name__=='__main__': main()
