diff --git a/src/local_transcriber/formatter.py b/src/local_transcriber/formatter.py index b887bab..b0d8b22 100644 --- a/src/local_transcriber/formatter.py +++ b/src/local_transcriber/formatter.py @@ -1,7 +1,44 @@ +from dataclasses import dataclass from datetime import datetime from pathlib import Path -from .transcriber import TranscribeResult +from .transcriber import Segment, TranscribeResult + +_PAUSE_THRESHOLD_S = 2.0 # пауза между сегментами для разбиения на абзацы +_MAX_PARAGRAPH_S = 60.0 # максимальная длительность абзаца + + +@dataclass +class _Paragraph: + start: float + end: float + text: str + + +def _group_segments(segments: list[Segment]) -> list[_Paragraph]: + """Объединяет мелкие сегменты в абзацы по паузам и макс. длительности.""" + if not segments: + return [] + + paragraphs: list[_Paragraph] = [] + cur_start = segments[0].start + cur_end = segments[0].end + cur_texts: list[str] = [segments[0].text.strip()] + + for seg in segments[1:]: + gap = seg.start - cur_end + duration = seg.end - cur_start + if gap >= _PAUSE_THRESHOLD_S or duration > _MAX_PARAGRAPH_S: + paragraphs.append(_Paragraph(cur_start, cur_end, " ".join(cur_texts))) + cur_start = seg.start + cur_end = seg.end + cur_texts = [seg.text.strip()] + else: + cur_end = seg.end + cur_texts.append(seg.text.strip()) + + paragraphs.append(_Paragraph(cur_start, cur_end, " ".join(cur_texts))) + return paragraphs def format_timestamp(seconds: float, use_hours: bool = False) -> str: @@ -56,11 +93,11 @@ def format_transcript( lines.append("") lines.append("*Речь не обнаружена.*") else: - for seg in result.segments: - start = format_timestamp(seg.start, use_hours=use_hours) - end = format_timestamp(seg.end, use_hours=use_hours) + for para in _group_segments(result.segments): + start = format_timestamp(para.start, use_hours=use_hours) + end = format_timestamp(para.end, use_hours=use_hours) lines.append("") - lines.append(f"[{start} - {end}] {seg.text.strip()}") + lines.append(f"[{start} - {end}] {para.text}") lines.append("") return "\n".join(lines) diff --git a/tests/test_formatter.py b/tests/test_formatter.py index efa41bd..cb9dc2e 100644 --- a/tests/test_formatter.py +++ b/tests/test_formatter.py @@ -2,6 +2,7 @@ from datetime import datetime from pathlib import Path from local_transcriber.formatter import ( + _group_segments, format_timestamp, format_transcript, write_transcript, @@ -50,9 +51,8 @@ def test_format_transcript_basic(): assert "**Длительность**: 02:00" in content assert "**Устройство**: CUDA (NVIDIA GeForce RTX 3060)" in content assert "---" in content - # Проверяем пробел между ] и текстом независимо от ведущих пробелов в seg.text - assert "[00:00.00 - 00:04.82] Добрый день, коллеги." in content - assert "[00:04.82 - 00:09.15] Первый вопрос." in content + # Соседние сегменты без паузы объединяются в один абзац + assert "[00:00.00 - 00:09.15] Добрый день, коллеги. Первый вопрос." in content def test_format_transcript_segment_no_leading_space(): @@ -124,6 +124,50 @@ def test_format_transcript_long(): assert "[01:01:40.00 - 01:01:50.25] Конец." in content +def test_group_segments_merges_adjacent(): + """Соседние сегменты без паузы объединяются.""" + segments = [ + Segment(start=0.0, end=3.0, text=" Первый."), + Segment(start=3.0, end=6.0, text=" Второй."), + Segment(start=6.0, end=9.0, text=" Третий."), + ] + groups = _group_segments(segments) + assert len(groups) == 1 + assert groups[0].start == 0.0 + assert groups[0].end == 9.0 + assert groups[0].text == "Первый. Второй. Третий." + + +def test_group_segments_splits_on_pause(): + """Пауза > 2с разбивает на отдельные абзацы.""" + segments = [ + Segment(start=0.0, end=3.0, text=" Первый."), + Segment(start=3.0, end=6.0, text=" Второй."), + Segment(start=9.0, end=12.0, text=" После паузы."), + ] + groups = _group_segments(segments) + assert len(groups) == 2 + assert groups[0].text == "Первый. Второй." + assert groups[1].text == "После паузы." + + +def test_group_segments_splits_on_max_duration(): + """Абзац разбивается при превышении макс. длительности.""" + segments = [ + Segment(start=0.0, end=30.0, text=" Длинный."), + Segment(start=30.0, end=55.0, text=" Ещё."), + Segment(start=55.0, end=80.0, text=" Перелив."), + ] + groups = _group_segments(segments) + assert len(groups) == 2 + assert groups[0].text == "Длинный. Ещё." + assert groups[1].text == "Перелив." + + +def test_group_segments_empty(): + assert _group_segments([]) == [] + + def test_write_transcript(tmp_path): out = tmp_path / "output.md" write_transcript("# Test content\n", out)