SaarthiAI / tests /test_api.py
parthmax24's picture
working proto 5
8b96826
Raw
History Blame Contribute Delete
6.35 kB
"""FastAPI endpoints via TestClient, with the orchestrator mocked."""
import json
from datetime import datetime, timedelta
from fastapi.testclient import TestClient
import main
from app.providers import llm
client = TestClient(main.app)
# A valid deadline relative to "now" — hardcoded dates would start tripping
# the past-deadline validation as soon as they age out.
ARRIVE_BY = (datetime.now() + timedelta(hours=20)).strftime("%Y-%m-%dT%H:%M")
def test_index_renders():
response = client.get("/")
assert response.status_code == 200
assert "Saarthi" in response.text
assert "plan-form" in response.text
def test_health():
response = client.get("/api/health")
assert response.status_code == 200
assert response.json()["status"] == "ok"
def test_plan_stream_emits_sse_events(monkeypatch):
def fake_run_plan(from_text, to_text, arrive_by, mode, **kwargs):
yield {"type": "status", "message": "working"}
yield {"type": "verdict", "data": {"risk": {"score": 10, "level": "LOW"}}}
monkeypatch.setattr(main.orchestrator, "run_plan", fake_run_plan)
response = client.get(
"/api/plan/stream",
params={"from": "Gomti Nagar", "to": "Charbagh", "arrive_by": ARRIVE_BY},
)
assert response.status_code == 200
assert response.headers["content-type"].startswith("text/event-stream")
payloads = [
json.loads(line[6:])
for line in response.text.splitlines()
if line.startswith("data: ")
]
types = [payload["type"] for payload in payloads]
assert types == ["status", "verdict", "done"]
def test_plan_stream_converts_crashes_to_error_events(monkeypatch):
def exploding_run_plan(*args, **kwargs):
raise RuntimeError("unexpected crash")
yield # pragma: no cover
monkeypatch.setattr(main.orchestrator, "run_plan", exploding_run_plan)
response = client.get(
"/api/plan/stream",
params={"from": "Aliganj", "to": "Chowk", "arrive_by": ARRIVE_BY},
)
assert response.status_code == 200
assert '"type": "error"' in response.text or '"error"' in response.text
def test_plan_stream_validates_mode():
response = client.get(
"/api/plan/stream",
params={"from": "Aliganj", "to": "Chowk", "arrive_by": ARRIVE_BY, "mode": "rocket"},
)
assert response.status_code == 422
def test_plan_stream_requires_params():
response = client.get("/api/plan/stream", params={"from": "A"})
assert response.status_code == 422
def test_ask_returns_503_without_llm(monkeypatch):
monkeypatch.setattr(llm, "available", lambda: False)
monkeypatch.setattr(main.llm, "available", lambda: False)
response = client.post("/api/ask", json={"question": "test"})
assert response.status_code == 503
def test_autocomplete_returns_suggestions(monkeypatch):
monkeypatch.setattr(
main.geocode, "autocomplete_place",
lambda q: [{"name": "Hazratganj", "address": "MG Marg, Lucknow", "lat": 26.84, "lon": 80.94}],
)
response = client.get("/api/autocomplete", params={"q": "hazra"})
assert response.status_code == 200
suggestions = response.json()["suggestions"]
assert suggestions[0]["name"] == "Hazratganj"
def test_autocomplete_requires_min_length():
response = client.get("/api/autocomplete", params={"q": "h"})
assert response.status_code == 422
def test_plan_stream_passes_coords_to_orchestrator(monkeypatch):
captured = {}
def fake_run_plan(from_text, to_text, arrive_by, mode, origin_coords=None, dest_coords=None):
captured["origin_coords"] = origin_coords
captured["dest_coords"] = dest_coords
yield {"type": "verdict", "data": {}}
monkeypatch.setattr(main.orchestrator, "run_plan", fake_run_plan)
response = client.get(
"/api/plan/stream",
params={
"from": "IIIT Lucknow", "to": "Hazratganj", "arrive_by": ARRIVE_BY,
"from_lat": 26.8003, "from_lon": 81.0242,
"to_lat": 26.8488, "to_lon": 80.9436,
},
)
assert response.status_code == 200
assert captured["origin_coords"] == (26.8003, 81.0242)
assert captured["dest_coords"] == (26.8488, 80.9436)
def test_ask_stream_emits_sse_events(monkeypatch):
monkeypatch.setattr(main.llm, "available", lambda: True)
def fake_stream(question, history=None):
yield {"type": "tool", "name": "get_festivals", "args": {}}
yield {"type": "answer", "text": "Bada Mangal tomorrow.", "provider": "gemini", "steps": []}
monkeypatch.setattr(main.orchestrator, "agent_ask_stream", fake_stream)
response = client.get(
"/api/ask/stream",
params={"question": "Any festival tomorrow?", "history": '[{"role":"user","content":"hi"}]'},
)
assert response.status_code == 200
payloads = [
json.loads(line[6:])
for line in response.text.splitlines()
if line.startswith("data: ")
]
types = [payload["type"] for payload in payloads]
assert types == ["tool", "answer", "done"]
def test_ask_stream_handles_bad_history_json(monkeypatch):
monkeypatch.setattr(main.llm, "available", lambda: True)
captured = {}
def fake_stream(question, history=None):
captured["history"] = history
yield {"type": "answer", "text": "ok", "provider": "gemini", "steps": []}
monkeypatch.setattr(main.orchestrator, "agent_ask_stream", fake_stream)
response = client.get(
"/api/ask/stream",
params={"question": "test question", "history": "{not json"},
)
assert response.status_code == 200
assert captured["history"] == []
def test_ask_stream_503_without_llm(monkeypatch):
monkeypatch.setattr(main.llm, "available", lambda: False)
response = client.get("/api/ask/stream", params={"question": "test question"})
assert response.status_code == 503
def test_ask_calls_agent(monkeypatch):
monkeypatch.setattr(main.llm, "available", lambda: True)
monkeypatch.setattr(
main.orchestrator, "agent_ask",
lambda question: {"answer": f"echo: {question}", "steps": [], "provider": "gemini"},
)
response = client.post("/api/ask", json={"question": "When should I leave?"})
assert response.status_code == 200
assert response.json()["answer"] == "echo: When should I leave?"