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"