my_voice_assistant_v3/tests/test_resilience.py

129 lines
4.2 KiB
Python
Raw Normal View History

import asyncio
import pytest
from fastapi.testclient import TestClient
import app.dependencies as deps
from app.main import app
from app.config import settings
from app.metrics import metrics
from app.providers.fallback import FallbackLLMProvider
client = TestClient(app)
def _run(coro):
return asyncio.run(coro)
# --- Fallback-Einheiten ----------------------------------------------------
def test_llm_fallback_uses_second_on_error():
class BadLLM:
async def complete(self, text, history=None, session_id=None):
raise RuntimeError("down")
class GoodLLM:
async def complete(self, text, history=None, session_id=None):
return "ok"
chain = FallbackLLMProvider("llm", [("bad", BadLLM()), ("good", GoodLLM())])
assert _run(chain.complete("x")) == "ok"
counters = metrics.snapshot()["counters"]
assert any("provider_fallback_total" in key for key in counters)
assert any('provider_error_total{module="llm",provider="bad"}' in key for key in counters)
def test_llm_fallback_all_fail_raises():
class BadLLM:
async def complete(self, text, history=None, session_id=None):
raise RuntimeError("x")
chain = FallbackLLMProvider("llm", [("a", BadLLM()), ("b", BadLLM())])
with pytest.raises(RuntimeError):
_run(chain.complete("x"))
def test_llm_stream_fallback_before_first_token():
class BadStream:
async def complete(self, text, history=None, session_id=None):
return "x"
async def stream(self, text, history=None, session_id=None):
raise RuntimeError("boom")
yield # macht die Funktion zum Generator
class GoodStream:
async def complete(self, text, history=None, session_id=None):
return "ok"
async def stream(self, text, history=None, session_id=None):
yield "he"
yield "llo"
chain = FallbackLLMProvider("llm", [("bad", BadStream()), ("good", GoodStream())])
async def collect():
return [delta async for delta in chain.stream("x")]
assert _run(collect()) == ["he", "llo"]
# --- Fallback ueber Config + Endpunkt --------------------------------------
def test_config_llm_fallback_applied(monkeypatch):
class BadLLM:
async def complete(self, text, history=None, session_id=None):
raise RuntimeError("primary down")
class GoodLLM:
async def complete(self, text, history=None, session_id=None):
return "rescued"
class StubTTS:
async def synthesize(self, text, voice=None, audio_format="pcm"):
return b"A"
monkeypatch.setitem(deps.LLM_REGISTRY, "bad", lambda s: BadLLM())
monkeypatch.setitem(deps.LLM_REGISTRY, "good", lambda s: GoodLLM())
monkeypatch.setitem(deps.TTS_REGISTRY, "t", lambda s: StubTTS())
monkeypatch.setattr(settings, "llm_fallback", "good")
resp = client.post(
"/api/chat?debug=true",
json={"text": "x", "llm_provider": "bad", "tts_provider": "t"},
)
assert resp.status_code == 200
assert resp.json()["trace"]["semantic_response"] == "rescued"
# --- Metriken --------------------------------------------------------------
def test_metrics_endpoint_records_requests_and_stages(monkeypatch):
class StubLLM:
async def complete(self, text, history=None, session_id=None):
return "hi"
class StubTTS:
async def synthesize(self, text, voice=None, audio_format="pcm"):
return b"A"
monkeypatch.setitem(deps.LLM_REGISTRY, "l", lambda s: StubLLM())
monkeypatch.setitem(deps.TTS_REGISTRY, "t", lambda s: StubTTS())
resp = client.post(
"/api/chat?debug=true", json={"text": "x", "llm_provider": "l", "tts_provider": "t"}
)
assert resp.status_code == 200
snap = client.get("/api/metrics").json()
assert any("http_requests_total" in k and "chat" in k for k in snap["counters"])
assert any('stage_duration_seconds{stage="llm"}' in k for k in snap["timers"])
assert any('stage_duration_seconds{stage="tts"}' in k for k in snap["timers"])
def test_metrics_prometheus_format(monkeypatch):
client.get("/health")
text = client.get("/api/metrics?format=prometheus").text
assert "http_requests_total" in text