Files
meet/src/summary/tests/unit/test_user_assign.py
T
leo 5ba1885411 (transcription) fix broken speaker assignment tests
Fix broken speaker assignement tests following #1522.
2026-07-22 14:23:10 +02:00

706 lines
25 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""Tests for the speaker-to-user assignment service."""
import math
from datetime import datetime
from summary.core import user_assign
from summary.core.user_assign import (
AssignmentResult,
Interval,
SpeakerAssignment,
_build_speaker_timelines,
_merge_intervals,
_overlap_duration,
_total_duration,
resolve_speaker_identities,
)
def make_transcription(segments: list | None = None) -> dict:
"""Build a WhisperX transcription dict, matching WhisperXResponse.model_dump()."""
return {"segments": segments or []}
RECORDING_START = datetime.fromisoformat("2026-03-17T15:30:33.000001")
RECORDING_END = datetime.fromisoformat("2026-03-17T15:31:33.000001")
METADATA_SINGLE_USER = {
"events": [
{
"participant_id": "da8d39ff-3b1c-4e8d-9a70-c630c9871bcb",
"type": "participant_connected",
"timestamp": "2026-03-17T15:30:33.000001",
},
{
"participant_id": "da8d39ff-3b1c-4e8d-9a70-c630c9871bcb",
"type": "speech_start",
"timestamp": "2026-03-17T15:30:36.039456",
},
{
"participant_id": "da8d39ff-3b1c-4e8d-9a70-c630c9871bcb",
"type": "speech_end",
"timestamp": "2026-03-17T15:30:36.589114",
},
{
"participant_id": "da8d39ff-3b1c-4e8d-9a70-c630c9871bcb",
"type": "speech_start",
"timestamp": "2026-03-17T15:30:38.887518",
},
{
"participant_id": "da8d39ff-3b1c-4e8d-9a70-c630c9871bcb",
"type": "speech_end",
"timestamp": "2026-03-17T15:30:39.438141",
},
{
"participant_id": "da8d39ff-3b1c-4e8d-9a70-c630c9871bcb",
"type": "participant_disconnected",
"timestamp": "2026-03-17T15:30:43.223255",
},
],
"participants": [
{
"participantId": "da8d39ff-3b1c-4e8d-9a70-c630c9871bcb",
"name": "cameledev",
}
],
}
DIARIZATION_SINGLE_SPEAKER = make_transcription(
segments=[
{
"start": 1.363,
"end": 3.545,
"text": " The stale smell.",
"speaker": "SPEAKER_00",
"words": [
{"word": "The", "start": 1.363, "end": 1.8},
{"word": "stale", "start": 1.8, "end": 2.7},
{"word": "smell.", "start": 2.7, "end": 3.545},
],
},
{
"start": 4.466,
"end": 6.247,
"text": "It takes heat.",
"speaker": "SPEAKER_00",
"words": [
{"word": "It", "start": 4.466, "end": 4.7},
{"word": "takes", "start": 4.7, "end": 5.5},
{"word": "heat.", "start": 5.5, "end": 6.247},
],
},
],
)
USER_ID = "da8d39ff-3b1c-4e8d-9a70-c630c9871bcb"
class TestMergeIntervals:
"""Tests for _merge_intervals."""
def test_empty(self):
"""Empty input returns empty list."""
assert _merge_intervals([]) == []
def test_no_overlap(self):
"""Non-overlapping intervals stay separate."""
result = _merge_intervals([Interval(1, 2), Interval(3, 4)])
assert len(result) == 2
def test_overlap(self):
"""Overlapping intervals are merged."""
result = _merge_intervals([Interval(1, 3), Interval(2, 4)])
assert len(result) == 1
assert result[0].start == 1
assert result[0].end == 4
def test_adjacent(self):
"""Adjacent intervals are merged."""
result = _merge_intervals([Interval(1, 2), Interval(2, 3)])
assert len(result) == 1
assert result[0].start == 1
assert result[0].end == 3
def test_unsorted(self):
"""Unsorted input is sorted before merging."""
result = _merge_intervals([Interval(5, 6), Interval(1, 2)])
assert len(result) == 2
assert result[0].start == 1
class TestOverlapDuration:
"""Tests for _overlap_duration."""
def test_no_overlap(self):
"""Disjoint intervals have zero overlap."""
a = [Interval(1, 2)]
b = [Interval(3, 4)]
assert math.isclose(_overlap_duration(a, b), 0.0)
def test_full_overlap(self):
"""Identical intervals have full overlap."""
a = [Interval(1, 3)]
b = [Interval(1, 3)]
assert math.isclose(_overlap_duration(a, b), 2.0)
def test_partial_overlap(self):
"""Partially overlapping intervals."""
a = [Interval(1, 3)]
b = [Interval(2, 4)]
assert math.isclose(_overlap_duration(a, b), 1.0)
def test_multiple_intervals(self):
"""Multiple intervals with a spanning interval."""
a = [Interval(1, 3), Interval(5, 7)]
b = [Interval(2, 6)]
assert math.isclose(_overlap_duration(a, b), 2.0)
def test_empty(self):
"""Empty input yields zero overlap."""
assert math.isclose(_overlap_duration([], [Interval(1, 2)]), 0.0)
assert math.isclose(_overlap_duration([Interval(1, 2)], []), 0.0)
class TestTotalDuration:
"""Tests for _total_duration."""
def test_basic(self):
"""Sum of durations for multiple intervals."""
ivs = [Interval(0, 1), Interval(2, 5)]
assert math.isclose(_total_duration(ivs), 4.0)
def test_empty(self):
"""Empty input returns zero."""
assert math.isclose(_total_duration([]), 0.0)
class TestBuildSpeakerTimelines:
"""Tests for _build_speaker_timelines."""
def test_segment_without_words_falls_back_to_segment_bounds(self):
"""Segments missing a `words` key use the segment start/end as one interval."""
transcription = make_transcription(
segments=[{"start": 1.5, "end": 3.5, "speaker": "SPEAKER_00"}],
)
result = _build_speaker_timelines(transcription)
assert result == {"SPEAKER_00": [Interval(1.5, 3.5)]}
def test_segment_with_only_none_word_timestamps_falls_back(self):
"""If every word has None start/end, fall back to segment bounds."""
transcription = make_transcription(
segments=[
{
"start": 1.0,
"end": 4.0,
"speaker": "SPEAKER_00",
"words": [
{"word": "hi", "start": None, "end": None},
{"word": "there", "start": None, "end": None},
],
},
],
)
result = _build_speaker_timelines(transcription)
assert result == {"SPEAKER_00": [Interval(1.0, 4.0)]}
def test_short_words_only_uses_segment_start_and_last_word_end(self):
"""With no overly long words, the interval runs segment start to end."""
transcription = make_transcription(
segments=[
{
"start": 1.0,
"end": 5.0,
"speaker": "SPEAKER_00",
"words": [
{"word": "a", "start": 1.0, "end": 1.3},
{"word": "b", "start": 1.4, "end": 1.7},
{"word": "c", "start": 1.8, "end": 2.1},
],
},
],
)
result = _build_speaker_timelines(transcription)
# Tail: min(2.1, 1.8 + 1.0) = 2.1
assert result == {"SPEAKER_00": [Interval(1.0, 2.1)]}
def test_long_word_caps_interval_at_max_duration(self):
"""A word longer than the max-word-duration cap truncates the segment."""
max_word_duration = (
user_assign.settings.resolve_speaker_identities_max_word_duration
)
transcription = make_transcription(
segments=[
{
"start": 0.0,
"end": max_word_duration + 7,
"speaker": "SPEAKER_00",
"words": [
{
"word": "pause",
"start": 0.0,
"end": max_word_duration + 7,
},
],
},
],
)
result = _build_speaker_timelines(transcription)
assert result == {"SPEAKER_00": [Interval(0.0, max_word_duration)]}
def test_long_word_in_middle_splits_segment(self):
"""Short words around a long word produce two intervals (before-cap + after)."""
transcription = make_transcription(
segments=[
{
"start": 0.0,
"end": 20.0,
"speaker": "SPEAKER_00",
"words": [
{"word": "a", "start": 0.0, "end": 0.5},
{"word": "long", "start": 1.0, "end": 15.0},
{"word": "z", "start": 18.0, "end": 18.4},
],
},
],
)
result = _build_speaker_timelines(transcription)
# First emit: (0.0, 1.0 + 1.0). Then start_time resets, picks up at "z" (18.0).
# Tail: min(18.4, 18.0 + 1.0) = 18.4. So second interval is (18.0, 18.4).
assert result == {
"SPEAKER_00": [Interval(0.0, 2.0), Interval(18.0, 18.4)],
}
def test_tail_word_is_capped_at_max_duration(self):
"""The trailing word's end is capped at word.start + max_word_duration."""
transcription = make_transcription(
segments=[
{
"start": 0.0,
"end": 50.0,
"speaker": "SPEAKER_00",
"words": [
{"word": "a", "start": 0.0, "end": 0.4},
# Last word ends inside the cap, so the cap doesn't apply.
{"word": "b", "start": 1.0, "end": 1.5},
],
},
],
)
result = _build_speaker_timelines(transcription)
# Tail: min(1.5, 1.0 + 1.0) = 1.5
assert result == {"SPEAKER_00": [Interval(0.0, 1.5)]}
def test_split_on_words_disabled_keeps_segment_as_one_interval(self, monkeypatch):
"""With splitting disabled, long words don't split the interval."""
monkeypatch.setattr(
user_assign,
"settings",
user_assign.settings.model_copy(
update={"resolve_speaker_identities_enable_split_on_words": False},
),
)
transcription = make_transcription(
segments=[
{
"start": 0.0,
"end": 20.0,
"speaker": "SPEAKER_00",
"words": [
{"word": "a", "start": 0.0, "end": 0.5},
{"word": "long", "start": 1.0, "end": 15.0},
{"word": "z", "start": 18.0, "end": 18.4},
],
},
],
)
result = _build_speaker_timelines(transcription)
# No mid-segment split; tail caps at min(18.4, 18.0 + 1.0) = 18.4.
assert result == {"SPEAKER_00": [Interval(0.0, 18.4)]}
def test_multiple_speakers_keep_separate_timelines(self):
"""Segments from different speakers populate independent timeline entries."""
transcription = make_transcription(
segments=[
{
"start": 0.0,
"end": 1.0,
"speaker": "SPEAKER_00",
"words": [{"word": "hi", "start": 0.0, "end": 0.5}],
},
{
"start": 2.0,
"end": 3.0,
"speaker": "SPEAKER_01",
"words": [{"word": "yo", "start": 2.0, "end": 2.5}],
},
],
)
result = _build_speaker_timelines(transcription)
assert result == {
"SPEAKER_00": [Interval(0.0, 0.5)],
"SPEAKER_01": [Interval(2.0, 2.5)],
}
class TestResolveSpeakerIdentities:
"""Tests for resolve_speaker_identities."""
def test_single_speaker_single_user(self):
"""Single speaker assigned to single user with low threshold."""
result = resolve_speaker_identities(
METADATA_SINGLE_USER,
DIARIZATION_SINGLE_SPEAKER,
RECORDING_START,
RECORDING_END,
overlap_threshold=0.2,
)
assert len(result.assignments) == 1
assert result.assignments[0].speaker_label == "SPEAKER_00"
assert result.assignments[0].participant_id == USER_ID
assert result.assignments[0].participant_name == "cameledev"
assert result.assignments[0].score > 0
assert result.unassigned_speakers == []
def test_no_vad_events(self):
"""Participant with no speech leaves speakers unassigned."""
metadata = {
"events": [
{
"participant_id": "user-a",
"type": "participant_connected",
"timestamp": "2026-03-17T15:30:33.000000",
},
],
"participants": [{"participantId": "user-a", "name": "Silent User"}],
}
result = resolve_speaker_identities(
metadata,
DIARIZATION_SINGLE_SPEAKER,
RECORDING_START,
RECORDING_END,
)
assert len(result.assignments) == 0
assert "SPEAKER_00" in result.unassigned_speakers
def test_multiple_speakers_same_user(self):
"""Two speakers from same mic both assigned to same user."""
metadata = {
"events": [
{
"participant_id": "user-a",
"type": "speech_start",
"timestamp": "2026-03-17T15:30:35.000000",
},
{
"participant_id": "user-a",
"type": "speech_end",
"timestamp": "2026-03-17T15:30:50.000000",
},
],
"participants": [{"participantId": "user-a", "name": "Shared Mic"}],
}
transcription = make_transcription(
segments=[
{"start": 1.0, "end": 3.0, "speaker": "SPEAKER_00"},
{"start": 5.0, "end": 7.0, "speaker": "SPEAKER_01"},
],
)
result = resolve_speaker_identities(
metadata, transcription, RECORDING_START, RECORDING_END
)
assert len(result.assignments) == 2
pids = {a.participant_id for a in result.assignments}
assert pids == {"user-a"}
def test_two_users_two_speakers(self):
"""Each speaker maps to correct user by VAD overlap."""
metadata = {
"events": [
{
"participant_id": "user-a",
"type": "speech_start",
"timestamp": "2026-03-17T15:30:34.000001",
},
{
"participant_id": "user-a",
"type": "speech_end",
"timestamp": "2026-03-17T15:30:37.000001",
},
{
"participant_id": "user-b",
"type": "speech_start",
"timestamp": "2026-03-17T15:30:38.000001",
},
{
"participant_id": "user-b",
"type": "speech_end",
"timestamp": "2026-03-17T15:30:41.000001",
},
],
"participants": [
{"participantId": "user-a", "name": "Alice"},
{"participantId": "user-b", "name": "Bob"},
],
}
transcription = make_transcription(
segments=[
{"start": 1.5, "end": 3.5, "speaker": "SPEAKER_00"},
{"start": 5.5, "end": 7.5, "speaker": "SPEAKER_01"},
],
)
result = resolve_speaker_identities(
metadata, transcription, RECORDING_START, RECORDING_END
)
assert len(result.assignments) == 2
by_speaker = {a.speaker_label: a for a in result.assignments}
assert by_speaker["SPEAKER_00"].participant_name == "Alice"
assert by_speaker["SPEAKER_01"].participant_name == "Bob"
def test_overlapping_speech_two_users(self):
"""Simultaneous speech from two users still assigns each speaker correctly."""
# user-a speaks from t=1s to t=6s, user-b speaks from t=3s to t=8s
# (3s overlap where both are speaking)
# SPEAKER_00 diarization covers t=1.55.5 (mostly user-a)
# SPEAKER_01 diarization covers t=4.07.5 (mostly user-b)
metadata = {
"events": [
{
"participant_id": "user-a",
"type": "speech_start",
"timestamp": "2026-03-17T15:30:34.000001",
},
{
"participant_id": "user-b",
"type": "speech_start",
"timestamp": "2026-03-17T15:30:36.000001",
},
{
"participant_id": "user-a",
"type": "speech_end",
"timestamp": "2026-03-17T15:30:39.000001",
},
{
"participant_id": "user-b",
"type": "speech_end",
"timestamp": "2026-03-17T15:30:41.000001",
},
],
"participants": [
{"participantId": "user-a", "name": "Alice"},
{"participantId": "user-b", "name": "Bob"},
],
}
transcription = make_transcription(
segments=[
{"start": 1.5, "end": 5.5, "speaker": "SPEAKER_00"},
{"start": 4.0, "end": 7.5, "speaker": "SPEAKER_01"},
],
)
result = resolve_speaker_identities(
metadata,
transcription,
RECORDING_START,
RECORDING_END,
overlap_threshold=0.3,
)
assert len(result.assignments) == 2
by_speaker = {a.speaker_label: a for a in result.assignments}
assert by_speaker["SPEAKER_00"].participant_name == "Alice"
assert by_speaker["SPEAKER_01"].participant_name == "Bob"
assert result.unassigned_speakers == []
def test_below_threshold(self):
"""Speaker with minimal overlap stays unassigned."""
metadata = {
"events": [
{
"participant_id": "user-a",
"type": "speech_start",
"timestamp": "2026-03-17T15:30:34.000001",
},
{
"participant_id": "user-a",
"type": "speech_end",
"timestamp": "2026-03-17T15:30:35.169950",
},
],
"participants": [{"participantId": "user-a", "name": "Brief User"}],
}
transcription = make_transcription(
segments=[
{"start": 1.0, "end": 10.0, "speaker": "SPEAKER_00"},
],
)
result = resolve_speaker_identities(
metadata,
transcription,
RECORDING_START,
RECORDING_END,
overlap_threshold=0.5,
)
assert len(result.assignments) == 0
assert "SPEAKER_00" in result.unassigned_speakers
def test_events_before_recording_start_clamped(self):
"""Speech events before recording start are clamped to t=0."""
metadata = {
"events": [
{
"participant_id": "user-a",
"type": "speech_start",
"timestamp": "2026-03-17T15:30:31.000001", # before RECORDING_START
},
{
"participant_id": "user-a",
"type": "speech_end",
"timestamp": "2026-03-17T15:30:36.000001", # after RECORDING_START
},
],
"participants": [{"participantId": "user-a", "name": "Early User"}],
}
transcription = make_transcription(
segments=[
{"start": 0.0, "end": 3.0, "speaker": "SPEAKER_00"},
],
)
result = resolve_speaker_identities(
metadata, transcription, RECORDING_START, RECORDING_END
)
assert len(result.assignments) == 1
assert result.assignments[0].participant_name == "Early User"
def test_empty_diarization(self):
"""No segments produces empty result."""
result = resolve_speaker_identities(
METADATA_SINGLE_USER,
make_transcription(segments=[]),
RECORDING_START,
RECORDING_END,
)
assert result == AssignmentResult()
def test_segment_without_speaker_ignored(self):
"""Segments missing speaker key are skipped."""
transcription = make_transcription(
segments=[
{"start": 1.0, "end": 3.0, "text": "no speaker"},
],
)
result = resolve_speaker_identities(
METADATA_SINGLE_USER, transcription, RECORDING_START, RECORDING_END
)
assert result == AssignmentResult()
def test_unclosed_speech_closed_at_recording_end(self):
"""Open speech_start without speech_end is closed at recording end."""
recording_end = datetime.fromisoformat("2026-03-17T15:30:43.000001")
metadata = {
"events": [
{
"participant_id": "user-a",
"type": "speech_start",
"timestamp": "2026-03-17T15:30:35.000001",
},
# No speech_end — participant kept speaking until recording stopped
],
"participants": [{"participantId": "user-a", "name": "Still Talking"}],
}
transcription = make_transcription(
segments=[
{"start": 2.0, "end": 9.0, "speaker": "SPEAKER_00"},
],
)
result = resolve_speaker_identities(
metadata,
transcription,
RECORDING_START,
recording_end,
overlap_threshold=0.5,
)
assert len(result.assignments) == 1
assert result.assignments[0].participant_name == "Still Talking"
assert result.unassigned_speakers == []
class TestApply:
"""Tests for AssignmentResult.apply_to."""
def test_replaces_speakers_in_segments_and_words(self):
"""Speaker labels replaced in segments, words, and word_segments."""
diarization = {
"segments": [
{
"start": 1.0,
"end": 3.0,
"text": "Hello",
"speaker": "SPEAKER_00",
"words": [
{
"word": "Hello",
"start": 1.0,
"end": 1.5,
"speaker": "SPEAKER_00",
},
],
},
{
"start": 4.0,
"end": 6.0,
"text": "World",
"speaker": "SPEAKER_01",
"words": [
{
"word": "World",
"start": 4.0,
"end": 4.5,
"speaker": "SPEAKER_01",
},
],
},
],
"word_segments": [
{"word": "Hello", "start": 1.0, "end": 1.5, "speaker": "SPEAKER_00"},
{"word": "World", "start": 4.0, "end": 4.5, "speaker": "SPEAKER_01"},
],
}
assignment = AssignmentResult(
assignments=[
SpeakerAssignment("SPEAKER_00", "id-a", "Alice", 0.9),
SpeakerAssignment("SPEAKER_01", "id-b", "Bob", 0.8),
],
)
result = assignment.apply_to(diarization)
assert result["segments"][0]["speaker"] == "Alice"
assert result["segments"][0]["words"][0]["speaker"] == "Alice"
assert result["segments"][1]["speaker"] == "Bob"
assert result["word_segments"][0]["speaker"] == "Alice"
assert result["word_segments"][1]["speaker"] == "Bob"
def test_unassigned_speakers_unchanged(self):
"""Unassigned speaker labels are left as-is."""
diarization = {
"segments": [
{"start": 1.0, "end": 3.0, "speaker": "SPEAKER_02"},
],
}
assignment = AssignmentResult(
assignments=[
SpeakerAssignment("SPEAKER_00", "id-a", "Alice", 0.9),
],
unassigned_speakers=["SPEAKER_02"],
)
result = assignment.apply_to(diarization)
assert result["segments"][0]["speaker"] == "SPEAKER_02"
def test_preserves_extra_keys(self):
"""Non-segment keys in diarization are preserved."""
diarization = {
"segments": [],
"language": "en",
"custom_field": 42,
}
result = AssignmentResult().apply_to(diarization)
assert result["language"] == "en"
assert result["custom_field"] == 42