Support known speaker hints in REST API
This commit is contained in:
@@ -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()
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user