palimpseste-max / tests /test_api.py
thefinalboss's picture
Upload tests/test_api.py with huggingface_hub
a1621b5 verified
Raw
History Blame Contribute Delete
2.54 kB
"""Tests for the FastAPI server (``examples.api_server``).
These tests use FastAPI's TestClient and don't require a running server.
Skipped if fastapi is not installed.
"""
from __future__ import annotations
import pytest
try:
from fastapi.testclient import TestClient
HAS_FASTAPI = True
except ImportError:
HAS_FASTAPI = False
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "examples"))
from palimseste.hf import HFPalimpsesteLM
from palimseste.lm import PalimpsesteConfig
import numpy as np
pytestmark = pytest.mark.skipif(not HAS_FASTAPI,
reason="fastapi not installed")
@pytest.fixture
def app_and_lm():
cfg = PalimpsesteConfig(D=2000, context_window=16, kernel_radius=100,
temperature=0.3)
lm = HFPalimpsesteLM(config=cfg, rng=np.random.default_rng(0))
text = "bonjour salut je suis palimpseste. qui es-tu je suis un cortex. merci je vous en prie."
lm.build_tokenizer(text)
lm.train_on_text(text)
from api_server import create_app
app = create_app(lm)
return app, lm
def test_health(app_and_lm):
app, _ = app_and_lm
client = TestClient(app)
r = client.get("/health")
assert r.status_code == 200
data = r.json()
assert data["status"] == "ok"
assert "D" in data
assert "|M|" in data
def test_stats(app_and_lm):
app, _ = app_and_lm
client = TestClient(app)
r = client.get("/stats")
assert r.status_code == 200
data = r.json()
assert "D" in data
assert "n_traces" in data
assert "vocab_size" in data
assert "turns" in data
def test_chat(app_and_lm):
app, _ = app_and_lm
client = TestClient(app)
r = client.post("/chat", json={"message": "bonjour"})
assert r.status_code == 200
data = r.json()
assert "response" in data
assert "elapsed_ms" in data
assert isinstance(data["response"], str)
def test_teach(app_and_lm):
app, lm = app_and_lm
client = TestClient(app)
n_before = len(lm.mem)
r = client.post("/teach", json={"question": "what color is the sky",
"answer": "blue"})
assert r.status_code == 200
data = r.json()
assert data["success"] is True
assert data["tokens_written"] > 0
assert len(lm.mem) > n_before
def test_root_html(app_and_lm):
app, _ = app_and_lm
client = TestClient(app)
r = client.get("/")
assert r.status_code == 200
assert "PALIMPSESTE" in r.text