# Copyright © 2026 Swift contributors. """Offline HTTP regression checks on tiny random models, never Swift weights.""" import argparse import contextlib import copy import functools import importlib.metadata import json import logging import os import platform import shutil import sys import tempfile import threading import urllib.request from http.server import ThreadingHTTPServer from pathlib import Path from unittest.mock import patch os.environ["HF_HUB_OFFLINE"] = "1" os.environ["TRANSFORMERS_OFFLINE"] = "1" import mlx.core as mx import mlx.nn as nn from mlx.utils import tree_flatten from mlx_lm import server from mlx_lm.models.cache import LRUPromptCache from mlx_lm.models.qwen3_5_full import Model, ModelArgs class RecordingCache(LRUPromptCache): def __init__(self, *args, **kwargs): super().__init__(*args, **kwargs) self.peak_bytes = 0 self.insertions = 0 def insert_cache(self, *args, **kwargs): super().insert_cache(*args, **kwargs) self.insertions += 1 self.peak_bytes = max(self.peak_bytes, self.nbytes) assert self.nbytes <= self.max_bytes assert len(self) <= self.max_size def make_fixture(folder, assets, config_path, bits): config = json.loads(config_path.read_text()) source = json.loads((assets / "config.json").read_text()) config["text_config"]["vocab_size"] = source["text_config"]["vocab_size"] config.pop("quantization", None) mx.random.seed(25) model = Model(ModelArgs.from_dict(copy.deepcopy(config))) model.language_model.lm_head.weight = mx.zeros_like( model.language_model.lm_head.weight ) model.apply( lambda value: ( value.astype(mx.bfloat16) if mx.issubdtype(value.dtype, mx.floating) else value ) ) nn.quantize( model, bits=bits, group_size=64, mode="affine", class_predicate=lambda path, layer: hasattr(layer, "to_quantized") and layer.weight.shape[-1] % 64 == 0, ) mx.eval(model.parameters()) weights = dict(tree_flatten(model.parameters())) mx.save_safetensors( str(folder / "model.safetensors"), weights, metadata={"purpose": "synthetic regression only"}, ) index = { "metadata": {"total_size": sum(w.nbytes for w in weights.values())}, "weight_map": {name: "model.safetensors" for name in weights}, } (folder / "model.safetensors.index.json").write_text(json.dumps(index)) config["quantization"] = {"bits": bits, "group_size": 64, "mode": "affine"} (folder / "config.json").write_text(json.dumps(config)) for name in ( "tokenizer.json", "tokenizer_config.json", "vocab.json", "merges.txt", "chat_template.jinja", "generation_config.json", ): shutil.copyfile(assets / name, folder / name) del model, weights mx.clear_cache() @contextlib.contextmanager def running(folder, options): captured = [] argv = ["mlx_lm.server", "--model", str(folder), "--port", "0", *options] with patch.object(sys, "argv", argv), patch.object( server, "run", side_effect=lambda h, p, m: captured.append(m) ), patch.object(server, "maybe_set_recommended_wired_limit", return_value=None): server.main() provider = captured[0] instances = [] def start_http(host, port, generator): handler = functools.partial( server.APIHandler, generator, system_fingerprint="synthetic-offline-test" ) httpd = ThreadingHTTPServer((host, port), handler) thread = threading.Thread(target=httpd.serve_forever, daemon=True) instances.append((httpd, thread, generator)) thread.start() with patch.object(server, "LRUPromptCache", RecordingCache), patch.object( server, "_run_http_server", start_http ): server.run("127.0.0.1", 0, provider) httpd, thread, generator = instances[0] try: yield f"http://127.0.0.1:{httpd.server_port}", generator, provider.cli_args finally: httpd.shutdown() httpd.server_close() thread.join(timeout=5) generator.stop_and_join() def request(url, messages, stream=False, seeded=False): data = { "model": "default_model", "messages": messages, "max_tokens": 4, "temperature": 0, "chat_template_kwargs": {"enable_thinking": False}, "stream": stream, } if seeded: data["seed"] = 25 req = urllib.request.Request( url + "/v1/chat/completions", data=json.dumps(data).encode(), headers={"Content-Type": "application/json"}, ) with urllib.request.urlopen(req, timeout=60) as response: body = response.read().decode() assert response.status == 200 if stream: assert "data: [DONE]" in body chunks = [ json.loads(line[6:]) for line in body.splitlines() if line.startswith("data: ") and line != "data: [DONE]" ] text = "".join( choice.get("delta", {}).get("content", "") or "" for chunk in chunks for choice in chunk.get("choices", []) ) assert text == "!!!!", text return {"http_status": 200, "stream_complete": True} result = json.loads(body) assert result["choices"][0]["message"]["content"] == "!!!!" usage = result["usage"] return { "http_status": 200, "prompt_tokens": usage["prompt_tokens"], "cached_tokens": usage["prompt_tokens_details"]["cached_tokens"], } def main(): parser = argparse.ArgumentParser(description=__doc__) parser.add_argument( "--assets", type=Path, required=True, help="Local Swift snapshot; only tokenizer and config assets are read", ) parser.add_argument("--config", type=Path, required=True) parser.add_argument("--bits", type=int, choices=(4, 5), required=True) parser.add_argument("--output", type=Path, required=True) args = parser.parse_args() logging.basicConfig(level=logging.WARNING) report = { "scope": "Tiny random weights with real local tokenizer assets. " "No released model weights loaded or changed; no full-model capacity claim.", "bits": args.bits, "offline_hub": True, "memory_limit_overrides": False, "machine": platform.machine(), "python": platform.python_version(), "packages": { name: importlib.metadata.version(name) for name in ("mlx", "mlx-lm", "transformers", "huggingface_hub") }, "profiles": [], } profiles = [ ("defaults", []), ("explicit_bytes", ["--prompt-cache-bytes", "100000"]), ("no_retention", ["--prompt-cache-size", "0"]), ] with tempfile.TemporaryDirectory(prefix="swift-cache-test-") as temp: folder = Path(temp) make_fixture(folder, args.assets, args.config, args.bits) for name, options in profiles: with running(folder, options) as (url, generator, cli): assert cli.prompt_concurrency == cli.decode_concurrency == 1 assert cli.prefill_step_size == 512 messages = [ { "role": "system", "content": "Help with coding. " + "Read the session carefully. " * 20, }, { "role": "user", "content": "Remember this context. " + "sample words " * 40, }, {"role": "assistant", "content": "Previous response."}, {"role": "user", "content": "Continue briefly."}, ] row = {"name": name, "requests": []} for index in range(12): result = request( url, messages, stream=index == 10, seeded=index == 11 ) assert generator.generation_available() row["requests"].append(result) messages.extend( [ {"role": "assistant", "content": "!!!!"}, { "role": "user", "content": f"Another short answer {index}.", }, ] ) cache = generator.prompt_cache assert cache.insertions > 0 row.update( byte_limit=cache.max_bytes, peak_retained_bytes=cache.peak_bytes, retained_sequences=len(cache), insertions=cache.insertions, ) nonstream = [r for r in row["requests"] if "cached_tokens" in r] if name == "no_retention": assert cache.peak_bytes == 0 assert all(r["cached_tokens"] == 0 for r in nonstream) else: assert any(r["cached_tokens"] > 0 for r in nonstream[1:]) assert 0 < cache.max_bytes < 1 << 63 if name == "explicit_bytes": assert cache.max_bytes == 100000 report["profiles"].append(row) report["peak_mlx_bytes"] = mx.get_peak_memory() report["status"] = "PASS" args.output.write_text(json.dumps(report, indent=2) + "\n") print(json.dumps(report, indent=2)) if __name__ == "__main__": main()