From 6422444017243c259d960399e08ef7e3e80b1c63 Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Dieter=20Schl=C3=BCter?= Date: Wed, 17 Jun 2026 05:19:07 +0200 Subject: [PATCH] feat: Resilienz (Fallback-Ketten) und Metriken (#5) - Fallback-Provider (app/providers/fallback.py) fuer STT/LLM/TTS: Provider-Kette der Reihe nach; Config *_FALLBACK; build_orchestrator baut Ketten (dedupliziert) - LLM-Stream-Fallback nur solange kein Token gesendet wurde - Metriken (app/metrics.py): In-Memory Counter/Timer, keine externe Dependency - HTTP-Middleware (Requests/Latenz/Status je Pfad); Pipeline-Stufen-Timing stt/llm/tts; Fallback-/Fehlerzaehler; GET /api/metrics (JSON + Prometheus) - Tests: 58 gruen (+6); Doku aktualisiert (README, BEDIENUNGSANLEITUNG, Architektur, .env.example) Co-Authored-By: Claude Opus 4.8 --- .env.example | 6 ++ BEDIENUNGSANLEITUNG.md | 27 +++++- Docs/voice-assistant-architecture.md | 6 +- README.md | 26 ++++++ app/api/metrics.py | 13 +++ app/config.py | 3 + app/core/orchestrator.py | 37 +++++--- app/dependencies.py | 32 ++++++- app/main.py | 25 +++++- app/metrics.py | 82 +++++++++++++++++ app/providers/fallback.py | 85 ++++++++++++++++++ tests/conftest.py | 4 +- tests/test_resilience.py | 129 +++++++++++++++++++++++++++ 13 files changed, 452 insertions(+), 23 deletions(-) create mode 100644 app/api/metrics.py create mode 100644 app/metrics.py create mode 100644 app/providers/fallback.py create mode 100644 tests/test_resilience.py diff --git a/.env.example b/.env.example index 491b84a..46e1324 100644 --- a/.env.example +++ b/.env.example @@ -36,3 +36,9 @@ DEFAULT_OUTPUT_ENDPOINT=local-default LOCAL_LLM_BASE_URL=http://127.0.0.1:11434/v1 LOCAL_LLM_API_KEY=dummy LOCAL_LLM_MODEL=llama3.1 + +# --- Resilienz: Fallback-Ketten (kommaseparierte Provider-Namen) ------------ +# Faellt der primaere Provider aus, uebernimmt der naechste. +# STT_FALLBACK=faster-whisper +# LLM_FALLBACK=local-openai-compatible +# TTS_FALLBACK=piper diff --git a/BEDIENUNGSANLEITUNG.md b/BEDIENUNGSANLEITUNG.md index da8b369..8a698c8 100644 --- a/BEDIENUNGSANLEITUNG.md +++ b/BEDIENUNGSANLEITUNG.md @@ -292,7 +292,30 @@ senden — die laufende Antwort wird abgebrochen (`interrupted`-Event). --- -## 11. Fehlerbehebung +## 11. Resilienz & Metriken (Betrieb) + +**Fallback bei Provider-Ausfall:** Pro Modul eine Ersatzliste setzen (in `.env`). +Fällt der primäre Provider aus, übernimmt der nächste automatisch: + +``` +LLM_FALLBACK=local-openai-compatible +STT_FALLBACK=faster-whisper +TTS_FALLBACK=piper +``` + +**Metriken ansehen:** + +```bash +curl http://localhost:8080/api/metrics # JSON +curl http://localhost:8080/api/metrics?format=prometheus +``` + +Enthält Request-Zahlen/-Laufzeiten, Pipeline-Stufen (`stt`/`llm`/`tts`) und +Fallback-/Fehlerzähler. Die Werte gelten pro laufendem Prozess. + +--- + +## 12. Fehlerbehebung | Symptom | Ursache | Lösung | |---|---|---| @@ -313,7 +336,7 @@ Logs erscheinen im Terminal, in dem `make run` läuft. Für mehr Details --- -## 12. Tests ausführen +## 13. Tests ausführen ```bash make test diff --git a/Docs/voice-assistant-architecture.md b/Docs/voice-assistant-architecture.md index 228732d..e968e41 100644 --- a/Docs/voice-assistant-architecture.md +++ b/Docs/voice-assistant-architecture.md @@ -171,6 +171,7 @@ ein No-op; `LoopbackOutput` sammelt die Chunks (testbar ohne Hardware). | `POST /api/admin/users` | Nutzer anlegen (Admin-Key) → Token einmalig | | `GET /api/me` · `PUT /api/me/prefs` | aktueller Nutzer + dauerhafte Präferenzen | | `GET/POST/DELETE /api/me/memories` | Langzeit-Erinnerungen des Nutzers | +| `GET /api/metrics` | Metriken (JSON / Prometheus) | | `WS /ws/chat` | Echtzeit-Chat (Text rein, Streaming-Events) | | `WS /ws/voice` | Echtzeit-Sprache (Audio rein → STT → Antwort) | @@ -206,7 +207,7 @@ Reihenfolge der Weiterentwicklung: 2. **(erledigt)** Cloud-Fundament: Bearer-Token-Auth, Mehrbenutzer, persistenter SQLite-Store, Mandanten-Trennung, dauerhafte Nutzer-Präferenzen. Offen: Skalierung auf gemeinsamen Store (Postgres/Redis) für mehrere Instanzen. 3. **(erledigt)** Konversationsgedächtnis: Kurzzeit-Gesprächsverlauf pro Session + Langzeit-Erinnerungen pro Nutzer (manuell gepflegt, als LLM-Kontext). Offen: **automatische** Extraktion/Zusammenfassung von Erinnerungen aus Gesprächen. 4. **(weitgehend erledigt)** Echtzeit: WebSocket-Streaming-Chat (`/ws/chat`), **Token-Level-LLM-Streaming (SSE, `stream:true`)**, **Audio-Streaming (chunked TTS satzweise, `audio_stream:true`)**, **Audio-Eingang (`/ws/voice`)**, **Barge-in/Turn-Manager (`interrupt` bricht laufende Antwort ab)** und **VAD-Aeusserungserkennung (energie-basiert, opt-in)** sind umgesetzt. Offen: **echte partielle Live-Transkripte (Streaming-STT-Dienst, wortweise)** und **WebRTC (aiortc)** — beide brauchen schwere Abhaengigkeiten/Dienste. Heute laeuft STT pro Aeusserung. -5. **Resilienz:** Fallback-Policy (remote KI fällt aus → lokaler/alternativer Provider), Metriken/Tracing. +5. **(weitgehend erledigt)** Resilienz: Fallback-Ketten je Modul (`*_FALLBACK`, Provider faellt aus → naechster) und In-Memory-Metriken (`/api/metrics`: Request/Latenz, Pipeline-Stufen, Fallback/Fehler; JSON + Prometheus). Offen: verteiltes Tracing, Alerting. 6. **Betrieb:** Kosten-/Quota-Kontrolle pro Nutzer; Notfall-/Eskalationskonzept (Senioren-Kontext). 7. **TransportRouter** als eigene lokal/remote-Achse aktivieren. @@ -224,13 +225,14 @@ voice-assistant-scaffold/ │ ├── dependencies.py # Registries, ResolvedRoute, resolve_route, Store-/Router-Singleton │ ├── store.py # Persistenz: Store-Interface + SQLiteStore (Nutzer/Sessions/Verlauf) │ ├── auth.py # Bearer-Token-Auth (require_user) + Admin-Schutz +│ ├── metrics.py # In-Memory-Metriken (Counter/Timer, JSON + Prometheus) │ ├── errors.py # RoutingError -> HTTP 422 │ ├── schemas.py # Pydantic-Modelle │ ├── api/ # health, chat, speak, transcribe, devices, sessions, config, admin, me, ws │ ├── core/ # orchestrator │ ├── audio/ # router, transport_router, vad, endpoints/input|output/* │ ├── pipeline/ # input_cleaner, spoken_response_adapter, tts_normalizer, sentence_chunker -│ └── providers/ # stt/ llm/ tts/ (openrouter + lokale Stubs) +│ └── providers/ # stt/ llm/ tts/ (openrouter + lokale Stubs) + fallback.py ├── config/ # voice-assistant.example.toml (+ lokale .toml, gitignored) ├── data/ # SQLite-DB (gitignored) ├── deploy/ # systemd unit + env-Beispiel diff --git a/README.md b/README.md index 6929aa8..483e2ad 100644 --- a/README.md +++ b/README.md @@ -16,6 +16,7 @@ Praktische Bedienung: [`BEDIENUNGSANLEITUNG.md`](BEDIENUNGSANLEITUNG.md). - **Geschichtete Konfiguration** mit Profilen (`local-dev` / `hybrid` / `cloud`) - **Routing auf jeder Ebene:** Default → Profil → Nutzer → Session → Request - **Authentifizierung** (Bearer-Token) + persistente Nutzer/Sessions (SQLite) +- **Resilienz:** Fallback-Ketten je Modul (Provider fällt aus → nächster) + Metriken - **Gesprächsgedächtnis pro Session:** Verlauf wird gespeichert und fließt ins LLM - **Langzeit-Erinnerungen pro Nutzer:** dauerhafte Fakten/Vorlieben als LLM-Kontext - **WebSocket-Streaming-Chat** (`/ws/chat`) als Echtzeit-Transport @@ -90,6 +91,7 @@ Aktive Konfiguration prüfen: `curl http://localhost:8080/api/config`. | `GET/POST/DELETE /api/me/memories` | Langzeit-Erinnerungen des Nutzers verwalten | | `WS /ws/chat` | Echtzeit-Chat über WebSocket (Text rein, Streaming-Events) | | `WS /ws/voice` | Echtzeit-Sprache (Audio rein → Transkript → Antwort) | +| `GET /api/metrics` | Metriken (JSON, oder `?format=prometheus`) | Beispiel (Sprachausgabe an den Test-Loopback, lokaler TTS-Stub): @@ -162,6 +164,30 @@ eine neue Eingabe) abbrechen — der Server stoppt das Streaming und meldet > Sprechens) und **WebRTC** sind als nächste Increments vorgesehen (siehe > Architektur-Dokument). Heute läuft STT pro Äußerung. +## Resilienz & Metriken + +**Fallback-Ketten:** Pro Modul lässt sich eine Ersatz-Provider-Liste setzen. Fällt +der primäre Provider aus (Timeout/Fehler), übernimmt transparent der nächste: + +```bash +# z. B. Cloud-LLM mit lokalem Fallback +LLM_FALLBACK=local-openai-compatible +STT_FALLBACK=faster-whisper +TTS_FALLBACK=piper +``` + +Die Kette ist `Route-Provider` + `*_FALLBACK` (dedupliziert). Erfolgreiche Fallbacks +und Provider-Fehler werden gezählt. + +**Metriken** (`GET /api/metrics`): Request-Counts/-Latenzen pro Pfad, Pipeline-Stufen +(`stt`/`llm`/`tts`), Fallback-/Fehlerzähler — als JSON oder Prometheus-Text +(`?format=prometheus`). In-Memory pro Prozess (keine externe Dependency). + +```bash +curl http://localhost:8080/api/metrics +curl http://localhost:8080/api/metrics?format=prometheus +``` + ## Authentifizierung Standardmäßig (`AUTH_ENABLED=true`) sind `chat`/`speak`/`transcribe`/`sessions`/`me` diff --git a/app/api/metrics.py b/app/api/metrics.py new file mode 100644 index 0000000..42d09a7 --- /dev/null +++ b/app/api/metrics.py @@ -0,0 +1,13 @@ +from fastapi import APIRouter, Query +from fastapi.responses import PlainTextResponse + +from app.metrics import metrics + +router = APIRouter() + + +@router.get("/metrics") +async def get_metrics(format: str = Query(default="json", description="json | prometheus")): + if format == "prometheus": + return PlainTextResponse(metrics.prometheus(), media_type="text/plain; version=0.0.4") + return metrics.snapshot() diff --git a/app/config.py b/app/config.py index 58b8926..ae9db7b 100644 --- a/app/config.py +++ b/app/config.py @@ -124,6 +124,9 @@ class Settings(BaseSettings): admin_api_key: str = "" auth_enabled: bool = True history_max_messages: int = 10 + stt_fallback: str = "" # kommaseparierte Provider-Namen (Fallback-Kette) + llm_fallback: str = "" + tts_fallback: str = "" model_config = SettingsConfigDict( env_file=ENV_FILE, case_sensitive=False, extra="ignore" ) diff --git a/app/core/orchestrator.py b/app/core/orchestrator.py index 15ef63b..10755c6 100644 --- a/app/core/orchestrator.py +++ b/app/core/orchestrator.py @@ -1,5 +1,10 @@ from app.schemas import AudioChunk, PipelineTrace from app.pipeline.sentence_chunker import SentenceChunker +from app.metrics import timer, metrics + + +def _stage(name: str): + return timer("stage_duration_seconds", {"stage": name}) # Festes Ausgabeformat der TTS-Stufe (s16le PCM, 24 kHz, mono). TTS_AUDIO_FORMAT = "pcm" @@ -48,11 +53,12 @@ class Orchestrator: # input dient hier nur der Validierung/Metadaten; das Audio kommt per Upload. if input is not None: await input.capabilities() - trace.raw_transcript = await self.stt.transcribe( - audio_bytes, - fmt=fmt, - language=language, - ) + with _stage("stt"): + trace.raw_transcript = await self.stt.transcribe( + audio_bytes, + fmt=fmt, + language=language, + ) trace.cleaned_transcript = await self.input_cleaner.run( trace.raw_transcript or "" ) @@ -67,7 +73,8 @@ class Orchestrator: ): spoken = await self.spoken_adapter.run(text, language=language) normalized = await self.tts_normalizer.run(spoken, language=language) - audio = await self.tts.synthesize(normalized, voice=voice) + with _stage("tts"): + audio = await self.tts.synthesize(normalized, voice=voice) await self._emit_to_output(audio, output) return audio @@ -84,10 +91,11 @@ class Orchestrator: trace.raw_transcript = text trace.cleaned_transcript = await self.input_cleaner.run(text or "") - trace.semantic_response = await self.llm.complete( - trace.cleaned_transcript or "", - history=history, - ) + with _stage("llm"): + trace.semantic_response = await self.llm.complete( + trace.cleaned_transcript or "", + history=history, + ) if not trace.semantic_response: raise RuntimeError("LLM returned an empty response") @@ -100,10 +108,11 @@ class Orchestrator: language=language, ) - audio = await self.tts.synthesize( - trace.tts_ready_text, - voice=voice, - ) + with _stage("tts"): + audio = await self.tts.synthesize( + trace.tts_ready_text, + voice=voice, + ) await self._emit_to_output(audio, output) return trace, audio diff --git a/app/dependencies.py b/app/dependencies.py index 51c68c6..c686400 100644 --- a/app/dependencies.py +++ b/app/dependencies.py @@ -19,6 +19,11 @@ from app.providers.llm.openrouter import OpenRouterLLMProvider from app.providers.tts.openrouter import OpenRouterTTSProvider from app.providers.tts.chatterbox import ChatterboxTTSProvider from app.providers.tts.piper import PiperTTSProvider +from app.providers.fallback import ( + FallbackSTTProvider, + FallbackLLMProvider, + FallbackTTSProvider, +) from app.pipeline.input_cleaner import InputCleaner from app.pipeline.spoken_response_adapter import SpokenResponseAdapter from app.pipeline.tts_normalizer import TTSNormalizer @@ -198,11 +203,32 @@ def resolve_route( return ResolvedRoute(**resolved) +_FALLBACK_CLASS = { + "stt": FallbackSTTProvider, + "llm": FallbackLLMProvider, + "tts": FallbackTTSProvider, +} + + +def _provider_chain(registry, primary: str, fallback_csv: str, module: str, cfg: Settings): + """Baut primaeren Provider + optionale Fallback-Kette (dedupliziert, Reihenfolge erhalten).""" + names = [primary] + [n.strip() for n in (fallback_csv or "").split(",") if n.strip()] + seen, ordered = set(), [] + for name in names: + if name not in seen: + seen.add(name) + ordered.append(name) + entries = [(name, _from_registry(registry, name, module.upper(), cfg)) for name in ordered] + if len(entries) == 1: + return entries[0][1] + return _FALLBACK_CLASS[module](module, entries) + + def build_orchestrator(route: ResolvedRoute, cfg: Settings = settings) -> Orchestrator: return Orchestrator( - stt=get_stt_provider(route.stt_provider, cfg), - llm=get_llm_provider(route.llm_provider, cfg), - tts=get_tts_provider(route.tts_provider, cfg), + stt=_provider_chain(STT_REGISTRY, route.stt_provider, cfg.stt_fallback, "stt", cfg), + llm=_provider_chain(LLM_REGISTRY, route.llm_provider, cfg.llm_fallback, "llm", cfg), + tts=_provider_chain(TTS_REGISTRY, route.tts_provider, cfg.tts_fallback, "tts", cfg), input_cleaner=InputCleaner(), spoken_adapter=SpokenResponseAdapter(), tts_normalizer=TTSNormalizer(), diff --git a/app/main.py b/app/main.py index 961b113..0b12281 100644 --- a/app/main.py +++ b/app/main.py @@ -1,4 +1,8 @@ -from fastapi import FastAPI +import time + +from fastapi import FastAPI, Request + +from app.metrics import metrics from app.api.health import router as health_router from app.api.chat import router as chat_router from app.api.transcribe import router as transcribe_router @@ -8,9 +12,27 @@ from app.api.sessions import router as sessions_router from app.api.config import router as config_router from app.api.admin import router as admin_router from app.api.me import router as me_router +from app.api.metrics import router as metrics_router from app.api.ws import router as ws_router app = FastAPI(title="Voice Assistant Gateway") + + +@app.middleware("http") +async def record_metrics(request: Request, call_next): + start = time.perf_counter() + response = await call_next(request) + duration = time.perf_counter() - start + # Route-Template (z. B. /api/sessions/{session_id}/route) statt konkreter URL, + # um die Label-Kardinalitaet niedrig zu halten. + route = request.scope.get("route") + path = getattr(route, "path", request.url.path) + labels = {"method": request.method, "path": path} + metrics.inc("http_requests_total", {**labels, "status": response.status_code}) + metrics.observe("http_request_duration_seconds", duration, labels) + return response + + app.include_router(health_router) app.include_router(chat_router, prefix="/api") app.include_router(transcribe_router, prefix="/api") @@ -20,4 +42,5 @@ app.include_router(sessions_router, prefix="/api") app.include_router(config_router, prefix="/api") app.include_router(admin_router, prefix="/api") app.include_router(me_router, prefix="/api") +app.include_router(metrics_router, prefix="/api") app.include_router(ws_router) diff --git a/app/metrics.py b/app/metrics.py new file mode 100644 index 0000000..4fe21a7 --- /dev/null +++ b/app/metrics.py @@ -0,0 +1,82 @@ +"""Schlanke In-Memory-Metriken (Counter + Timer) fuer einen Prozess. + +Bewusst ohne externe Dependency. Fuer mehrere Instanzen/Prozesse spaeter durch +einen gemeinsamen Backend (z. B. Prometheus-Exporter) ersetzbar. +""" + +import threading +import time +from collections import defaultdict + + +class Metrics: + def __init__(self): + self._lock = threading.Lock() + self._counters: dict[str, float] = defaultdict(float) + self._timers: dict[str, list] = defaultdict(lambda: [0.0, 0]) # [sum, count] + + @staticmethod + def _key(name: str, labels: dict | None) -> str: + if not labels: + return name + rendered = ",".join(f'{k}="{v}"' for k, v in sorted(labels.items())) + return f"{name}{{{rendered}}}" + + def inc(self, name: str, labels: dict | None = None, value: float = 1.0) -> None: + with self._lock: + self._counters[self._key(name, labels)] += value + + def observe(self, name: str, seconds: float, labels: dict | None = None) -> None: + with self._lock: + agg = self._timers[self._key(name, labels)] + agg[0] += seconds + agg[1] += 1 + + def snapshot(self) -> dict: + with self._lock: + counters = dict(self._counters) + timers = { + key: { + "sum": agg[0], + "count": agg[1], + "avg": (agg[0] / agg[1] if agg[1] else 0.0), + } + for key, agg in self._timers.items() + } + return {"counters": counters, "timers": timers} + + def prometheus(self) -> str: + snap = self.snapshot() + lines = [] + for key, value in sorted(snap["counters"].items()): + lines.append(f"{key} {value}") + for key, agg in sorted(snap["timers"].items()): + base, _, labels = key.partition("{") + suffix = ("{" + labels) if labels else "" + lines.append(f"{base}_sum{suffix} {agg['sum']}") + lines.append(f"{base}_count{suffix} {agg['count']}") + return "\n".join(lines) + "\n" + + def reset(self) -> None: + with self._lock: + self._counters.clear() + self._timers.clear() + + +metrics = Metrics() + + +class timer: + """Context-Manager: misst die Dauer und schreibt sie als Timer-Beobachtung.""" + + def __init__(self, name: str, labels: dict | None = None): + self.name = name + self.labels = labels + + def __enter__(self): + self._start = time.perf_counter() + return self + + def __exit__(self, *exc): + metrics.observe(self.name, time.perf_counter() - self._start, self.labels) + return False diff --git a/app/providers/fallback.py b/app/providers/fallback.py new file mode 100644 index 0000000..3e7e71d --- /dev/null +++ b/app/providers/fallback.py @@ -0,0 +1,85 @@ +"""Fallback-Ketten: versuchen mehrere Provider der Reihe nach. + +Faellt der primaere Provider aus (Timeout/Fehler), wird transparent der naechste +versucht. Erfolgreicher Fallback und Provider-Fehler werden als Metrik erfasst. +""" + +from collections.abc import AsyncIterator + +from app.metrics import metrics + + +class _Chain: + def __init__(self, module: str, entries: list[tuple[str, object]]): + self.module = module + self.entries = entries # [(provider_name, provider), ...] + + def _on_error(self, name: str) -> None: + metrics.inc("provider_error_total", {"module": self.module, "provider": name}) + + def _on_fallback(self) -> None: + metrics.inc("provider_fallback_total", {"module": self.module}) + + +class FallbackSTTProvider(_Chain): + async def transcribe(self, audio_bytes, fmt, language=None) -> str: + last_exc = None + for index, (name, provider) in enumerate(self.entries): + try: + result = await provider.transcribe(audio_bytes, fmt, language=language) + if index > 0: + self._on_fallback() + return result + except Exception as exc: # noqa: BLE001 - bewusst breit fuer Resilienz + last_exc = exc + self._on_error(name) + raise last_exc + + +class FallbackLLMProvider(_Chain): + async def complete(self, text, history=None, session_id=None) -> str: + last_exc = None + for index, (name, provider) in enumerate(self.entries): + try: + result = await provider.complete(text, history=history, session_id=session_id) + if index > 0: + self._on_fallback() + return result + except Exception as exc: # noqa: BLE001 + last_exc = exc + self._on_error(name) + raise last_exc + + async def stream(self, text, history=None, session_id=None) -> AsyncIterator[str]: + last_exc = None + for index, (name, provider) in enumerate(self.entries): + produced = False + try: + async for delta in provider.stream(text, history=history, session_id=session_id): + produced = True + yield delta + if index > 0: + self._on_fallback() + return + except Exception as exc: # noqa: BLE001 + last_exc = exc + self._on_error(name) + if produced: + # Schon Token gesendet -> kein Fallback mehr moeglich. + raise + raise last_exc + + +class FallbackTTSProvider(_Chain): + async def synthesize(self, text, voice=None, audio_format="pcm") -> bytes: + last_exc = None + for index, (name, provider) in enumerate(self.entries): + try: + result = await provider.synthesize(text, voice=voice, audio_format=audio_format) + if index > 0: + self._on_fallback() + return result + except Exception as exc: # noqa: BLE001 + last_exc = exc + self._on_error(name) + raise last_exc diff --git a/tests/conftest.py b/tests/conftest.py index e1bfbef..cf9e574 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -3,16 +3,18 @@ import pytest import app.dependencies as deps from app.config import settings from app.store import SQLiteStore +from app.metrics import metrics @pytest.fixture(autouse=True) def reset_state(tmp_path, monkeypatch): - """Pro Test: frische SQLite-DB, frischer Singleton-Audio-Router, Auth aus. + """Pro Test: frische SQLite-DB, frischer Singleton-Audio-Router, Auth aus, Metriken leer. Auth-Tests schalten `settings.auth_enabled` selbst wieder ein. """ deps._store = SQLiteStore(str(tmp_path / "test.db")) deps._audio_router = None + metrics.reset() monkeypatch.setattr(settings, "auth_enabled", False) monkeypatch.setattr(settings, "admin_api_key", "") yield diff --git a/tests/test_resilience.py b/tests/test_resilience.py new file mode 100644 index 0000000..7d25933 --- /dev/null +++ b/tests/test_resilience.py @@ -0,0 +1,129 @@ +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