Gargi_AI / tests /test_failures.py
Sameer Singh
Added
5aaf5ba
Raw
History Blame Contribute Delete
2.66 kB
from collections.abc import AsyncIterator
from gargi_ai.main import app
from gargi_ai.providers import LLMProvider, TTSProvider
from gargi_ai.schemas import TeachingPayload
class FailingLLM(LLMProvider):
async def stream_explanation(self, prompt: str) -> AsyncIterator[str]:
raise RuntimeError("provider unavailable")
yield
async def generate_artifact(self, prompt: str) -> TeachingPayload:
raise RuntimeError("not reached")
class FailingTTS(TTSProvider):
async def synthesize(self, text: str, voice: str):
raise RuntimeError("tts unavailable")
class FailingArtifactLLM(LLMProvider):
async def stream_explanation(self, prompt: str) -> AsyncIterator[str]:
yield "The teacher response still works."
async def generate_artifact(self, prompt: str) -> TeachingPayload:
raise RuntimeError("artifact model overloaded")
def create_lesson(client):
return client.post(
"/api/v1/sessions", json={"topic": "Physics"}
).json()["id"]
def test_gemini_failure_is_a_recoverable_sse_error(client):
session_id = create_lesson(client)
app.state.llm = FailingLLM()
response = client.post(
f"/api/v1/sessions/{session_id}/teach",
json={"text": "Teach me."},
)
assert response.status_code == 200
assert "PROVIDER_ERROR" in response.text
assert '"ok": false' in response.text
def test_tts_failure_keeps_lesson_payload(client):
session_id = create_lesson(client)
app.state.tts = FailingTTS()
response = client.post(
f"/api/v1/sessions/{session_id}/teach",
json={"text": "Teach me."},
)
assert "event: lesson_payload" in response.text
assert "TTS_GENERATION_FAILED" in response.text
assert 'event: done' in response.text
def test_live_artifact_failure_returns_fallback_payload(client):
session_id = create_lesson(client)
app.state.llm = FailingArtifactLLM()
with client.websocket_connect(
f"/api/v1/sessions/{session_id}/live"
) as websocket:
assert websocket.receive_json()["type"] == "ready"
websocket.send_bytes(b"fake-microphone-pcm")
assert websocket.receive_json()["type"] == "input_transcription"
assert websocket.receive_json()["type"] == "output_transcription"
assert websocket.receive_bytes() == b"fake-live-pcm"
lesson_payload = websocket.receive_json()
turn_complete = websocket.receive_json()
assert lesson_payload["type"] == "lesson_payload"
assert lesson_payload["fallback"] is True
assert len(lesson_payload["quiz"]) == 3
assert turn_complete["type"] == "turn_complete"