import unittest

import numpy as np

from dictation.audio import AudioEventKind, AudioSegmenter


class AudioSegmenterTest(unittest.TestCase):
    def test_emits_partial_and_final_events(self):
        segmenter = AudioSegmenter(
            sample_rate=1000,
            speech_threshold=0.01,
            silence_ms=200,
            partial_interval_ms=500,
            max_utterance_seconds=5,
            pre_roll_ms=100,
        )
        silence = np.zeros(100, dtype=np.float32)
        speech = np.full(250, 0.2, dtype=np.float32)

        self.assertEqual(segmenter.feed(silence), [])
        events = segmenter.feed(speech)
        self.assertEqual(events[0].kind, AudioEventKind.SPEECH_STARTED)
        events = segmenter.feed(speech)
        self.assertIn(AudioEventKind.PARTIAL_READY, [event.kind for event in events])
        events = segmenter.feed(np.zeros(200, dtype=np.float32))
        self.assertEqual(events[-1].kind, AudioEventKind.FINAL_READY)
        self.assertFalse(segmenter.speaking)

    def test_max_duration_bounds_utterance(self):
        segmenter = AudioSegmenter(
            sample_rate=100,
            speech_threshold=0.01,
            silence_ms=500,
            partial_interval_ms=10000,
            max_utterance_seconds=3,
            pre_roll_ms=0,
        )
        events = []
        for _ in range(3):
            events.extend(
                segmenter.feed(np.full(100, 0.2, dtype=np.float32))
            )
        final = [event for event in events if event.kind == AudioEventKind.FINAL_READY]
        self.assertEqual(len(final), 1)
        self.assertLessEqual(final[0].samples.size, 300)


if __name__ == "__main__":
    unittest.main()
