← back to Local Model Leaderboard Watch

test-model-preflight.py

97 lines

#!/usr/bin/env python3
"""Run the actual wrapper in a temporary filesystem with all external calls stubbed."""
import json
import os
from pathlib import Path
import shutil
import struct
import subprocess
import tempfile

SOURCE = Path(__file__).resolve().parent

def shard(path, dtype='F32', shape=None, size=4, header=None):
    if header is None:
        header = {'weight': {'dtype': dtype, 'shape': [1] if shape is None else shape, 'data_offsets': [0, size]}}
    encoded = json.dumps(header).encode()
    path.write_bytes(struct.pack('<Q', len(encoded)) + encoded + b'\0'*size)


def case(name, setup, expected_posts, chosen='org/model'):
    with tempfile.TemporaryDirectory(prefix='tk11401-model-') as directory:
        root = Path(directory)
        app, home, bins = root/'app', root/'home', root/'bin'
        app.mkdir(); home.mkdir(); bins.mkdir(); (app/'data').mkdir()
        for source in ['review.sh', 'verify-model-weights.py']:
            shutil.copy2(SOURCE/source, app/source)
        (app/'review-prompt.md').write_text('offline fixture only')
        (app/'data/chosen-model.txt').write_text(chosen+'\n')
        model = home/'.exo/models/org--model'; model.mkdir(parents=True)
        (model/'config.json').write_text('{"model_type":"fixture"}')
        setup(model)
        log = root/'requests.jsonl'
        curl = bins/'curl'
        curl.write_text('#!/usr/bin/env python3\nimport json,os,sys\nwith open(os.environ["REQUEST_LOG"],"a") as f: f.write(json.dumps(sys.argv[1:])+"\\n")\nprint("{\\"instances\\": {}}")\n')
        curl.chmod(0o755)
        timeout = bins/'timeout'; timeout.write_text('#!/bin/sh\nexit 0\n'); timeout.chmod(0o755)
        env = dict(os.environ, EXO_MODELS_DIR=str(model.parent), PATH=str(bins)+':'+os.environ['PATH'], REQUEST_LOG=str(log))
        result = subprocess.run(['bash', str(app/'review.sh')], env=env, capture_output=True, text=True, timeout=10)
        calls = [json.loads(line) for line in log.read_text().splitlines()] if log.exists() else []
        posts = [args for args in calls if any('/v1/chat/completions' in arg for arg in args)]
        assert result.returncode == 0, (name, result.stderr)
        assert len(posts) == expected_posts, (name, calls, (app/'data/run.log').read_text())
        run_log = (app/'data/run.log').read_text()
        assert 'Traceback' not in run_log, (name, run_log)
        if not expected_posts and chosen:
            assert 'REFUSED' in run_log, name
        return {'name': name, 'verdict': 'PASS', 'model_load_requests': len(posts), 'external_calls': 'all mocked'}


def missing(root):
    (root/'model.safetensors.index.json').write_text('{"weight_map":{"weight":"model-00001-of-00001.safetensors"}}')

def complete(root):
    missing(root); shard(root/'model-00001-of-00001.safetensors')

def truncated(root):
    complete(root); p=root/'model-00001-of-00001.safetensors'; p.write_bytes(p.read_bytes()[:-1])

def wrong_tensor(root):
    complete(root); (root/'model.safetensors.index.json').write_text('{"weight_map":{"absent":"model-00001-of-00001.safetensors"}}')

def traversal(root):
    (root/'model.safetensors.index.json').write_text('{"weight_map":{"weight":"../outside.safetensors"}}')

results = [
    case('config-only stub refuses model POST', lambda p: None, 0),
    case('indexed missing shard refuses model POST', missing, 0),
    case('truncated shard refuses model POST', truncated, 0),
    case('malformed index refuses model POST', lambda p: (p/'model.safetensors.index.json').write_text('{'), 0),
    case('missing mapped tensor refuses model POST', wrong_tensor, 0),
    case('path traversal refuses model POST', traversal, 0),
    case('complete indexed model permits mocked POST', complete, 1),
    case('complete single shard permits mocked POST', lambda p: shard(p/'model.safetensors'), 1),
    case('shape payload mismatch refuses model POST', lambda p: shard(p/'model.safetensors', shape=[1000]), 0),
    case('unknown dtype refuses model POST', lambda p: shard(p/'model.safetensors', dtype='NEW_UNKNOWN'), 0),
    case('negative dimension refuses model POST', lambda p: shard(p/'model.safetensors', shape=[-1]), 0),
    case('boolean dimension refuses model POST', lambda p: shard(p/'model.safetensors', shape=[True]), 0),
    case('nonobject index refuses without traceback', lambda p: (p/'model.safetensors.index.json').write_text('[]'), 0),
    case('nonobject header refuses without traceback', lambda p: shard(p/'model.safetensors', header=[]), 0),
    case('nonobject tensor refuses without traceback', lambda p: shard(p/'model.safetensors', header={'weight': None}), 0),
    case('unhashable index value refuses without traceback', lambda p: (p/'model.safetensors.index.json').write_text('{"weight_map":{"weight":[]}}'), 0),
    case('zero-size tensor permits mocked POST', lambda p: shard(p/'model.safetensors', shape=[0,1000], size=0), 1),
    case('scalar tensor permits mocked POST', lambda p: shard(p/'model.safetensors', shape=[]), 1),
    case('blank choice remains no-op', lambda p: None, 0, chosen=''),
]
# Exercise each supported dtype against the real wrapper, including packed widths.
for dtype, count, size in [
    ('BOOL', 1, 1), ('I8', 1, 1), ('U8', 1, 1), ('I16', 1, 2), ('U16', 1, 2),
    ('F16', 1, 2), ('BF16', 1, 2), ('I32', 1, 4), ('U32', 1, 4), ('F32', 1, 4),
    ('I64', 1, 8), ('U64', 1, 8), ('F64', 1, 8), ('C64', 1, 8),
    ('F8_E4M3', 1, 1), ('F8_E5M2', 1, 1), ('F8_E8M0', 1, 1),
    ('F8_E4M3FNUZ', 1, 1), ('F8_E5M2FNUZ', 1, 1),
    ('F4', 2, 1), ('F6_E2M3', 4, 3), ('F6_E3M2', 4, 3),
]:
    results.append(case('supported dtype '+dtype, lambda p, d=dtype, c=count, s=size: shard(p/'model.safetensors', dtype=d, shape=[c], size=s), 1))
print(json.dumps(results, indent=2))