"""Bounded upstream-MLX research harness, not a production training service.

Cosmos-Reason2 local conversion; frozen visual features; explicit two-stage
forward/backward over TCP. Removal condition: qualified fleet training runtime.
No pickle or remote executable payloads. No trained-model quality claim.
"""
import argparse, hashlib, io, json, os, resource, socket, struct, time
from pathlib import Path
import numpy as np
import mlx.core as mx
import mlx.nn as nn
import mlx.optimizers as optim
from mlx.utils import tree_flatten, tree_unflatten
from mlx_vlm import load
from mlx_vlm.trainer.lora_layers import LoRALinear

def emit(event, **kw):
    print(json.dumps(dict(event=event, **kw)), flush=True)

def flat(tree):
    return dict(tree_flatten(tree))

def save_tree(path, tree):
    mx.save_safetensors(str(path), flat(tree))

def load_tree(path):
    return tree_unflatten(list(mx.load(str(path)).items()))

def exact(s, n):
    result = bytearray()
    while len(result) < n:
        b = s.recv(n - len(result))
        if not b: raise EOFError('peer disconnected before message completed')
        result.extend(b)
    return result

def send(s, meta, tensors=None):
    b = io.BytesIO()
    np.savez(b, **{k:np.array(v) for k,v in (tensors or {}).items()})
    header = json.dumps(meta).encode(); payload = b.getvalue()
    s.sendall(struct.pack('!QQ', len(header), len(payload)) + header + payload)
    return 16 + len(header) + len(payload)

def recv(s):
    a,b = struct.unpack('!QQ', exact(s,16))
    if a > 65536 or b > 128*1024*1024: raise ValueError('message exceeds bounds')
    meta = json.loads(exact(s,a))
    with np.load(io.BytesIO(exact(s,b)),allow_pickle=False) as z:
        data = {k:mx.array(z[k]) for k in z.files}
    return meta,data

class Stage(nn.Module):
    def __init__(self, lm, start, end, last=False):
        super().__init__()
        self.layers = lm.model.layers[start:end]
        self.start = start
        if last:
            self.norm = lm.model.norm
            self.head = lm.lm_head

    def __call__(self, h, batch):
        for offset, layer in enumerate(self.layers):
            h = layer(h, mask='causal', position_ids=batch['position_ids'])
            idx = self.start + offset
            if idx < 3 and 'deepstack' in batch:
                mask = batch['visual_mask']
                # Data preparation emits dense additions, avoiding host indexing
                # inside the differentiated path.
                h = h + batch['deepstack'][idx].astype(h.dtype)
        if 'head' in self:
            h = self.head(self.norm(h))
        return h

def loss(logits, batch):
    ce = nn.losses.cross_entropy(logits.astype(mx.float32), batch['labels'], reduction='none')
    return (ce * batch['loss_mask']).sum() / batch['loss_mask'].sum()

def make_stages(base):
    mx.random.seed(20260909)
    model, processor = load(base)
    model.freeze()
    for layer in model.language_model.model.layers:
        for key in ('q_proj','v_proj'):
            original = getattr(layer.self_attn,key)
            setattr(layer.self_attn,key,LoRALinear.from_base(original,r=4,dropout=0,scale=2.0))
    first=Stage(model.language_model,0,14)
    last=Stage(model.language_model,14,28,last=True)
    mx.eval(first.parameters(),last.parameters())
    return model,processor,first,last

