"""Test fixtures: an app wired to a temp root with a fake runtime adapter (no MLX needed).""" from __future__ import annotations import asyncio import json import os import struct import threading import time from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from pathlib import Path import httpx import pytest import pytest_asyncio os.environ.setdefault("LLM_API_TEST", "1") # --------------------------------------------------------------------------- # Fake model files # --------------------------------------------------------------------------- def make_safetensors(path: Path, tensors: dict[str, tuple[str, list[int]]]) -> None: """Write a minimal safetensors file with zeroed data.""" header = {} offset = 0 bits = {"U32": 4, "F16": 2, "BF16": 2, "F32": 4} for name, (dtype, shape) in tensors.items(): n = 1 for s in shape: n *= s size = n * bits[dtype] header[name] = {"dtype": dtype, "shape": shape, "data_offsets": [offset, offset + size]} offset += size hb = json.dumps(header).encode() with open(path, "wb") as f: f.write(struct.pack(" None: d.mkdir(parents=True, exist_ok=True) cfg = {"architectures": ["Qwen3ForCausalLM"], "model_type": model_type, "num_hidden_layers": layers, "hidden_size": hidden, "num_attention_heads": heads, "num_key_value_heads": kv_heads, "head_dim": hidden // heads, "max_position_embeddings": max_pos, "vocab_size": 1000, "quantization": {"bits": bits, "group_size": 64}, "torch_dtype": "bfloat16"} (d / "config.json").write_text(json.dumps(cfg)) (d / "tokenizer_config.json").write_text(json.dumps({"chat_template": "{% for m in messages %}{{m.content}}{% endfor %}"})) (d / "README.md").write_text("---\ntags:\n- mlx\n- 4-bit\nbase_model: Qwen/Qwen3-Test\n---\n# test") tensors = {} for i in range(layers): tensors[f"model.layers.{i}.mlp.up_proj.weight"] = ("U32", [hidden * 4, hidden * bits // 32]) tensors[f"model.layers.{i}.mlp.up_proj.scales"] = ("F16", [hidden * 4, hidden // 64]) tensors["model.embed_tokens.weight"] = ("F16", [1000, hidden]) make_safetensors(d / "model.safetensors", tensors) def make_gguf(path: Path, arch: str = "llama", layers: int = 4, kv_heads: int = 2, head_dim: int = 64, ctx: int = 8192, file_type: int = 15, n_params: int = 1_000_000) -> None: """Minimal GGUF v3 with metadata + one tensor info (no data).""" def s(v: str) -> bytes: b = v.encode() return struct.pack("= self.ready_at else "loading" return self._json(200, {"status": st, "memory": {"active_gb": 0.5, "cache_gb": 0.0}}) self._json(404, {"error": {"message": "nf"}}) def do_POST(self): n = int(self.headers.get("Content-Length") or 0) body = json.loads(self.rfile.read(n) or b"{}") if self.path == "/warmup": return self._json(200, {"ok": True, "ttft_ms": 12.3}) if self.path == "/shutdown": self._json(200, {"ok": True}) threading.Thread(target=lambda: (time.sleep(0.1), os._exit(0)), daemon=True).start() return if self.path == "/v1/chat/completions": if _FakeHandler.crash_on_generate: os._exit(3) msgs = body.get("messages", []) text = "echo: " + str(msgs[-1].get("content", "")) usage = {"prompt_tokens": 7, "completion_tokens": 3, "total_tokens": 10} timings = {"ttft_ms": 20.0, "generation_tps": 55.5, "prompt_tps": 300.0} if body.get("stream"): self.send_response(200) self.send_header("Content-Type", "text/event-stream") self.end_headers() for tok in text.split(" "): ch = {"id": "x", "object": "chat.completion.chunk", "created": 1, "model": body.get("model"), "choices": [{"index": 0, "delta": {"content": tok + " "}, "finish_reason": None}]} self.wfile.write(f"data: {json.dumps(ch)}\n\n".encode()) self.wfile.flush() time.sleep(0.01) fin = {"id": "x", "object": "chat.completion.chunk", "created": 1, "model": body.get("model"), "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}], "usage": usage, "timings": timings} self.wfile.write(f"data: {json.dumps(fin)}\n\ndata: [DONE]\n\n".encode()) return return self._json(200, {"id": "x", "object": "chat.completion", "created": 1, "model": body.get("model"), "choices": [{"index": 0, "message": {"role": "assistant", "content": text}, "finish_reason": "stop"}], "usage": usage, "timings": timings}) if self.path == "/v1/completions": return self._json(200, {"id": "c", "object": "text_completion", "created": 1, "model": body.get("model"), "choices": [{"index": 0, "text": " world", "finish_reason": "stop"}], "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, "timings": {"generation_tps": 10}}) if self.path == "/v1/embeddings": inp = body.get("input") n = len(inp) if isinstance(inp, list) else 1 return self._json(200, {"object": "list", "model": body.get("model"), "data": [{"object": "embedding", "index": i, "embedding": [0.1, 0.2, 0.3]} for i in range(n)], "usage": {"prompt_tokens": n, "total_tokens": n}}) self._json(404, {"error": {"message": "nf", "type": "x", "code": "NF"}}) class _Server(ThreadingHTTPServer): daemon_threads = True def server_bind(self): # HTTPServer.server_bind calls socket.getfqdn(), which can hang for a long time on Macs with slow # reverse DNS; skip it. import socketserver socketserver.TCPServer.server_bind(self) self.server_name = "localhost" self.server_port = self.server_address[1] def fake_worker_main(): import sys port = int(sys.argv[1]) delay = float(sys.argv[2]) if len(sys.argv) > 2 else 0.0 _FakeHandler.ready_at = time.time() + delay _FakeHandler.crash_on_generate = os.environ.get("FAKE_CRASH") == "1" _Server(("127.0.0.1", port), _FakeHandler).serve_forever() if __name__ == "__main__": fake_worker_main() # --------------------------------------------------------------------------- # Fake adapter # --------------------------------------------------------------------------- def install_fake_adapters(manager, load_delay: float = 0.0, crash: bool = False): import sys from llm_api.runtimes.base import RuntimeAdapter class FakeAdapter(RuntimeAdapter): name = "mlx" def available(self): return True def build_command(self, model, port, context): return [sys.executable, __file__, str(port), str(load_delay)] async def is_ready(self, handle, client): try: r = await client.get(f"{handle.base_url}/health", timeout=2) except Exception: return "loading", None return r.json().get("status", "loading"), None def spawn(self, model, port, context): if crash: os.environ["FAKE_CRASH"] = "1" else: os.environ.pop("FAKE_CRASH", None) return super().spawn(model, port, context) class FakeLlama(FakeAdapter): name = "llamacpp" manager.adapters["mlx"] = FakeAdapter(manager.settings, manager.settings.logs_path / "workers") manager.adapters["llamacpp"] = FakeLlama(manager.settings, manager.settings.logs_path / "workers") # --------------------------------------------------------------------------- # Fixtures # --------------------------------------------------------------------------- @pytest.fixture def root(tmp_path: Path) -> Path: models = tmp_path / "models" make_mlx_model(models / "mlx" / "qwen" / "Qwen3-Test-4bit") make_mlx_model(models / "mlx" / "llama" / "Llama-Test-4bit", layers=2, model_type="llama") make_mlx_model(models / "mlx" / "big" / "Huge-Test-4bit", layers=400, hidden=8192, heads=64, kv_heads=8, model_type="llama") (models / "gguf" / "gemma").mkdir(parents=True) make_gguf(models / "gguf" / "gemma" / "gemma-test-Q4_K_M.gguf") return tmp_path @pytest_asyncio.fixture async def app(root: Path, monkeypatch): from llm_api.config import Settings, get_settings from llm_api.main import create_app from llm_api.models import compat monkeypatch.setattr(compat, "mlx_available", lambda: True) monkeypatch.setattr(compat, "mlx_lm_model_types", lambda: {"qwen3", "llama"}) monkeypatch.setattr(compat, "mlx_vlm_available", lambda: False) monkeypatch.setattr(compat, "mlx_vlm_model_types", lambda: set()) import llm_api.models.scanner as sc monkeypatch.setattr(sc, "llamacpp_available", lambda b: True) settings = Settings(LLM_API_ROOT=str(root), MODEL_ROOT=str(root / "models"), MAX_MODEL_MEMORY_GB=45, ABSOLUTE_MAX_MEMORY_GB=50, ADMIN_EMAIL="admin@test.local", ADMIN_PASSWORD="correct-horse-battery", SECRET_KEY="test-secret", LOAD_TIMEOUT_SECONDS=20, WORKER_PORT_START=18400, WORKER_PORT_END=18450, METRICS_INTERVAL_SECONDS=60, _env_file=None) get_settings.cache_clear() application = create_app(settings) yield application @pytest_asyncio.fixture async def client(app): from asgi_lifespan import LifespanManager async with LifespanManager(app, startup_timeout=60, shutdown_timeout=60): install_fake_adapters(app.state.manager) transport = httpx.ASGITransport(app=app) async with httpx.AsyncClient(transport=transport, base_url="http://test") as c: yield c @pytest_asyncio.fixture async def admin(client): from llm_api.auth import api_limiter, login_limiter login_limiter.hits.clear() api_limiter.hits.clear() r = await client.post("/api/auth/login", json={"email": "admin@test.local", "password": "correct-horse-battery"}) assert r.status_code == 200, r.text client.headers["X-LLM-CSRF"] = "1" return client @pytest_asyncio.fixture async def api_key(admin): r = await admin.post("/api/keys", json={"name": "t", "scopes": ["inference"]}) assert r.status_code == 200, r.text return r.json()["key"]