"""Trusted, supervised CPU proposer worker. Sends bounded JSON bytes, never edits the workspace."""
import importlib.metadata
import json
import os
import socket
import sys
import time
from tiny_agent_boundary import Denied
from tiny_agent_artifacts import verify

RUNTIME={'torch':'2.9.0','transformers':'4.57.1','tokenizers':'0.22.1','safetensors':'0.6.2','huggingface-hub':'0.35.3'}


def worker(connection,directory):
    try:
        os.environ.update(HF_HUB_OFFLINE='1',TRANSFORMERS_OFFLINE='1',HF_HUB_DISABLE_IMPLICIT_TOKEN='1',TOKENIZERS_PARALLELISM='false')
        def deny_connect(*args,**kwargs):raise RuntimeError('Offline worker forbids socket connections')
        socket.socket.connect=deny_connect
        if sys.version_info[:3]!=(3,13,13):raise Denied('python_runtime_mismatch')
        observed={name:importlib.metadata.version(name) for name in RUNTIME}
        if observed!=RUNTIME:raise Denied('runtime_mismatch')
        artifacts=verify(directory)
        import torch
        from transformers import AutoTokenizer,AutoModelForCausalLM,GenerationConfig
        torch.set_num_threads(2);torch.set_num_interop_threads(1);torch.manual_seed(0)
        loaded=time.monotonic()
        tokenizer=AutoTokenizer.from_pretrained(directory,local_files_only=True,trust_remote_code=False)
        model=AutoModelForCausalLM.from_pretrained(directory,local_files_only=True,trust_remote_code=False,
                   use_safetensors=True,torch_dtype=torch.float32,attn_implementation='eager').to('cpu').eval()
        count=sum(p.numel() for p in model.parameters())
        if count!=494032768 or type(model).__name__!='Qwen2ForCausalLM' or any(p.device.type!='cpu' or p.dtype!=torch.float32 for p in model.parameters()):
            raise Denied('loaded_model_mismatch')
        generation=GenerationConfig(do_sample=False,num_beams=1,num_return_sequences=1,max_new_tokens=128,
             repetition_penalty=1.0,use_cache=True,bos_token_id=151643,pad_token_id=151643,eos_token_id=[151645,151643])
        connection.send_bytes(json.dumps({'status':'ready','runtime':observed,'parameters':count,
                'load_seconds':time.monotonic()-loaded,'artifacts':artifacts,'requested_generation':generation.to_dict()}).encode())
        while True:
            raw=connection.recv_bytes(32768);request=json.loads(raw)
            if request.get('stop'):return
            messages=request['messages']
            inputs=tokenizer.apply_chat_template(messages,tokenize=True,add_generation_prompt=True,return_dict=True,return_tensors='pt')
            length=inputs['input_ids'].shape[1]
            if length>2048:
                connection.send_bytes(b'{"status":"context_limit"}');continue
            # Capture the ordinary call's prepared effective config without altering generated defaults.
            effective=[];original=model._prepare_generation_config
            def capture(*args,**kwargs):
                result=original(*args,**kwargs);effective.append(result[0].to_dict());return result
            model._prepare_generation_config=capture
            start=time.monotonic()
            try:
                with torch.inference_mode():output=model.generate(**inputs,generation_config=generation)
            finally:model._prepare_generation_config=original
            ids=output[0,length:].tolist()
            if len(ids)>128:raise Denied('output_token_limit')
            decoded=ids[:-1] if ids and ids[-1] in (151645,151643) else ids
            text=tokenizer.decode(decoded,skip_special_tokens=False,clean_up_tokenization_spaces=False)
            result={'status':'proposal','input_tokens':length,'input_ids':inputs['input_ids'][0].tolist(),
                    'raw_output_ids':ids,'output_tokens':len(ids),'proposal_text':text,
                    'generation_seconds':time.monotonic()-start,'effective_generation':effective}
            encoded=json.dumps(result).encode()
            if len(encoded)>131072:raise Denied('worker_result_limit')
            connection.send_bytes(encoded)
    except BaseException as error:
        try:connection.send_bytes(json.dumps({'status':'model_error','diagnostic_type':type(error).__name__,'code':getattr(error,'code','model_error')}).encode())
        except (BrokenPipeError,EOFError,OSError):pass
    finally:connection.close()