def prepare(args):
    from mlx_vlm.prompt_utils import apply_chat_template
    from mlx_vlm.utils import prepare_inputs
    model, processor=load(args.base)
    row=json.loads(Path(args.dataset).read_text().splitlines()[0])
    full=apply_chat_template(processor,model.config,row['messages'],num_images=1,add_generation_prompt=False)
    prefix=apply_chat_template(processor,model.config,row['messages'][:-1],num_images=1,add_generation_prompt=True)
    kwargs=dict(processor=processor,images=row['images'],image_token_index=model.config.image_token_index)
    inp=prepare_inputs(prompts=[full],**kwargs)
    pre=prepare_inputs(prompts=[prefix],**kwargs)
    ids=inp.pop('input_ids')
    feat=model.get_input_embeddings(ids,**inp)
    n=ids.shape[1]-1
    if n>512: raise ValueError(f'bounded pilot length exceeded: {n}')
    data={'h':feat.inputs_embeds[:,:-1].astype(mx.float32),
          'position_ids':feat.position_ids[...,:-1], 'labels':ids[:,1:],
          'loss_mask':(mx.arange(n)[None,:]>=pre['input_ids'].shape[1]-1).astype(mx.float32),
          'visual_mask':feat.visual_pos_masks[:,:-1]}
    if feat.deepstack_visual_embeds is not None:
        mask=np.array(feat.visual_pos_masks[:,:-1])
        additions=[]
        for values in feat.deepstack_visual_embeds:
            dense=np.zeros(data['h'].shape,dtype=np.float32)
            dense[mask]=np.array(values.astype(mx.float32))
            additions.append(mx.array(dense))
        data['deepstack']=mx.stack(additions)
    mx.eval(data)
    mx.savez(str(Path(args.out)/'batch.npz'),**data)
    info={'row_id':row['id'],'image_sha256':row['image_sha256'],'tokens':n,
          'supervised_tokens':float(data['loss_mask'].sum().item()),'scope':'synthetic diagram; frozen vision; one-example numerical pilot'}
    (Path(args.out)/'batch.json').write_text(json.dumps(info,indent=2))
    emit('prepared',**info)

def checkpoint(out,stage,opt,step):
    folder=out/f'step-{step:04}'
    if folder.exists():
        committed=json.loads((out/'committed.json').read_text())['step'] if (out/'committed.json').exists() else 0
        if step<=committed:raise FileExistsError('refusing to overwrite committed checkpoint')
        folder.rename(out/(folder.name+'-uncommitted-'+str(time.time_ns())))
    folder.mkdir(exist_ok=False)
    save_tree(folder/'adapter.safetensors',stage.trainable_parameters())
    save_tree(folder/'optimizer.safetensors',opt.state)
    (folder/'meta.json').write_text(json.dumps({'step':step,'optimizer':'Adam','lr':1e-4,'rank':4,'alpha':8}))
    return folder

def stats():
    return {'mlx_peak_bytes':mx.get_peak_memory(),'process_maxrss_bytes':resource.getrusage(resource.RUSAGE_SELF).ru_maxrss}

