"""Full source-aligned rational audit and reusable sparse finite-horizon LP."""
from fractions import Fraction as Q
import hashlib
import json
from pathlib import Path
import warnings
import numpy as np
from scipy.optimize import linprog, OptimizeWarning
from scipy.sparse import coo_matrix, vstack, hstack, eye, csr_matrix
from freeze_source import medium_masks, BUDGET_ROWS

OBJECTIVES={'service':{'R_NaKt':1},'turnover':{'R_CAT':2,'R_GTHP':1},
            'joint':{'R_NaKt':1,'R_CAT':8,'R_GTHP':4}}


def read_json(path):
    def unique(pairs):
        out={}
        for k,v in pairs:
            if k in out:raise ValueError(f'Duplicate JSON identifier: {k}')
            out[k]=v
        return out
    return json.loads(Path(path).read_text(encoding='utf-8'),object_pairs_hook=unique)


class SourceModel:
    def __init__(self,data):
        self.data=data;self.rows=data['rows'];self.columns=data['columns'];self.masks=data['medium_masks']
        self.ri={r['id']:i for i,r in enumerate(self.rows)};self.ci={c['id']:i for i,c in enumerate(self.columns)}
        if len(self.ri)!=len(self.rows) or len(self.ci)!=len(self.columns):raise ValueError('Duplicate row/column identifiers.')
        self.aux=set();self.chem=set()
        for r in self.rows:
            kind='auxiliary' if r['compartment']=='pc' or r['id'] in BUDGET_ROWS else 'chemical'
            if r['kind']!=kind:raise ValueError('Row classification mismatch.')
            (self.aux if kind=='auxiliary' else self.chem).add(r['id'])
        if medium_masks(self.rows,self.columns)!=self.masks:raise ValueError('Medium masks do not match source formulas and exchange stoichiometry.')
        for c in self.columns:
            if set(c['stoichiometry'])-self.ri.keys():raise ValueError('Unknown stoichiometric row.')
            if Q(c['lower'])>Q(c['upper']):raise ValueError('Reversed source bounds.')
        self.stoich=[{s:Q(v) for s,v in c['stoichiometry'].items()} for c in self.columns]

    @classmethod
    def load(cls,directory):
        directory=Path(directory);manifest=read_json(directory/'source_manifest.json')
        for name,key in [('s7_model.json','model_sha256'),('certificates.json','certificates_sha256')]:
            if hashlib.sha256((directory/name).read_bytes()).hexdigest()!=manifest[key]:raise ValueError(f'Frozen input hash mismatch: {name}')
        data=read_json(directory/'s7_model.json')
        if data['source_xml_sha256']!=manifest['source_xml_sha256']:raise ValueError('Source XML provenance mismatch.')
        model=cls(data)
        if (len(model.rows),len(model.columns),len(model.chem),len(model.aux))!=(10411,19620,2157,8254):raise ValueError('Frozen S7 dimensions changed.')
        return model,read_json(directory/'certificates.json'),manifest

    def bounds(self,medium='restricted',time=1):
        if medium not in ('original','sulfur_closed','restricted','internal_blocked'):raise ValueError('Unknown declared medium.')
        time=Q(time)
        if time<0:raise ValueError('Nonnegative horizon required.')
        closed=set(self.masks['nonionic_sink_imports'])
        if medium in ('sulfur_closed','restricted'):closed.update(self.masks['sulfur_exchange_imports'])
        if medium=='restricted':closed.update(self.masks['other_carbon_exchange_imports'])
        if medium=='internal_blocked':closed.update(('R_CYSGLTH','R_LHCYSTIN','R_HTSULGTHST','R_SSGTHRD'))
        return [(max(Q(c['lower']),Q(0)) * time if c['id'] in closed else Q(c['lower'])*time,Q(c['upper'])*time) for c in self.columns]

    def vector(self,sparse,kind='reaction'):
        universe=self.ci if kind=='reaction' else self.ri
        if set(sparse)-universe.keys():raise ValueError('Unknown source identifier in sparse vector.')
        return {k:Q(v) for k,v in sparse.items()}

    def chemical_vector(self,sparse):
        result=self.vector(sparse,'row')
        if set(result)-self.chem:raise ValueError('Chemical inventory contains auxiliary rows.')
        if any(v<0 for v in result.values()):raise ValueError('Negative chemical inventory/floor.')
        return result


