#!/usr/bin/env python3
"""Replay captured allocator requests without changing their constraints."""
import json, math, sys
from pathlib import Path
R1=524288
HOST=(327680,360448)
MAX=2**32-1

def align(n,a):return (n+a-1)//a*a

def subtract(ranges,lo,hi):
 out=[]
 for a,b in ranges:
  if hi<=a or b<=lo:out.append((a,b));continue
  if a<lo:out.append((a,lo))
  if hi<b:out.append((hi,b))
 return out

def forbidden_span(a,b):
 out=[]
 for lo,hi,elem in [(a,min(b,R1),16384),(max(a,R1),b,32768)]:
  if lo<hi:out.append((lo//elem*elem,align(hi,elem)))
 return out

def place(data,order='alignment',fit='first'):
 req=list(enumerate(data['requests']))
 def key(pair):
  i,r=pair;b=r['bytes'];a=r['alignment'];life=r['lifetime'];duration=life['last']-life['first']
  if order=='alignment':return (-a,-b,-duration,life['first'],i)
  if order=='size':return (-b,-a,-duration,life['first'],i)
  if order=='persistent':return (life['last']!=MAX,-b,-a,life['first'],i)
  if order=='resident':return (not(life['first']==0 and life['last']==MAX),-b,-a,life['first'],i)
  if order=='time':return (life['first'],r['class']!='Ipu21Interleaved',-a,-b,life['last'],i)
  if order.startswith('weighted'):return (-b*a**float(order[8:]),-duration,life['first'],i)
  raise ValueError(order)
 req.sort(key=key)
 history=[];roots={};allocation={}
 for i,r in req:
  life=r['lifetime'];first,last=life['first'],life['last'];free=list(map(tuple,data['ranges']))
  for j,a,b in history:
   q=data['requests'][j]['lifetime']
   if q['last']>=first and last>=q['first']:free=subtract(free,a,b)
  forbidden=sorted(s for root in r['conflicts'] if root in roots for s in forbidden_span(*roots[root]))
  candidates=[]
  for base,limit in free:
   if HOST[0]<=base and limit<=HOST[1] and (first==0 or last==MAX):continue
   for region1 in (False,True):
    if region1:
     lo=max(base,R1+(data['interleaved_offset'] if r['class']=='Ipu21Interleaved' else 0));hi=limit
    else:
     if r['class']=='Ipu21Interleaved':continue
     lo=base;hi=min(limit,R1)
    alignment=max(r['alignment'],(32768 if region1 else 16384) if r['region1_stride'] is not None else 1)
    size=r['region1_stride']*len(r['assignments']) if region1 and r['region1_stride'] is not None else r['bytes']
    start=align(lo,alignment)
    for a,b in forbidden:
     if start<b and a<start+size:start=align(b,alignment)
    if start+size>hi:continue
    score=(region1, start) if fit=='first' else (region1, hi-start-size, start)
    candidates.append((score,start,start+size))
  if not candidates:
   return {'fit':False,'placed':len(history),'request':i,'bytes':r['bytes'],'class':r['class'],'free':sum(b-a for a,b in free),'largest':max((b-a for a,b in free),default=0)}
  _,start,end=min(candidates)
  history.append((i,start,end));allocation[i]=(start,end)
  for root,_ in r['assignments']:roots[root]=(start,end)
 # Independently verify overlap and element separation after all requests.
 for i,(a,b) in allocation.items():
  r=data['requests'][i];life=r['lifetime']
  for j,(c,d) in allocation.items():
   if j>=i:continue
   other=data['requests'][j]['lifetime']
   assert b<=c or d<=a or life['last']<other['first'] or other['last']<life['first']
  for root in r['conflicts']:
   if root in roots:
    assert all(b<=c or d<=a for c,d in forbidden_span(*roots[root]))
 return {'fit':True,'placed':len(history)}

def split_sequences(data):
 import copy
 result=copy.deepcopy(data);groups={}
 for r in result['requests']:
  roots=[a[0] for a in r['assignments']]
  if len(roots)>1:
   for root in roots:groups[root]=roots
 requests=[]
 for r in result['requests']:
  r['conflicts']=sorted({member for root in r['conflicts'] for member in groups.get(root,[root])})
  if len(r['assignments'])<=1:requests.append(r);continue
  for root,_ in r['assignments']:
   part=copy.deepcopy(r);part['bytes']=r['bytes']//len(r['assignments']);part['assignments']=[[root,0]]
   if r['region1_stride'] is not None:part['conflicts']+= [other for other,_ in r['assignments'] if other!=root]
   requests.append(part)
 result['requests']=requests
 return result

if __name__=='__main__':
 results={}
 for path in sorted(Path(sys.argv[1]).glob('*.json')):
  d=json.loads(path.read_text());variants={}
  for order in ['alignment','time','size','persistent','resident','weighted0.25','weighted0.5','weighted1']:
   for fit in ['first','best']:
    variants[order+'/'+fit]=place(d,order,fit)
  split=split_sequences(d)
  for order in ['alignment','size','persistent','time']:
   variants['split/'+order]=place(split,order,'best')
  results[path.name]=variants
  print(path.name, 'baseline',variants['alignment/first'], 'fits',[k for k,v in variants.items() if v['fit']])
 Path(sys.argv[1]).parent.joinpath('results.json').write_text(json.dumps(results,indent=2)+'\n')