def run(args):
    out=Path(args.out);out.mkdir(parents=True,exist_ok=True)
    mx.set_memory_limit(8*1024**3);mx.set_cache_limit(128*1024**2)
    if args.mode=='prepare':return prepare(args)
    model,processor,a,b=make_stages(args.base)
    batch=mx.load(args.batch)
    if args.mode=='reference':
        deep = [mx.array(np.array(v)[np.array(batch['visual_mask'])]) for v in batch['deepstack']] if 'deepstack' in batch else None
        native_h = model.language_model.model(
            mx.zeros(batch['labels'].shape,dtype=mx.int32),inputs_embeds=batch['h'],
            mask='causal',position_ids=batch['position_ids'],
            visual_pos_masks=batch['visual_mask'],deepstack_visual_embeds=deep)
        native_loss=loss(model.language_model.lm_head(native_h),batch)
        mx.eval(native_loss);native_value=native_loss.item()
        del native_h,native_loss,model,processor
        class Whole(nn.Module):
            def __init__(self):super().__init__();self.a=a;self.b=b
            def __call__(self):return loss(self.b(self.a(batch['h'],batch),batch),batch)
        w=Whole(); opt=optim.Adam(1e-4)
        for step in range(args.steps):
            value,grads=nn.value_and_grad(w,lambda m:m())(w)
            mx.eval(value,grads)
            if step==0:
                delta=abs(value.item()-native_value)
                emit('native-forward-parity',native_loss=native_value,staged_loss=value.item(),absolute_error=delta)
                if delta>1e-5:raise ValueError('native/staged forward mismatch')
                save_tree(out/'first-grad-a.safetensors',grads['a'])
                save_tree(out/'first-grad-b.safetensors',grads['b'])
            opt.update(w,grads);mx.eval(w.parameters(),opt.state)
            emit('reference-step',step=step,loss=value.item(),**stats())
        save_tree(out/'final-a.safetensors',a.trainable_parameters())
        save_tree(out/'final-b.safetensors',b.trainable_parameters())
        (out/'result.json').write_text(json.dumps({'steps':args.steps,'last_loss':value.item(),**stats()},indent=2))
        return
    del model,processor
    stage=a if args.mode=='first' else b
    if args.mode=='first':del b
    else:del a
    opt=optim.Adam(1e-4);start=args.resume
    if start:
        folder=out/f'step-{start:04}'
        stage.update(load_tree(folder/'adapter.safetensors'))
        opt.state=load_tree(folder/'optimizer.safetensors')
        mx.eval(stage.parameters(),opt.state)
    config_id=hashlib.sha256(Path(args.batch).read_bytes()).hexdigest()
    if args.mode=='last':
        listener=socket.socket();listener.setsockopt(socket.SOL_SOCKET,socket.SO_REUSEADDR,1)
        listener.bind((args.host,args.port));listener.listen(1);listener.settimeout(300)
        emit('listening',host=args.host,port=args.port,start=start)
        s,addr=listener.accept();listener.close();emit('connected',peer=addr[0])
    else:s=socket.create_connection((args.host,args.port),timeout=120)
    s.settimeout(300)
    send(s,{'protocol':1,'start':start,'batch_sha256':config_id})
    peer,_=recv(s)
    if peer!={'protocol':1,'start':start,'batch_sha256':config_id}:raise ValueError('peer checkpoint/batch mismatch')
    try:
        for step in range(start,start+args.steps):
            begin=time.monotonic()
            if args.mode=='first':
                h=stage(batch['h'],batch);mx.eval(h)
                sent=send(s,{'op':'forward','step':step},{'h':h.astype(mx.float32)})
                if step==args.disconnect_step:
                    raise ConnectionAbortedError('intentional transport interruption after forward')
                meta,data=recv(s)
                if meta.get('op')!='gradient' or meta.get('step')!=step:raise ValueError('invalid gradient reply')
                def objective(params):
                    stage.update(params)
                    return (stage(batch['h'],batch).astype(mx.float32)*data['dh']).sum()
                _,grads=mx.value_and_grad(objective)(stage.trainable_parameters())
                value=meta['loss']
            else:
                meta,data=recv(s)
                if meta!={'op':'forward','step':step}:raise ValueError('invalid forward request')
                def objective(params,h):
                    stage.update(params);return loss(stage(h,batch),batch)
                value,(grads,dh)=mx.value_and_grad(objective,argnums=(0,1))(stage.trainable_parameters(),data['h'])
                mx.eval(value,grads,dh);value=float(value.item())
                sent=send(s,{'op':'gradient','step':step,'loss':value},{'dh':dh.astype(mx.float32)})
            mx.eval(grads)
            if step==0:save_tree(out/'first-grad.safetensors',grads)
            if not np.isfinite(value):raise ValueError('nonfinite loss')
            if not all(bool(mx.all(mx.isfinite(v)).item()) for v in flat(grads).values()):raise ValueError('nonfinite gradient')
            opt.update(stage,grads);mx.eval(stage.parameters(),opt.state)
            checkpoint(out,stage,opt,step+1)
            send(s,{'op':'saved','step':step+1});ack,_=recv(s)
            if ack!={'op':'saved','step':step+1}:raise ValueError('checkpoint acknowledgement mismatch')
            (out/'committed.json').write_text(json.dumps({'step':step+1,'batch_sha256':config_id}))
            emit('distributed-step',rank=args.mode,step=step,loss=value,seconds=time.monotonic()-begin,sent_bytes=sent,**stats())
            mx.clear_cache()
        save_tree(out/'final.safetensors',stage.trainable_parameters())
        (out/f'result-{start}.json').write_text(json.dumps({'start':start,'steps':args.steps,'last_loss':value,**stats()},indent=2))
    finally:s.close()

if __name__=='__main__':
    p=argparse.ArgumentParser();p.add_argument('mode',choices=['prepare','reference','first','last'])
    p.add_argument('--base',required=True);p.add_argument('--out',required=True)
    p.add_argument('--dataset');p.add_argument('--batch');p.add_argument('--steps',type=int,default=2)
    p.add_argument('--resume',type=int,default=0);p.add_argument('--host',default='127.0.0.1');p.add_argument('--port',type=int,default=29471)
    p.add_argument('--disconnect-step',type=int,default=-1)
    run(p.parse_args())
