test_llama_api.py (6565B)
1 """Offline HTTP contract tests; no model or server required.""" 2 import os 3 import sys 4 import types 5 import unittest 6 from unittest.mock import Mock, patch 7 8 from fastapi.testclient import TestClient 9 from semif_api.app import create_app 10 from semif_api.llama import BackendError, EXAMPLE 11 from test_llama import FakeBackend 12 13 14 class APITests(unittest.TestCase): 15 def setUp(self): 16 self.env = patch.dict(os.environ, {"SEMIF_BACKEND": "llama"}) 17 self.env.start() 18 self.backend = FakeBackend() 19 self.factory = patch("semif_api.app.LlamaBackend", return_value=self.backend) 20 self.factory.start() 21 self.loader = patch("semif_api.app.load_causal_model", side_effect=AssertionError("torch loaded")) 22 self.loader.start() 23 self.client = TestClient(create_app()) 24 self.client.__enter__() 25 26 def tearDown(self): 27 self.client.__exit__(None, None, None) 28 self.loader.stop() 29 self.factory.stop() 30 self.env.stop() 31 32 def test_health_does_not_contact_server(self): 33 data = self.client.get("/healthz").json() 34 self.assertEqual(data["backend"], "llama") 35 self.assertEqual(data["backend_status"], "not_checked") 36 self.assertFalse(self.backend.calls) 37 38 def test_decide_and_persistent_slots(self): 39 for _ in range(2): 40 result = self.client.post("/decide", json=EXAMPLE) 41 self.assertEqual(result.status_code, 200, result.text) 42 self.assertEqual(result.json()["choice"], "interrupt") 43 self.assertEqual(sum(path == "/detokenize" for path, _ in self.backend.calls), 3) 44 45 def test_sequential_batch(self): 46 decisions = [{k: v for k, v in EXAMPLE.items() if k != "state"}, 47 {k: v for k, v in {**EXAMPLE, "id": "second"}.items() if k != "state"}] 48 result = self.client.post("/decide-batch", json={"state": EXAMPLE["state"], "decisions": decisions}) 49 self.assertEqual(result.status_code, 200, result.text) 50 data = result.json() 51 self.assertEqual([row["id"] for row in data["results"]], [EXAMPLE["id"], "second"]) 52 self.assertEqual(data["timing"]["mode"], "llama-sequential") 53 self.assertEqual(data["timing"]["batch_size"], 2) 54 self.assertEqual(data["results"][0]["shared_timing"], data["timing"]) 55 56 def test_invalid_batches_rejected_before_inference(self): 57 for decisions in ([], [EXAMPLE, EXAMPLE], [EXAMPLE, {**EXAMPLE, "id": "bad", "options": []}]): 58 result = self.client.post("/decide-batch", json={"state": "s", "decisions": decisions}) 59 self.assertEqual(result.status_code, 422, result.text) 60 self.assertFalse(self.backend.calls) 61 62 def test_backend_error_is_502(self): 63 with patch.object(self.backend, "score", side_effect=BackendError("server unavailable")): 64 result = self.client.post("/decide", json=EXAMPLE) 65 self.assertEqual(result.status_code, 502) 66 self.assertIn("server unavailable", result.json()["detail"]) 67 68 def test_input_error_is_422(self): 69 with patch.object(self.backend, "score", side_effect=ValueError("too long")): 70 self.assertEqual(self.client.post("/decide", json=EXAMPLE).status_code, 422) 71 72 def test_ui_served(self): 73 self.assertEqual(self.client.get("/ui/").status_code, 200) 74 75 def test_plan_happy_path(self): 76 result = self.client.post("/plan", json={ 77 "id": "plan-1", 78 "prompt": "Write the rules.", 79 "transcript": [{"role": "user", "content": "Chosen action: run. Outcome: moved 1 space."}], 80 }) 81 self.assertEqual(result.status_code, 200, result.text) 82 data = result.json() 83 self.assertEqual(data["id"], "plan-1") 84 self.assertEqual(data["rules"], "1. Run moves one space.") 85 self.assertEqual(data["reasoning"], "the transcript shows…") 86 self.assertEqual(data["model"]["backend"], "llama") 87 path, payload = self.backend.calls[-1] 88 self.assertEqual(path, "/v1/chat/completions") 89 self.assertTrue(payload["chat_template_kwargs"]["enable_thinking"]) 90 91 def test_plan_serialized_with_scoring_lock(self): 92 self.client.post("/decide", json=EXAMPLE) 93 self.assertEqual(self.client.post("/plan", json={"id": "p", "prompt": "x"}).status_code, 200) 94 95 def test_plan_rejects_bad_turns(self): 96 for transcript in ([{"role": "tool", "content": "x"}], [{"role": "user"}], "not-a-list"): 97 result = self.client.post("/plan", json={"id": "p", "prompt": "x", "transcript": transcript}) 98 self.assertEqual(result.status_code, 422, result.text) 99 self.assertFalse([c for c in self.backend.calls if c[0] == "/v1/chat/completions"]) 100 101 def test_plan_backend_error_is_502(self): 102 with patch.object(self.backend, "post", side_effect=BackendError("chat offline")): 103 result = self.client.post("/plan", json={"id": "p", "prompt": "x"}) 104 self.assertEqual(result.status_code, 502) 105 self.assertIn("chat offline", result.json()["detail"]) 106 107 108 class TorchRoutingTests(unittest.TestCase): 109 def test_torch_loader_and_scorers_retained(self): 110 direct = types.ModuleType("semif_phase1.direct") 111 shared = types.ModuleType("semif_phase1.shared") 112 direct.score = Mock(return_value={"id": EXAMPLE["id"]}) 113 shared.score_shared = Mock(return_value=([{"id": EXAMPLE["id"]}], {"batch_size": 1})) 114 with patch.dict(os.environ, {"SEMIF_BACKEND": "torch"}), patch.dict( 115 sys.modules, {"semif_phase1.direct": direct, "semif_phase1.shared": shared} 116 ), patch("semif_api.app.load_causal_model", return_value=("model", "tokenizer", {"source": "test"})) as loader: 117 with TestClient(create_app()) as client: 118 self.assertEqual(client.get("/healthz").json()["backend"], "torch") 119 self.assertEqual(client.post("/decide", json=EXAMPLE).status_code, 200) 120 result = client.post("/decide-batch", json={"state": "s", "decisions": [EXAMPLE]}) 121 self.assertEqual(result.status_code, 200) 122 # Planning is text generation; the torch backend only reads logits. 123 plan = client.post("/plan", json={"id": "p", "prompt": "x"}) 124 self.assertEqual(plan.status_code, 422) 125 self.assertIn("llama", plan.json()["detail"]) 126 loader.assert_called_once() 127 direct.score.assert_called_once() 128 shared.score_shared.assert_called_once() 129 130 131 if __name__ == "__main__": 132 unittest.main()