File size: 6,596 Bytes
66ee87e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
6012dcc
 
 
 
 
 
 
66ee87e
6012dcc
66ee87e
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
"""/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()