semif-api-rocm

SemIf HTTP API and rocm flake
Log | Files | Refs | README | LICENSE

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()