From 8ac98dceec2e64053319d168dc796dffa2189118 Mon Sep 17 00:00:00 2001 From: Aaron Boxer Date: Fri, 17 Apr 2026 10:35:29 -0400 Subject: [PATCH] feat: add SSE streaming for REST transcription endpoint - stream=true now returns text/event-stream with per-segment SSE events - Each segment yields 'data: {json}' followed by 'data: [DONE]' - Error events streamed as 'data: {"error": ...}' - Temp files cleaned up in finally block - 5 new tests in test_server_extended.py (183 total passing) --- tests/test_server_extended.py | 153 ++++++++++++++++++++++++++++++++++ whisper_live/server.py | 76 +++++++++++++++-- 2 files changed, 224 insertions(+), 5 deletions(-) diff --git a/tests/test_server_extended.py b/tests/test_server_extended.py index 3ec4aa4..d369df2 100644 --- a/tests/test_server_extended.py +++ b/tests/test_server_extended.py @@ -325,5 +325,158 @@ class TestTranscriptionServerCleanup(unittest.TestCase): client.cleanup.assert_called_once() +class TestStreamTranscription(unittest.TestCase): + """Tests for the SSE streaming endpoint (stream=true).""" + + def _make_app(self): + """Create a FastAPI app with the transcribe endpoint that has streaming support.""" + from fastapi import FastAPI, UploadFile, Form + from fastapi.testclient import TestClient + from starlette.responses import StreamingResponse + import os + import tempfile + import shutil + + app = FastAPI() + server = TranscriptionServer() + + @app.post("/v1/audio/transcriptions") + async def transcribe( + file: UploadFile, + stream: bool = Form(default=False), + language: str = Form(default=None), + response_format: str = Form(default="json"), + ): + if stream: + return server._stream_transcription( + file, language, None, 0.0, None, None, None + ) + return {"text": "non-streamed"} + + return app + + @patch("whisper_live.server.WhisperModel") + def test_stream_returns_sse_content_type(self, mock_model_cls): + mock_seg = MagicMock() + mock_seg.id = 0 + mock_seg.start = 0.0 + mock_seg.end = 1.0 + mock_seg.text = " hello " + mock_seg.words = [] + + mock_info = MagicMock() + mock_info.language = "en" + mock_info.duration = 1.0 + + mock_model = MagicMock() + mock_model.transcribe.return_value = (iter([mock_seg]), mock_info) + mock_model_cls.return_value = mock_model + + import io + from fastapi.testclient import TestClient + + app = self._make_app() + client = TestClient(app) + resp = client.post( + "/v1/audio/transcriptions", + files={"file": ("test.wav", io.BytesIO(b"\x00" * 100), "audio/wav")}, + data={"stream": "true"}, + ) + self.assertEqual(resp.status_code, 200) + self.assertIn("text/event-stream", resp.headers.get("content-type", "")) + + @patch("whisper_live.server.WhisperModel") + def test_stream_yields_segment_and_done(self, mock_model_cls): + mock_seg = MagicMock() + mock_seg.id = 0 + mock_seg.start = 0.0 + mock_seg.end = 1.5 + mock_seg.text = " hello world " + mock_seg.words = [] + + mock_info = MagicMock() + mock_model = MagicMock() + mock_model.transcribe.return_value = (iter([mock_seg]), mock_info) + mock_model_cls.return_value = mock_model + + import io + from fastapi.testclient import TestClient + + app = self._make_app() + client = TestClient(app) + resp = client.post( + "/v1/audio/transcriptions", + files={"file": ("test.wav", io.BytesIO(b"\x00" * 100), "audio/wav")}, + data={"stream": "true"}, + ) + body = resp.text + self.assertIn('"text": "hello world"', body) + self.assertIn("[DONE]", body) + + @patch("whisper_live.server.WhisperModel") + def test_stream_multiple_segments(self, mock_model_cls): + segs = [] + for i in range(3): + s = MagicMock() + s.id = i + s.start = float(i) + s.end = float(i + 1) + s.text = f" segment {i} " + s.words = [] + segs.append(s) + + mock_info = MagicMock() + mock_model = MagicMock() + mock_model.transcribe.return_value = (iter(segs), mock_info) + mock_model_cls.return_value = mock_model + + import io + from fastapi.testclient import TestClient + + app = self._make_app() + client = TestClient(app) + resp = client.post( + "/v1/audio/transcriptions", + files={"file": ("test.wav", io.BytesIO(b"\x00" * 100), "audio/wav")}, + data={"stream": "true"}, + ) + body = resp.text + events = [line for line in body.split("\n") if line.startswith("data: ") and "[DONE]" not in line] + self.assertEqual(len(events), 3) + for i, event in enumerate(events): + data = json.loads(event.removeprefix("data: ")) + self.assertEqual(data["text"], f"segment {i}") + + @patch("whisper_live.server.WhisperModel", side_effect=RuntimeError("model error")) + def test_stream_error_yields_error_event(self, mock_model_cls): + import io + from fastapi.testclient import TestClient + + app = self._make_app() + client = TestClient(app) + resp = client.post( + "/v1/audio/transcriptions", + files={"file": ("test.wav", io.BytesIO(b"\x00" * 100), "audio/wav")}, + data={"stream": "true"}, + ) + body = resp.text + self.assertIn('"error"', body) + self.assertIn("model error", body) + + def test_non_stream_still_works(self): + import io + from fastapi.testclient import TestClient + + app = self._make_app() + client = TestClient(app) + resp = client.post( + "/v1/audio/transcriptions", + files={"file": ("test.wav", io.BytesIO(b"\x00" * 100), "audio/wav")}, + data={"stream": "false"}, + ) + self.assertEqual(resp.status_code, 200) + self.assertEqual(resp.json()["text"], "non-streamed") + + if __name__ == "__main__": unittest.main() diff --git a/whisper_live/server.py b/whisper_live/server.py index 9798618..1e9acfb 100644 --- a/whisper_live/server.py +++ b/whisper_live/server.py @@ -10,7 +10,8 @@ import tempfile from typing import Optional, List from fastapi import FastAPI, UploadFile, Form from fastapi.middleware.cors import CORSMiddleware -from starlette.responses import PlainTextResponse, JSONResponse +from fastapi.responses import JSONResponse +from starlette.responses import PlainTextResponse, JSONResponse, StreamingResponse import uvicorn from faster_whisper import WhisperModel import torch @@ -463,6 +464,58 @@ class TranscriptionServer: wl_metrics.track_connection_closed() del websocket + def _stream_transcription(self, file, language, prompt, temperature, + timestamp_granularities, + faster_whisper_custom_model_path): + """Return a StreamingResponse that yields SSE events per segment.""" + + async def _sse_generator(): + tmp_path = None + try: + suffix = os.path.splitext(file.filename)[1] or ".wav" + with tempfile.NamedTemporaryFile(delete=False, suffix=suffix) as tmp: + shutil.copyfileobj(file.file, tmp) + tmp_path = tmp.name + + device = "cuda" if torch.cuda.is_available() else "cpu" + compute_type = "float16" if device == "cuda" else "int8" + model_name = faster_whisper_custom_model_path or "small" + transcriber = WhisperModel(model_name, device=device, compute_type=compute_type) + segments, info = transcriber.transcribe( + tmp_path, + language=language, + initial_prompt=prompt, + temperature=temperature, + vad_filter=False, + word_timestamps=(timestamp_granularities and "word" in timestamp_granularities), + ) + + for seg in segments: + seg_dict = { + "id": seg.id, + "start": seg.start, + "end": seg.end, + "text": seg.text.strip(), + } + if timestamp_granularities and "word" in timestamp_granularities: + seg_dict["words"] = [ + {"word": w.word, "start": w.start, "end": w.end, "probability": w.probability} + for w in seg.words + ] + yield f"data: {json.dumps(seg_dict)}\n\n" + + yield "data: [DONE]\n\n" + wl_metrics.track_rest_request(endpoint="transcriptions_stream", status=200) + except Exception as e: + yield f"data: {json.dumps({'error': str(e)})}\n\n" + wl_metrics.track_rest_request(endpoint="transcriptions_stream", status=500) + wl_metrics.track_error("rest_stream") + finally: + if tmp_path and os.path.exists(tmp_path): + os.unlink(tmp_path) + + return StreamingResponse(_sse_generator(), media_type="text/event-stream") + def run(self, host, port=9090, @@ -581,10 +634,23 @@ class TranscriptionServer: hotwords: Optional[str] = Form(default=None), ): if stream: - wl_metrics.track_rest_request(endpoint="transcriptions", status=400) - return JSONResponse({"error": "Streaming not supported in this backend."}, status_code=400) - if chunking_strategy or known_speaker_names or known_speaker_references: - logging.warning("Diarization/chunking params ignored; not supported.") + return self._stream_transcription( + file, language, prompt, temperature, + timestamp_granularities, + faster_whisper_custom_model_path, + ) + + ignored_params = [] + if chunking_strategy: + ignored_params.append(f"chunking_strategy='{chunking_strategy}'") + if known_speaker_names: + ignored_params.append("known_speaker_names") + if known_speaker_references: + ignored_params.append("known_speaker_references") + if include: + ignored_params.append(f"include={include}") + if ignored_params: + logging.warning(f"Unsupported OpenAI params ignored: {', '.join(ignored_params)}") supported_formats = ["json", "text", "srt", "verbose_json", "vtt"] if response_format not in supported_formats: