DecisionLab / tests /test_api.py
Michael Stattelman
Add application file
6012dcc
Raw History Blame Contribute Delete
6.6 kB
"""/api/decide input errors reach the client as HTTP 422 with the validation message.
Needs fastapi, pydantic and torch (the real app imports), so it runs in the DecisionLab container.
Where those packages are missing, it is skipped with the reason printed.
"""
import importlib.util
import unittest
MISSING = [m for m in ("fastapi", "pydantic", "torch") if importlib.util.find_spec(m) is None]
@unittest.skipIf(MISSING, f"needs the container runtime; missing: {', '.join(MISSING)}")
class DecideEndpointTest(unittest.TestCase):
def setUp(self):
from fastapi import HTTPException
from app import main
self.main, self.HTTPException = main, HTTPException
def test_invalid_questions_return_422_with_message(self):
req = self.main.DecideRequest(state="x", questions={"q": {"type": "rank", "instructions": "x"}}, models=[])
with self.assertRaises(self.HTTPException) as ctx:
self.main.decide(req)
self.assertEqual(ctx.exception.status_code, 422)
self.assertEqual(ctx.exception.detail, "Question 'q': type must be choice, score or noul (got 'rank').")
def test_status_lists_the_hub_models_then_every_folder_model(self):
"""The three Hub models in the operator's order, then every model folder found, with its path, before any load."""
import os
from pathlib import Path
from app.registry import discover
st = self.main.status()
hub = ["lightdec_arthur", "lightdec_v2", "enterprise_reflux_laya_v21", "laya"]
self.assertEqual(st["order"][:4], hub)
self.assertEqual([st["models"][k]["name"] for k in hub],
["LightDec_Arthur", "LightDec_V2", "Enterprise Reflux Laya V2.1", "Laya"])
self.assertEqual([st["models"][k]["repo"] for k in hub],
["Falconsai/LightDec_Arthur", "Falconsai/LightDec_V2", "yasserrmd/enterprise-reflux-laya-v21",
"convaiinnovations/laya"])
found = [folder.name for folder, _ in discover(os.environ.get("MODELS_DIR") or "/models")]
self.assertEqual([Path(st["models"][k]["path"]).name for k in st["order"][4:]], found)
def test_importing_the_app_raises_no_deprecation_warning_of_its_own(self):
import subprocess
import sys
r = subprocess.run([sys.executable, "-W", "error::DeprecationWarning:app.main", "-c", "import app.main"],
capture_output=True, text=True)
self.assertEqual(r.returncode, 0, r.stderr[-800:])
def test_lifespan_starts_the_model_loader(self):
import asyncio
calls = []
original = self.main.load_all
self.main.load_all = lambda: calls.append("load_all")
try:
async def run():
async with self.main.lifespan(self.main.app):
pass
asyncio.run(run())
for t in __import__("threading").enumerate():
if t.name == "model-loader":
t.join(timeout=5)
finally:
self.main.load_all = original
self.assertEqual(calls, ["load_all"])
def test_default_decide_request_runs_every_listed_model(self):
from app.validation import validate_models
req = self.main.DecideRequest(state="x", questions={"q": {"type": "noul", "instructions": "x"}})
self.assertEqual(validate_models(req.models, list(self.main.BACKENDS)), self.main.status()["order"])
def test_unknown_model_is_422(self):
req = self.main.DecideRequest(state="x", questions={"q": {"type": "noul", "instructions": "x"}}, models=["zzz"])
with self.assertRaises(self.HTTPException) as ctx:
self.main.decide(req)
self.assertEqual(ctx.exception.status_code, 422)
def test_busy_server_answers_429(self):
taken = 0
while self.main.DECIDE_SLOTS.try_enter():
taken += 1
try:
req = self.main.DecideRequest(state="x", questions={"q": {"type": "noul", "instructions": "x"}}, models=[])
with self.assertRaises(self.HTTPException) as ctx:
self.main.decide(req)
self.assertEqual(ctx.exception.status_code, 429)
finally:
for _ in range(taken):
self.main.DECIDE_SLOTS.leave()
def test_untrusted_modeling_code_is_never_executed(self):
import tempfile
from pathlib import Path
from app.models import LightDecBackend
with tempfile.TemporaryDirectory() as d:
(Path(d) / "falcondec_config.json").write_text("{}")
marker = Path(d) / "ran"
(Path(d) / "falcondec_modeling.py").write_text(f"open({str(marker)!r}, 'w').write('x')\n")
b = LightDecBackend({"key": "t", "name": "T", "side": "c0", "kind": "lightdec", "source": "local",
"path": d, "path_env": "MODELS_DIR", "variant": "fp16"})
b.load()
self.assertEqual(b.status, "error")
self.assertIn("Refusing to run", b.error)
self.assertFalse(marker.exists())
def test_models_run_through_the_replaceable_step_and_the_limit_stays_outside_it(self):
calls = []
original = self.main.RUN_MODELS
self.main.RUN_MODELS = lambda state, q, keys: (calls.append(keys), {"results": {}, "server_ms": 0})[1]
try:
q = {"q": {"type": "noul", "instructions": "x"}}
self.main.run_decision("s", q, [])
self.main.run_decision("s", q, None)
finally:
self.main.RUN_MODELS = original
self.assertEqual(calls, [[], list(self.main.BACKENDS)])
taken = 0
while self.main.DECIDE_SLOTS.try_enter(): # every slot was released again
taken += 1
for _ in range(taken):
self.main.DECIDE_SLOTS.leave()
self.assertEqual(taken, int(__import__("os").getenv("MAX_PENDING_DECIDES", "4")))
def test_security_middleware_is_installed_and_docs_are_off(self):
from app.security import SecurityMiddleware
self.assertIn(SecurityMiddleware, [m.cls for m in self.main.app.user_middleware])
self.assertEqual((self.main.app.docs_url, self.main.app.redoc_url, self.main.app.openapi_url), (None, None, None))
def test_valid_questions_with_no_models_return_empty_results(self):
req = self.main.DecideRequest(state="x", questions={"q": {"type": "noul", "instructions": "x"}}, models=[])
self.assertEqual(self.main.decide(req)["results"], {})
if __name__ == "__main__":
unittest.main()