SPB Git forge

spb/llm-api

Public
0commits 0branches 0releases
0 Bsize
maindefault branch
—last push
12.7 KB · 295 lines python
Raw Blame History
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