import asyncio from fastapi.testclient import TestClient import app.dependencies as deps from app.main import app from app.providers.llm.base import sse_delta, LLMProvider from app.pipeline.sentence_chunker import SentenceChunker client = TestClient(app) def test_sentence_chunker_incremental(): ch = SentenceChunker() emitted = [] for tok in ["Hallo", " Anna", ". ", "Wie", " geht", " es", "? ", "Tschuess"]: emitted += ch.feed(tok) assert emitted == ["Hallo Anna.", "Wie geht es?"] assert ch.flush() == "Tschuess" assert SentenceChunker().feed("Eins. Zwei! Drei? Vier") == ["Eins.", "Zwei!", "Drei?"] def test_sentence_chunker_keeps_ordinals_and_abbreviations(): # "1." (Ziffer+Punkt) darf KEINE Satzgrenze sein. ch = SentenceChunker() assert ch.feed("Am 1. Mai ist frei. ") == ["Am 1. Mai ist frei."] # Abkuerzung "z. B." darf den Satz nicht zerschneiden. ch2 = SentenceChunker() assert ch2.feed("Obst, z. B. Äpfel und Birnen. ") == ["Obst, z. B. Äpfel und Birnen."] def test_sse_delta_parsing(): assert sse_delta('data: {"choices":[{"delta":{"content":"Hal"}}]}') == "Hal" assert sse_delta("data: [DONE]") is None assert sse_delta("") is None assert sse_delta(": keep-alive") is None assert sse_delta('data: {"choices":[{"delta":{}}]}') is None assert sse_delta("data: nicht-json") is None def test_base_stream_default_yields_full_completion(): class P(LLMProvider): async def complete(self, text, history=None, session_id=None): return "ganze Antwort" async def run(): return [delta async for delta in P().stream("x")] assert asyncio.run(run()) == ["ganze Antwort"] def _install_streaming(monkeypatch, tokens): class StreamLLM: async def complete(self, text, history=None, session_id=None): return "".join(tokens) async def stream(self, text, history=None, session_id=None): for tok in tokens: yield tok class StubTTS: async def synthesize(self, text, voice=None, audio_format="pcm"): return b"AUD" monkeypatch.setitem(deps.LLM_REGISTRY, "stream", lambda s: StreamLLM()) monkeypatch.setitem(deps.TTS_REGISTRY, "stub", lambda s: StubTTS()) return {"llm_provider": "stream", "tts_provider": "stub", "output_endpoint": "loopback"} def test_ws_stream_emits_token_events(monkeypatch): base = _install_streaming(monkeypatch, ["Gu", "ten ", "Tag"]) with client.websocket_connect("/ws/chat") as ws: ws.send_json({"text": "Hallo", "stream": True, **base}) assert ws.receive_json()["type"] == "ack" tokens = [] event = ws.receive_json() while event["type"] == "token": tokens.append(event["text"]) event = ws.receive_json() assert tokens == ["Gu", "ten ", "Tag"] assert event["type"] == "semantic" and event["text"] == "Guten Tag" assert ws.receive_bytes() == b"AUD" assert ws.receive_json()["type"] == "done" def test_ws_without_stream_flag_has_no_tokens(monkeypatch): base = _install_streaming(monkeypatch, ["a", "b"]) with client.websocket_connect("/ws/chat") as ws: ws.send_json({"text": "Hallo", **base}) # kein stream-Flag assert ws.receive_json()["type"] == "ack" assert ws.receive_json()["type"] == "semantic" # direkt, keine token-Events def test_ws_stream_fallback_for_nonstreaming_llm(monkeypatch): class OnlyComplete: async def complete(self, text, history=None, session_id=None): return "komplett" class StubTTS: async def synthesize(self, text, voice=None, audio_format="pcm"): return b"X" monkeypatch.setitem(deps.LLM_REGISTRY, "oc", lambda s: OnlyComplete()) monkeypatch.setitem(deps.TTS_REGISTRY, "stub", lambda s: StubTTS()) base = {"llm_provider": "oc", "tts_provider": "stub", "output_endpoint": "loopback"} with client.websocket_connect("/ws/chat") as ws: ws.send_json({"text": "x", "stream": True, **base}) assert ws.receive_json()["type"] == "ack" token = ws.receive_json() assert token["type"] == "token" and token["text"] == "komplett" assert ws.receive_json()["type"] == "semantic" def test_ws_audio_stream_sends_chunks_per_sentence(monkeypatch): tts_calls = [] class StreamLLM: async def complete(self, text, history=None, session_id=None): return "Satz eins. Satz zwei." async def stream(self, text, history=None, session_id=None): for tok in ["Satz ", "eins. ", "Satz ", "zwei."]: yield tok class CountTTS: async def synthesize(self, text, voice=None, audio_format="pcm"): tts_calls.append(text) return b"A" * len(tts_calls) monkeypatch.setitem(deps.LLM_REGISTRY, "stream", lambda s: StreamLLM()) monkeypatch.setitem(deps.TTS_REGISTRY, "cnt", lambda s: CountTTS()) base = {"llm_provider": "stream", "tts_provider": "cnt", "output_endpoint": "loopback"} with client.websocket_connect("/ws/chat") as ws: ws.send_json({"text": "x", "audio_stream": True, **base}) assert ws.receive_json()["type"] == "ack" audio_events = 0 event = ws.receive_json() while event["type"] != "semantic": assert event["type"] == "audio" assert ws.receive_bytes() # binärer Audio-Chunk folgt audio_events += 1 event = ws.receive_json() assert audio_events == 2 # zwei Sätze -> zwei Chunks # Kein finales Vollaudio mehr -> direkt done. assert ws.receive_json()["type"] == "done" assert len(tts_calls) == 2 # TTS pro Satz