1"""Test fixtures: an app wired to a temp root with a fake runtime adapter (no MLX needed)."""23from __future__ import annotations45import asyncio6import json7import os8import struct9import threading10import time11from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer12from pathlib import Path1314import httpx15import pytest16import pytest_asyncio1718os.environ.setdefault("LLM_API_TEST", "1")192021# ---------------------------------------------------------------------------22# Fake model files23# ---------------------------------------------------------------------------2425def make_safetensors(path: Path, tensors: dict[str, tuple[str, list[int]]]) -> None:26 """Write a minimal safetensors file with zeroed data."""27 header = {}28 offset = 029 bits = {"U32": 4, "F16": 2, "BF16": 2, "F32": 4}30 for name, (dtype, shape) in tensors.items():31 n = 132 for s in shape:33 n *= s34 size = n * bits[dtype]35 header[name] = {"dtype": dtype, "shape": shape, "data_offsets": [offset, offset + size]}36 offset += size37 hb = json.dumps(header).encode()38 with open(path, "wb") as f:39 f.write(struct.pack("<Q", len(hb)))40 f.write(hb)41 # Sparse data region: the scanner only reads headers, so do not materialize gigabytes of zeros.42 if offset:43 f.seek(len(hb) + 8 + offset - 1)44 f.write(b"\0")454647def make_mlx_model(d: Path, layers: int = 4, hidden: int = 256, heads: int = 4, kv_heads: int = 2,48 bits: int = 4, max_pos: int = 32768, model_type: str = "qwen3") -> None:49 d.mkdir(parents=True, exist_ok=True)50 cfg = {"architectures": ["Qwen3ForCausalLM"], "model_type": model_type, "num_hidden_layers": layers,51 "hidden_size": hidden, "num_attention_heads": heads, "num_key_value_heads": kv_heads,52 "head_dim": hidden // heads, "max_position_embeddings": max_pos, "vocab_size": 1000,53 "quantization": {"bits": bits, "group_size": 64}, "torch_dtype": "bfloat16"}54 (d / "config.json").write_text(json.dumps(cfg))55 (d / "tokenizer_config.json").write_text(json.dumps({"chat_template": "{% for m in messages %}{{m.content}}{% endfor %}<think>"}))56 (d / "README.md").write_text("---\ntags:\n- mlx\n- 4-bit\nbase_model: Qwen/Qwen3-Test\n---\n# test")57 tensors = {}58 for i in range(layers):59 tensors[f"model.layers.{i}.mlp.up_proj.weight"] = ("U32", [hidden * 4, hidden * bits // 32])60 tensors[f"model.layers.{i}.mlp.up_proj.scales"] = ("F16", [hidden * 4, hidden // 64])61 tensors["model.embed_tokens.weight"] = ("F16", [1000, hidden])62 make_safetensors(d / "model.safetensors", tensors)636465def make_gguf(path: Path, arch: str = "llama", layers: int = 4, kv_heads: int = 2, head_dim: int = 64,66 ctx: int = 8192, file_type: int = 15, n_params: int = 1_000_000) -> None:67 """Minimal GGUF v3 with metadata + one tensor info (no data)."""68 def s(v: str) -> bytes:69 b = v.encode()70 return struct.pack("<Q", len(b)) + b7172 kv = []73 def add(key, t, val):74 kv.append(s(key) + struct.pack("<I", t) + val)75 add("general.architecture", 8, s(arch))76 add("general.file_type", 4, struct.pack("<I", file_type))77 add(f"{arch}.block_count", 4, struct.pack("<I", layers))78 add(f"{arch}.attention.head_count_kv", 4, struct.pack("<I", kv_heads))79 add(f"{arch}.attention.key_length", 4, struct.pack("<I", head_dim))80 add(f"{arch}.context_length", 4, struct.pack("<I", ctx))81 add("tokenizer.chat_template", 8, s("{{messages}}"))82 with open(path, "wb") as f:83 f.write(b"GGUF")84 f.write(struct.pack("<I", 3))85 f.write(struct.pack("<Q", 1)) # n_tensors86 f.write(struct.pack("<Q", len(kv)))87 for k in kv:88 f.write(k)89 # tensor info: name, ndim, dims, type (Q4_K=12), offset90 f.write(s("blk.0.weight"))91 f.write(struct.pack("<I", 2))92 f.write(struct.pack("<QQ", n_params // 100, 100))93 f.write(struct.pack("<I", 12))94 f.write(struct.pack("<Q", 0))95 f.write(b"\0" * (n_params // 2)) # pretend weight data (~4.5 bpw)969798# ---------------------------------------------------------------------------99# Fake worker (OpenAI-compatible echo server) used by the fake adapter100# ---------------------------------------------------------------------------101102class _FakeHandler(BaseHTTPRequestHandler):103 ready_at = 0.0104 crash_on_generate = False105106 def log_message(self, *a): # silence107 pass108109 def _json(self, code, obj):110 b = json.dumps(obj).encode()111 self.send_response(code)112 self.send_header("Content-Type", "application/json")113 self.send_header("Content-Length", str(len(b)))114 self.end_headers()115 self.wfile.write(b)116117 def do_GET(self):118 if self.path == "/health":119 st = "ready" if time.time() >= self.ready_at else "loading"120 return self._json(200, {"status": st, "memory": {"active_gb": 0.5, "cache_gb": 0.0}})121 self._json(404, {"error": {"message": "nf"}})122123 def do_POST(self):124 n = int(self.headers.get("Content-Length") or 0)125 body = json.loads(self.rfile.read(n) or b"{}")126 if self.path == "/warmup":127 return self._json(200, {"ok": True, "ttft_ms": 12.3})128 if self.path == "/shutdown":129 self._json(200, {"ok": True})130 threading.Thread(target=lambda: (time.sleep(0.1), os._exit(0)), daemon=True).start()131 return132 if self.path == "/v1/chat/completions":133 if _FakeHandler.crash_on_generate:134 os._exit(3)135 msgs = body.get("messages", [])136 text = "echo: " + str(msgs[-1].get("content", ""))137 usage = {"prompt_tokens": 7, "completion_tokens": 3, "total_tokens": 10}138 timings = {"ttft_ms": 20.0, "generation_tps": 55.5, "prompt_tps": 300.0}139 if body.get("stream"):140 self.send_response(200)141 self.send_header("Content-Type", "text/event-stream")142 self.end_headers()143 for tok in text.split(" "):144 ch = {"id": "x", "object": "chat.completion.chunk", "created": 1, "model": body.get("model"),145 "choices": [{"index": 0, "delta": {"content": tok + " "}, "finish_reason": None}]}146 self.wfile.write(f"data: {json.dumps(ch)}\n\n".encode())147 self.wfile.flush()148 time.sleep(0.01)149 fin = {"id": "x", "object": "chat.completion.chunk", "created": 1, "model": body.get("model"),150 "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}], "usage": usage, "timings": timings}151 self.wfile.write(f"data: {json.dumps(fin)}\n\ndata: [DONE]\n\n".encode())152 return153 return self._json(200, {"id": "x", "object": "chat.completion", "created": 1, "model": body.get("model"),154 "choices": [{"index": 0, "message": {"role": "assistant", "content": text}, "finish_reason": "stop"}],155 "usage": usage, "timings": timings})156 if self.path == "/v1/completions":157 return self._json(200, {"id": "c", "object": "text_completion", "created": 1, "model": body.get("model"),158 "choices": [{"index": 0, "text": " world", "finish_reason": "stop"}],159 "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, "timings": {"generation_tps": 10}})160 if self.path == "/v1/embeddings":161 inp = body.get("input")162 n = len(inp) if isinstance(inp, list) else 1163 return self._json(200, {"object": "list", "model": body.get("model"),164 "data": [{"object": "embedding", "index": i, "embedding": [0.1, 0.2, 0.3]} for i in range(n)],165 "usage": {"prompt_tokens": n, "total_tokens": n}})166 self._json(404, {"error": {"message": "nf", "type": "x", "code": "NF"}})167168169class _Server(ThreadingHTTPServer):170 daemon_threads = True171172 def server_bind(self):173 # HTTPServer.server_bind calls socket.getfqdn(), which can hang for a long time on Macs with slow174 # reverse DNS; skip it.175 import socketserver176 socketserver.TCPServer.server_bind(self)177 self.server_name = "localhost"178 self.server_port = self.server_address[1]179180181def fake_worker_main():182 import sys183 port = int(sys.argv[1])184 delay = float(sys.argv[2]) if len(sys.argv) > 2 else 0.0185 _FakeHandler.ready_at = time.time() + delay186 _FakeHandler.crash_on_generate = os.environ.get("FAKE_CRASH") == "1"187 _Server(("127.0.0.1", port), _FakeHandler).serve_forever()188189190if __name__ == "__main__":191 fake_worker_main()192193194# ---------------------------------------------------------------------------195# Fake adapter196# ---------------------------------------------------------------------------197198def install_fake_adapters(manager, load_delay: float = 0.0, crash: bool = False):199 import sys200 from llm_api.runtimes.base import RuntimeAdapter201202 class FakeAdapter(RuntimeAdapter):203 name = "mlx"204205 def available(self):206 return True207208 def build_command(self, model, port, context):209 return [sys.executable, __file__, str(port), str(load_delay)]210211 async def is_ready(self, handle, client):212 try:213 r = await client.get(f"{handle.base_url}/health", timeout=2)214 except Exception:215 return "loading", None216 return r.json().get("status", "loading"), None217218 def spawn(self, model, port, context):219 if crash:220 os.environ["FAKE_CRASH"] = "1"221 else:222 os.environ.pop("FAKE_CRASH", None)223 return super().spawn(model, port, context)224225 class FakeLlama(FakeAdapter):226 name = "llamacpp"227228 manager.adapters["mlx"] = FakeAdapter(manager.settings, manager.settings.logs_path / "workers")229 manager.adapters["llamacpp"] = FakeLlama(manager.settings, manager.settings.logs_path / "workers")230231232# ---------------------------------------------------------------------------233# Fixtures234# ---------------------------------------------------------------------------235236@pytest.fixture237def root(tmp_path: Path) -> Path:238 models = tmp_path / "models"239 make_mlx_model(models / "mlx" / "qwen" / "Qwen3-Test-4bit")240 make_mlx_model(models / "mlx" / "llama" / "Llama-Test-4bit", layers=2, model_type="llama")241 make_mlx_model(models / "mlx" / "big" / "Huge-Test-4bit", layers=400, hidden=8192, heads=64, kv_heads=8, model_type="llama")242 (models / "gguf" / "gemma").mkdir(parents=True)243 make_gguf(models / "gguf" / "gemma" / "gemma-test-Q4_K_M.gguf")244 return tmp_path245246247@pytest_asyncio.fixture248async def app(root: Path, monkeypatch):249 from llm_api.config import Settings, get_settings250 from llm_api.main import create_app251 from llm_api.models import compat252253 monkeypatch.setattr(compat, "mlx_available", lambda: True)254 monkeypatch.setattr(compat, "mlx_lm_model_types", lambda: {"qwen3", "llama"})255 monkeypatch.setattr(compat, "mlx_vlm_available", lambda: False)256 monkeypatch.setattr(compat, "mlx_vlm_model_types", lambda: set())257 import llm_api.models.scanner as sc258 monkeypatch.setattr(sc, "llamacpp_available", lambda b: True)259260 settings = Settings(LLM_API_ROOT=str(root), MODEL_ROOT=str(root / "models"), MAX_MODEL_MEMORY_GB=45, ABSOLUTE_MAX_MEMORY_GB=50,261 ADMIN_EMAIL="admin@test.local", ADMIN_PASSWORD="correct-horse-battery", SECRET_KEY="test-secret",262 LOAD_TIMEOUT_SECONDS=20, WORKER_PORT_START=18400, WORKER_PORT_END=18450, METRICS_INTERVAL_SECONDS=60,263 _env_file=None)264 get_settings.cache_clear()265 application = create_app(settings)266 yield application267268269@pytest_asyncio.fixture270async def client(app):271 from asgi_lifespan import LifespanManager272 async with LifespanManager(app, startup_timeout=60, shutdown_timeout=60):273 install_fake_adapters(app.state.manager)274 transport = httpx.ASGITransport(app=app)275 async with httpx.AsyncClient(transport=transport, base_url="http://test") as c:276 yield c277278279@pytest_asyncio.fixture280async def admin(client):281 from llm_api.auth import api_limiter, login_limiter282 login_limiter.hits.clear()283 api_limiter.hits.clear()284 r = await client.post("/api/auth/login", json={"email": "admin@test.local", "password": "correct-horse-battery"})285 assert r.status_code == 200, r.text286 client.headers["X-LLM-CSRF"] = "1"287 return client288289290@pytest_asyncio.fixture291async def api_key(admin):292 r = await admin.post("/api/keys", json={"name": "t", "scopes": ["inference"]})293 assert r.status_code == 200, r.text294 return r.json()["key"]295