"""GET/POST /admin/api/classifier-config (cloud_fallback field) and POST /admin/api/cloud-fallback-config. The GET side returns the ``cloud_fallback`` key from the merged config (None when unset, a dict when configured). The POST side writes or clears classifier.cloud_fallback in the local overlay atomically. """ from __future__ import annotations import sqlite3 from pathlib import Path import pytest import yaml from fastapi import FastAPI from starlette.testclient import TestClient from admin import build_router from config import load_config ROOT = Path(__file__).resolve().parent.parent SCHEMA_SQL = (ROOT / "config" / "schema.sql").read_text() ADMIN_TABLE_SQL = """ CREATE TABLE IF NOT EXISTS admin_model_overrides ( model_id TEXT NOT NULL, provider TEXT NOT NULL, availability TEXT NOT NULL, reason TEXT, updated_at TEXT NOT NULL, PRIMARY KEY (model_id, provider) ); """ def _client(tmp_path, overlay_yaml: str | None = None): (tmp_path / "config").mkdir(parents=True, exist_ok=True) config_yaml = tmp_path / "config" / "config.yaml" local_yaml = tmp_path / "config" / "config.local.yaml" config_yaml.write_text((ROOT / "config" / "config.yaml").read_text()) if overlay_yaml is not None: local_yaml.write_text(overlay_yaml) cfg = load_config(str(config_yaml)) conn = sqlite3.connect(str(tmp_path / "admin.db")) conn.row_factory = sqlite3.Row conn.executescript(SCHEMA_SQL) conn.executescript(ADMIN_TABLE_SQL) conn.close() def _db_factory() -> sqlite3.Connection: c = sqlite3.connect(str(tmp_path / "admin.db")) c.row_factory = sqlite3.Row return c router = build_router(cfg, _db_factory, base_dir=str(tmp_path)) app = FastAPI() app.include_router(router, prefix="/admin") return TestClient(app), config_yaml, local_yaml # --- GET: cloud_fallback field on the classifier-config endpoint ---------- def test_get_cloud_fallback_is_none_when_unset(tmp_path): """GET /admin/api/classifier-config reports cloud_fallback as None when the base config has it unset (no overlay overrides it).""" client, _config_yaml, _local_yaml = _client(tmp_path) body = client.get("/admin/api/classifier-config").json() assert body["cloud_fallback"] is None def test_get_cloud_fallback_reports_configured_block(tmp_path): """When the overlay has classifier.cloud_fallback set, GET returns the full configured block.""" client, _config_yaml, _local_yaml = _client( tmp_path, overlay_yaml=( "classifier:\n" " cloud_fallback:\n" " base_url: https://api.neuralwatt.com/v1\n" " model: deepseek-v4-flash\n" " api_key_env: NEURALWATT_API_KEY\n" " timeout_seconds: 2\n" " max_output_tokens: 1024\n" ), ) body = client.get("/admin/api/classifier-config").json() cf = body["cloud_fallback"] assert cf["base_url"] == "https://api.neuralwatt.com/v1" assert cf["model"] == "deepseek-v4-flash" assert cf["api_key_env"] == "NEURALWATT_API_KEY" assert cf["timeout_seconds"] == 2 assert cf["max_output_tokens"] == 1024 # --- POST /admin/api/cloud-fallback-config -------------------------------- def test_post_writes_valid_block_atomically(tmp_path): """POST a full cloud_fallback block writes all five fields as a nested block under classifier.cloud_fallback -- NOT five separate sibling classifier.* keys.""" client, _config_yaml, local_yaml = _client(tmp_path) resp = client.post( "/admin/api/cloud-fallback-config", json={ "cloud_fallback": { "base_url": "https://api.neuralwatt.com/v1", "model": "deepseek-v4-flash", "api_key_env": "NEURALWATT_API_KEY", "timeout_seconds": 5, "max_output_tokens": 2048, } }, ) assert resp.status_code == 200 written = yaml.safe_load(local_yaml.read_text()) # The block is nested, not spread cf = written["classifier"]["cloud_fallback"] assert cf["base_url"] == "https://api.neuralwatt.com/v1" assert cf["model"] == "deepseek-v4-flash" assert cf["api_key_env"] == "NEURALWATT_API_KEY" assert cf["timeout_seconds"] == 5 assert cf["max_output_tokens"] == 2048 # Assert there are NOT five separate sibling classifier.* keys classifier_keys = list(written["classifier"].keys()) assert "cloud_fallback" in classifier_keys # No individual classifier.cloud_fallback.* keys spread at classifier level assert "cloud_fallback_base_url" not in classifier_keys assert "cloud_fallback_model" not in classifier_keys assert "cloud_fallback_api_key_env" not in classifier_keys assert "cloud_fallback_timeout_seconds" not in classifier_keys assert "cloud_fallback_max_output_tokens" not in classifier_keys def test_post_rejects_present_but_blank_block(tmp_path): """A cloud_fallback block with empty base_url and model is rejected with 422, and the overlay is NOT touched.""" client, _config_yaml, local_yaml = _client(tmp_path) resp = client.post( "/admin/api/cloud-fallback-config", json={"cloud_fallback": {"base_url": "", "model": ""}}, ) assert resp.status_code == 422 assert not local_yaml.exists(), "a rejected save must not touch the overlay" def test_post_rejects_block_with_only_model_blank(tmp_path): """A cloud_fallback block with base_url set but model blank is rejected.""" client, _config_yaml, local_yaml = _client(tmp_path) resp = client.post( "/admin/api/cloud-fallback-config", json={ "cloud_fallback": { "base_url": "https://api.neuralwatt.com/v1", "model": "", } }, ) assert resp.status_code == 422 assert not local_yaml.exists(), "a rejected save must not touch the overlay" def test_post_clearing_removes_cloud_fallback_from_overlay(tmp_path): """POST cloud_fallback=null after it was previously set removes the classifier.cloud_fallback key from the overlay.""" client, _config_yaml, local_yaml = _client(tmp_path) # First, set a valid block client.post( "/admin/api/cloud-fallback-config", json={ "cloud_fallback": { "base_url": "https://api.neuralwatt.com/v1", "model": "deepseek-v4-flash", } }, ) # Now clear it resp = client.post( "/admin/api/cloud-fallback-config", json={"cloud_fallback": None}, ) assert resp.status_code == 200 written = yaml.safe_load(local_yaml.read_text()) assert "cloud_fallback" not in written.get("classifier", {}) def test_post_response_includes_restart_message(tmp_path): """The POST response mentions 'restart' in its message.""" client, _config_yaml, _local_yaml = _client(tmp_path) resp = client.post( "/admin/api/cloud-fallback-config", json={ "cloud_fallback": { "base_url": "https://api.neuralwatt.com/v1", "model": "deepseek-v4-flash", } }, ) assert "restart" in resp.json()["message"].lower()