110 lines
3.5 KiB
Python
110 lines
3.5 KiB
Python
from __future__ import annotations
|
|
|
|
import sqlite3
|
|
from pathlib import Path
|
|
|
|
import pytest
|
|
from starlette.testclient import TestClient
|
|
|
|
import dispatcher
|
|
from dispatcher import Classification, app
|
|
|
|
ROOT = Path(__file__).resolve().parent.parent
|
|
SCHEMA_SQL = (ROOT / "config" / "schema.sql").read_text()
|
|
|
|
KEEP = "keep-model"
|
|
DROP = "drop-model"
|
|
|
|
|
|
@pytest.fixture
|
|
def client(tmp_path, monkeypatch):
|
|
db_path = tmp_path / "blk.db"
|
|
conn = sqlite3.connect(db_path)
|
|
conn.executescript(SCHEMA_SQL)
|
|
import admin
|
|
admin.ensure_admin_tables(conn)
|
|
for model_id in (KEEP, DROP):
|
|
conn.execute(
|
|
"""
|
|
INSERT INTO models (
|
|
model_id, provider, base_model_id, tier, context_window,
|
|
effective_context_window, max_output_tokens,
|
|
cost_per_1m_prompt, cost_per_1m_completion,
|
|
supports_vision, supports_json_mode,
|
|
latency_class, reasoning_mode, context_variant,
|
|
access_level, availability, last_updated
|
|
) VALUES (?, 'neuralwatt', ?, 2, 262128, 192500, 16384, 1.0, 2.0,
|
|
0, 1, 'standard', 'default', 'full', 'public', 'active',
|
|
'2026-08-22T00:00:00+00:00')
|
|
""",
|
|
(model_id, model_id),
|
|
)
|
|
conn.commit()
|
|
conn.close()
|
|
monkeypatch.setattr(dispatcher.cfg.database, "path", str(db_path))
|
|
monkeypatch.setattr(dispatcher.cfg.verification, "local_llm_enabled", False)
|
|
monkeypatch.setattr(dispatcher.cfg.local_vision, "enabled", False)
|
|
monkeypatch.setattr(dispatcher.cfg.session_cache, "enabled", False)
|
|
monkeypatch.setenv("NEURALWATT_API_KEY", "test-key")
|
|
monkeypatch.setattr(
|
|
dispatcher, "classify",
|
|
lambda task, context: Classification(
|
|
task_category="coding_general", task_tier=2,
|
|
required_context_tokens=100, confidence=0.9,
|
|
),
|
|
)
|
|
return TestClient(app)
|
|
|
|
|
|
def _block(client, model_id):
|
|
resp = client.post(
|
|
f"/admin/api/models/{model_id}/neuralwatt/availability",
|
|
json={"availability": "blocked"},
|
|
)
|
|
assert resp.status_code == 200, resp.text
|
|
|
|
|
|
def _listed(client):
|
|
return {m["id"] for m in client.get("/v1/models").json()["data"]}
|
|
|
|
|
|
def test_blocked_removes_model_from_v1_models(client):
|
|
assert DROP in _listed(client)
|
|
_block(client, DROP)
|
|
assert DROP not in _listed(client)
|
|
assert KEEP in _listed(client), "only the blocked model should go"
|
|
|
|
|
|
def test_blocked_refuses_a_pinned_request(client):
|
|
_block(client, DROP)
|
|
resp = client.post(
|
|
"/v1/chat/completions",
|
|
json={"model": DROP, "messages": [{"role": "user", "content": "hi"}]},
|
|
)
|
|
assert resp.status_code == 503, resp.text
|
|
assert "override" in resp.text
|
|
|
|
|
|
def test_blocked_drops_model_from_routing(client):
|
|
client.post("/route", json={"task": "write a function"}).json()
|
|
_block(client, DROP)
|
|
after = client.post("/route", json={"task": "write a function"}).json()
|
|
assert after.get("selected_model") != DROP
|
|
|
|
|
|
def test_clearing_blocked_restores_model(client):
|
|
_block(client, DROP)
|
|
assert DROP not in _listed(client)
|
|
resp = client.delete(f"/admin/api/models/{DROP}/neuralwatt/availability")
|
|
assert resp.status_code == 200, resp.text
|
|
assert DROP in _listed(client)
|
|
|
|
|
|
def test_unblocked_model_is_still_pinnable(client):
|
|
_block(client, DROP)
|
|
resp = client.post(
|
|
"/v1/chat/completions",
|
|
json={"model": KEEP, "messages": [{"role": "user", "content": "hi"}]},
|
|
)
|
|
assert resp.status_code != 503 or "override" not in resp.text
|