summaryrefslogtreecommitdiff
path: root/tests/test_segmenter.py
blob: d1f7ee5dbdac756f7772d03148c31bfac5ca9e4b (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
from datetime import datetime, timedelta, timezone

import numpy as np

from mic_clipper.segmenter import Audio, End, Segmenter, Start


SAMPLE_RATE = 16_000
FRAME_SAMPLES = 512
FRAME_DURATION = timedelta(seconds=FRAME_SAMPLES / SAMPLE_RATE)
START = datetime(2026, 9, 29, 23, 59, 59, tzinfo=timezone.utc)


def frame(value: float = 0.0) -> np.ndarray:
    return np.full(FRAME_SAMPLES, value, dtype=np.float32)


def push(segmenter: Segmenter, index: int, speech: bool, value: float = 0.0):
    return segmenter.push(frame(value), speech, START + index * FRAME_DURATION)


def test_speech_starts_segment_with_half_second_preroll():
    segmenter = Segmenter(pre_roll_seconds=0.5, silence_seconds=1.5)

    for index in range(16):
        assert push(segmenter, index, False) == ()

    events = push(segmenter, 16, True, 1.0)

    assert len(events) == 1
    assert isinstance(events[0], Start)
    assert events[0].started_at == START + 16 * FRAME_DURATION - timedelta(seconds=0.5)
    assert events[0].samples.shape == (8_512,)
    assert np.all(events[0].samples[:8_000] == 0.0)
    assert np.all(events[0].samples[8_000:] == 1.0)


def test_short_silence_is_merged_into_active_segment():
    segmenter = Segmenter(pre_roll_seconds=0.5, silence_seconds=1.5)

    assert isinstance(push(segmenter, 0, True, 1.0)[0], Start)
    assert isinstance(push(segmenter, 1, False)[0], Audio)
    assert isinstance(push(segmenter, 2, True, 1.0)[0], Audio)
    assert segmenter.active


def test_one_point_five_seconds_of_silence_closes_segment():
    segmenter = Segmenter(pre_roll_seconds=0.5, silence_seconds=1.5)

    push(segmenter, 0, True, 1.0)
    events = ()
    for index in range(1, 48):
        events = push(segmenter, index, False)

    assert isinstance(events[0], Audio)
    assert isinstance(events[-1], End)
    assert not segmenter.active


def test_maximum_segment_rolls_without_losing_audio():
    segmenter = Segmenter(
        pre_roll_seconds=0.5,
        silence_seconds=1.5,
        max_segment_seconds=FRAME_SAMPLES * 2 / SAMPLE_RATE,
    )

    first = push(segmenter, 0, True, 1.0)
    second = push(segmenter, 1, True, 2.0)
    third = push(segmenter, 2, True, 3.0)

    assert isinstance(first[0], Start)
    assert isinstance(second[0], Audio)
    assert isinstance(third[0], End)
    assert isinstance(third[1], Start)
    assert np.all(third[1].samples == 3.0)


def test_segment_start_date_is_retained_across_midnight():
    segmenter = Segmenter(pre_roll_seconds=0.5, silence_seconds=1.5)

    events = push(segmenter, 16, True, 1.0)

    assert isinstance(events[0], Start)
    assert events[0].started_at.date().isoformat() == "2026-09-29"


def test_speech_after_a_closed_segment_keeps_terminal_silence_as_preroll():
    segmenter = Segmenter(pre_roll_seconds=0.5, silence_seconds=1.5)

    push(segmenter, 0, True, 1.0)
    for index in range(1, 48):
        push(segmenter, index, False)
    events = push(segmenter, 48, True, 1.0)

    assert isinstance(events[0], Start)
    assert events[0].samples.shape == (8_512,)