class RationalAudit:
    def __init__(self,model):self.model=model
    def primal(self,sparse,medium='restricted',time=1,initial=None,floors=None):
        m=self.model;x=m.vector(sparse);initial=m.chemical_vector(initial or {});floors=m.chemical_vector(floors or {})
        amounts={r:initial.get(r,Q(0)) for r in m.ri};bounds=m.bounds(medium,time);violations=[]
        for c,s,(lo,hi) in zip(m.columns,m.stoich,bounds):
            extent=x.get(c['id'],Q(0))
            if not lo<=extent<=hi:violations.append(c['id'])
            for row,value in s.items():amounts[row]+=value*extent
        bad=[r for r,v in amounts.items() if (r in m.aux and v!=0) or (r in m.chem and v<floors.get(r,Q(0)))]
        if violations or bad:raise ValueError(f'Exact primal infeasible: {len(violations)} bound and {len(bad)} balance/floor violations; {(violations+bad)[:5]}')
        return dict(nonzero_extents=sum(v!=0 for v in x.values()),checked_columns=len(m.columns),checked_rows=len(m.rows),
                    M=x.get('R_NaKt',Q(0)),Q=2*x.get('R_CAT',Q(0))+x.get('R_GTHP',Q(0)),
                    terminal={k:v for k,v in amounts.items() if v},extent_GTHP=x.get('R_GTHP',Q(0)))

    def upper(self,certificate,objective,medium='restricted',time=1,initial=None,floors=None):
        m=self.model;w=m.vector(certificate['chemical'],'row');y=m.vector(certificate['auxiliary'],'row')
        if set(w)-m.chem or set(y)-m.aux or any(v<0 for v in w.values()):raise ValueError('Invalid dual signs or row classification.')
        alpha=m.vector(certificate.get('reverse',{}));objective=m.vector(objective)
        if any(v<0 for v in alpha.values()):raise ValueError('Negative reverse charge.')
        initial=m.chemical_vector(initial or {});floors=m.chemical_vector(floors or {})
        costs=[];total=Q(0);residuals={}
        for c,s,(lo,hi) in zip(m.columns,m.stoich,m.bounds(medium,time)):
            rid=c['id'];d=objective.get(rid,Q(0))+alpha.get(rid,Q(0))+sum((w.get(row,Q(0))-y.get(row,Q(0)))*v for row,v in s.items())
            cost=max(d*lo,d*hi);total+=cost
            if d:residuals[rid]=d
            if cost:costs.append(dict(reaction=rid,residual=d,interval_cost=cost))
        stock=sum(v*(initial.get(r,Q(0))-floors.get(r,Q(0))) for r,v in w.items())
        return dict(bound_without_reverse=stock+total,interval_cost=total,stock=stock,reverse_charges=alpha,
            chemical_weights=len(w),auxiliary_weights=len(y),nonzero_costs=len(costs),costs=costs,
            residuals=residuals,checked_columns=len(m.columns))

    def bypass(self):
        m=self.model;route={'R_HCYSTRDX':1,'R_TRDRy':1,'R_LHCYSTIN':-1,'R_GTHP':1,'R_ESTRONEDHy':-1}
        net={}
        for rid,extent in route.items():
            for row,value in m.stoich[m.ci[rid]].items():net[row]=net.get(row,Q(0))+extent*value
        chemical={k:v for k,v in net.items() if v and k in m.chem};auxiliary={k:v for k,v in net.items() if v and k in m.aux}
        assert chemical=={'M_estradiol_c':Q(-1),'M_h2o2_c':Q(-1),'M_estrone_c':Q(1),'M_h2o_c':Q(2)} and len(auxiliary)==8
        return dict(route=route,chemical_net=chemical,auxiliary_net=auxiliary,scope='Chemical cancellation only; surviving enzyme-allocation terms prevent calling this a standalone feasible cycle.')


