Support known speaker hints in REST API

This commit is contained in:
Aaron Boxer
2026-06-01 19:46:41 -04:00
committed by Aaron Boxer
parent 32c1b18c9f
commit 44940e2834
5 changed files with 180 additions and 28 deletions
+49
View File
@@ -8,6 +8,7 @@ class TestSpeakerDiarizer(unittest.TestCase):
def _make_diarizer(self, **kwargs):
from whisper_live.diarization import SpeakerDiarizer
d = SpeakerDiarizer(**kwargs)
# Mock the embedding model to return deterministic embeddings
d._model = MagicMock()
@@ -82,8 +83,26 @@ class TestSpeakerDiarizer(unittest.TestCase):
self.assertEqual(len(d.speakers), 0)
self.assertEqual(d._speaker_count, 0)
def test_enroll_speaker_uses_known_name(self):
d = self._make_diarizer(similarity_threshold=0.8)
self._set_embedding(d, [1.0, 0.0, 0.0])
audio = np.zeros(16000, dtype=np.float32)
self.assertTrue(d.enroll_speaker("Alice", audio))
self._set_embedding(d, [0.99, 0.01, 0.0])
speaker = d.identify_speaker(audio)
self.assertEqual(speaker, "Alice")
def test_speaker_names_label_new_speakers(self):
d = self._make_diarizer(speaker_names=["Alice"])
self._set_embedding(d, [1.0, 0.0, 0.0])
audio = np.zeros(16000, dtype=np.float32)
speaker = d.identify_speaker(audio)
self.assertEqual(speaker, "Alice")
def test_import_error_without_pyannote(self):
from whisper_live.diarization import SpeakerDiarizer
d = SpeakerDiarizer()
with patch.dict("sys.modules", {"pyannote": None, "pyannote.audio": None}):
with self.assertRaises(ImportError):
@@ -100,8 +119,10 @@ class TestDiarizationInBase(unittest.TestCase):
def __init__(self, **kwargs):
super().__init__(**kwargs)
self.language = "en"
def transcribe_audio(self, input_sample):
return None
def handle_transcription_output(self, result, duration):
pass
@@ -148,5 +169,33 @@ class TestDiarizationInBase(unittest.TestCase):
mock_diarizer.identify_speaker.assert_called_once()
class TestRestDiarizationHelpers(unittest.TestCase):
def test_normalize_form_list_accepts_repeated_or_comma_separated_values(self):
from whisper_live.server import TranscriptionServer
self.assertEqual(
TranscriptionServer._normalize_form_list(["Alice,Bob", "Carol"]),
["Alice", "Bob", "Carol"],
)
def test_speaker_labels_for_segments(self):
from whisper_live.server import TranscriptionServer
segment = MagicMock()
segment.start = 0.0
segment.end = 1.0
diarizer = MagicMock()
diarizer.identify_speaker.return_value = "Alice"
speakers = TranscriptionServer._speaker_labels_for_segments(
[segment],
np.zeros(16000, dtype=np.float32),
diarizer,
)
self.assertEqual(speakers, {0: "Alice"})
diarizer.identify_speaker.assert_called_once()
if __name__ == "__main__":
unittest.main()
+5 -15
View File
@@ -343,11 +343,6 @@ class TestStreamTranscription(unittest.TestCase):
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()
@@ -504,10 +499,9 @@ class TestRESTAPIParamWarnings(unittest.TestCase):
def setUpClass(cls):
"""Build a FastAPI test app by extracting the endpoint definition."""
import logging
from fastapi import FastAPI, UploadFile, Form
from fastapi import FastAPI, UploadFile, Form, File
from fastapi.testclient import TestClient
from typing import Optional, List
from starlette.responses import PlainTextResponse, JSONResponse
app = FastAPI()
@@ -523,16 +517,12 @@ class TestRESTAPIParamWarnings(unittest.TestCase):
chunking_strategy: Optional[str] = Form(default=None),
include: Optional[List[str]] = Form(default=None),
known_speaker_names: Optional[List[str]] = Form(default=None),
known_speaker_references: Optional[List[str]] = Form(default=None),
known_speaker_references: Optional[List[UploadFile]] = File(default=None),
stream: bool = Form(default=False),
):
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:
@@ -565,17 +555,17 @@ class TestRESTAPIParamWarnings(unittest.TestCase):
ignored = resp.json()["ignored"]
self.assertTrue(any("include" in p for p in ignored))
def test_known_speaker_names_warning(self):
def test_known_speaker_names_supported(self):
resp = self._post(known_speaker_names="alice")
self.assertEqual(resp.status_code, 200)
ignored = resp.json()["ignored"]
self.assertTrue(any("known_speaker_names" in p for p in ignored))
self.assertFalse(any("known_speaker_names" in p for p in ignored))
def test_multiple_ignored_params(self):
resp = self._post(chunking_strategy="auto", known_speaker_names="bob")
self.assertEqual(resp.status_code, 200)
ignored = resp.json()["ignored"]
self.assertGreaterEqual(len(ignored), 2)
self.assertEqual(len(ignored), 1)
class TestAPIKeyAuth(unittest.TestCase):