class LinearRelaxation:
    """Numerical full-S7 LP for new scenarios. Results are not rational certificates."""
    def __init__(self,model,medium='restricted',time=1,initial=None,floors=None):
        self.model=model;self.time=float(time);self.medium=medium
        initial=model.chemical_vector(initial or {});floors=model.chemical_vector(floors or {})
        rr=[];cc=[];vv=[]
        for j,s in enumerate(model.stoich):
            for row,value in s.items():rr.append(model.ri[row]);cc.append(j);vv.append(float(value))
        self.N=coo_matrix((vv,(rr,cc)),shape=(len(model.rows),len(model.columns))).tocsr()
        self.chem=[i for i,r in enumerate(model.rows) if r['id'] in model.chem];self.aux=[i for i,r in enumerate(model.rows) if r['id'] in model.aux]
        self.initial=np.array([float(initial.get(r['id'],0)) for r in model.rows]);self.floor=np.array([float(floors.get(r['id'],0)) for r in model.rows])
        self.bounds=np.array([[float(lo),float(hi)] for lo,hi in model.bounds(medium,time)])

    def objective(self,name):
        if name=='peroxide':return self.N.getrow(self.model.ri['M_h2o2_c'])
        d=OBJECTIVES[name];return coo_matrix(([float(v) for v in d.values()],([0]*len(d),[self.model.ci[k] for k in d])),shape=(1,len(self.model.columns))).tocsr()

    def solve(self,objective='turnover',required_service=0.,fixed=None,capacity_fractions=None,parsimony=False,time_limit=40):
        n=len(self.model.columns);A=-self.N[self.chem];b=(self.initial-self.floor)[self.chem]
        eq=self.N[self.aux];bounds=self.bounds.copy()
        service=self.model.ci['R_NaKt'];bounds[service,0]=max(bounds[service,0],float(required_service))
        for rid,fraction in (capacity_fractions or {}).items():
            if rid not in self.model.ci or not np.isfinite(fraction) or not 0<=fraction<=1:raise ValueError('Capacity scaling needs a known reaction and fraction in [0,1].')
            j=self.model.ci[rid]
            if bounds[j,1]<0:raise ValueError('Upper capacity scaling requires a nonnegative upper bound.')
            bounds[j,1]*=fraction
        for name,(lo,hi) in (fixed or {}).items():
            row=self.objective(name);A=vstack((A,row,-row),format='csr');b=np.r_[b,hi,-lo]
        target=self.objective(objective).toarray().ravel()*(1 if objective=='peroxide' else -1)
        if parsimony:
            A=vstack((hstack((A,csr_matrix((A.shape[0],n)))),hstack((eye(n),-eye(n))),hstack((-eye(n),-eye(n)))),format='csr')
            b=np.r_[b,np.zeros(2*n)];eq=hstack((eq,csr_matrix((eq.shape[0],n))),format='csr');bounds=np.vstack((bounds,np.column_stack((np.zeros(n),np.full(n,np.inf)))));target=np.r_[np.zeros(n),np.ones(n)]
        with warnings.catch_warnings():
            warnings.filterwarnings('ignore',category=OptimizeWarning,message='Unrecognized options detected.*')
            fit=linprog(target,A_ub=A,b_ub=b,A_eq=eq,b_eq=np.zeros(len(self.aux)),bounds=bounds,method='highs',
                options={'threads':1,'time_limit':time_limit,'primal_feasibility_tolerance':1e-9,'dual_feasibility_tolerance':1e-9})
        if not fit.success:return dict(status=int(fit.status),message=fit.message,finished=False)
        x=fit.x[:n];terminal=self.initial+self.N@x;peroxide=self.N.getrow(self.model.ri['M_h2o2_c']).toarray().ravel()*x
        uptake=self.model.ci['R_H2O2t'];counted={self.model.ci['R_CAT'],self.model.ci['R_GTHP']}
        other=sum(-v for j,v in enumerate(peroxide) if v<0 and j not in counted and j!=uptake)
        production=sum(v for j,v in enumerate(peroxide) if v>0 and j!=uptake)
        result=dict(finished=True,status=0,M=float(x[service]),Q=float((self.objective('turnover')@x)[0]),
            peroxide_import=max(0.,float(peroxide[uptake])),peroxide_export=max(0.,-float(peroxide[uptake])),
            peroxide_other_consumption=float(other),peroxide_internal_production=float(production),terminal_peroxide=float(terminal[self.model.ri['M_h2o2_c']]),
            minimum_floor_slack=float(np.min((terminal-self.floor)[self.chem])),max_auxiliary_residual=float(np.max(abs(terminal[self.aux]))),
            peroxide_terms={self.model.columns[j]['id']:float(v) for j,v in enumerate(peroxide) if v!=0},
            scope='Floating-point integrated relaxation; no trajectory or exact feasibility claim.')
        return result
