From 46565233d8ce7a26e31b1814387c83056ce82fba Mon Sep 17 00:00:00 2001 From: Martin Tazl Date: Wed, 29 Jul 2026 15:34:01 +0200 Subject: [PATCH] Fix Whisper JSON chunk extraction --- src/meeting_lab/chunking/chunk_transcript.py | 325 +++++++++++++++++++ tests/test_chunking.py | 50 +++ 2 files changed, 375 insertions(+) create mode 100644 tests/test_chunking.py diff --git a/src/meeting_lab/chunking/chunk_transcript.py b/src/meeting_lab/chunking/chunk_transcript.py index e69de29..a163e34 100644 --- a/src/meeting_lab/chunking/chunk_transcript.py +++ b/src/meeting_lab/chunking/chunk_transcript.py @@ -0,0 +1,325 @@ +#!/usr/bin/env python3 +""" +Split a meeting transcript into reasonably sized chunks without cutting +through transcript blocks. + +The script treats paragraphs separated by blank lines as atomic blocks. +It tries to cut close to --target-chars and will not exceed --max-chars +unless a single block is already larger than that limit. + +Example: + python chunk_transcript.py meeting.txt + python chunk_transcript.py meeting.txt --target-chars 9000 --max-chars 11000 + python chunk_transcript.py meeting.txt --overlap-blocks 1 +""" + +from __future__ import annotations + +import argparse +import json +import re +import sys +from dataclasses import asdict, dataclass +from pathlib import Path +from typing import Any + + +TIMESTAMP_RE = re.compile( + r"(?m)^\s*(?:" + r"\[(?:\d{1,2}:)?\d{1,2}:\d{2}\]" + r"|(?:\d{1,2}:)?\d{1,2}:\d{2}\s*(?:-->|-)" + r")" +) + + +@dataclass(frozen=True) +class ChunkInfo: + number: int + filename: str + chars: int + blocks: int + first_timestamp: str | None + last_timestamp: str | None + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description="Split a transcript into block-aligned text chunks." + ) + parser.add_argument( + "input_file", + type=Path, + help="Transcript text file or Whisper JSON file", + ) + parser.add_argument( + "-o", + "--output-dir", + type=Path, + help="Output directory; default: _chunks", + ) + parser.add_argument( + "--target-chars", + type=int, + default=9000, + help="Preferred chunk size in characters (default: 9000)", + ) + parser.add_argument( + "--max-chars", + type=int, + default=11000, + help="Soft maximum chunk size in characters (default: 11000)", + ) + parser.add_argument( + "--min-chars", + type=int, + default=5000, + help="Preferred minimum before a chunk may be closed (default: 5000)", + ) + parser.add_argument( + "--overlap-blocks", + type=int, + default=0, + help="Repeat this many trailing blocks in the next chunk (default: 0)", + ) + return parser.parse_args() + + +def validate_args(args: argparse.Namespace) -> None: + if args.target_chars <= 0 or args.max_chars <= 0 or args.min_chars < 0: + raise ValueError("Character limits must be positive.") + if args.min_chars > args.target_chars: + raise ValueError("--min-chars must not exceed --target-chars.") + if args.target_chars > args.max_chars: + raise ValueError("--target-chars must not exceed --max-chars.") + if args.overlap_blocks < 0: + raise ValueError("--overlap-blocks must not be negative.") + + +def normalize_text(text: str) -> str: + text = text.replace("\r\n", "\n").replace("\r", "\n") + return text.strip() + + +def blocks_from_whisper_json(data: Any) -> list[str]: + if not isinstance(data, dict): + raise ValueError("Whisper JSON input must contain a JSON object.") + + segments = data.get("segments") + if not isinstance(segments, list): + text = data.get("text") + if isinstance(text, str) and text.strip(): + return split_into_blocks(normalize_text(text)) + raise ValueError("Whisper JSON input must contain 'segments' or 'text'.") + + blocks: list[str] = [] + for segment in segments: + if not isinstance(segment, dict): + continue + text = str(segment.get("text", "")).strip() + if text: + blocks.append(text) + + if not blocks: + raise ValueError("Whisper JSON input contains no transcript text.") + + return blocks + + +def read_transcript_blocks(input_file: Path) -> list[str]: + source = normalize_text(input_file.read_text(encoding="utf-8-sig")) + if not source: + raise ValueError("The transcript is empty.") + + if input_file.suffix.lower() == ".json": + return blocks_from_whisper_json(json.loads(source)) + + return split_into_blocks(source) + + +def split_into_blocks(text: str) -> list[str]: + """ + Prefer blank-line-delimited transcript blocks. + + If the file has no blank lines but contains line-start timestamps, + split before each timestamp. Otherwise, use non-empty lines as blocks. + """ + paragraphs = [part.strip() for part in re.split(r"\n\s*\n+", text) if part.strip()] + if len(paragraphs) > 1: + return paragraphs + + timestamp_starts = list(TIMESTAMP_RE.finditer(text)) + if len(timestamp_starts) > 1: + blocks: list[str] = [] + for index, match in enumerate(timestamp_starts): + start = match.start() + end = ( + timestamp_starts[index + 1].start() + if index + 1 < len(timestamp_starts) + else len(text) + ) + block = text[start:end].strip() + if block: + blocks.append(block) + prefix = text[: timestamp_starts[0].start()].strip() + if prefix: + blocks.insert(0, prefix) + return blocks + + return [line.strip() for line in text.splitlines() if line.strip()] + + +def rendered_length(blocks: list[str]) -> int: + if not blocks: + return 0 + return sum(len(block) for block in blocks) + 2 * (len(blocks) - 1) + + +def build_chunks( + blocks: list[str], + target_chars: int, + max_chars: int, + min_chars: int, + overlap_blocks: int, +) -> list[list[str]]: + chunks: list[list[str]] = [] + current: list[str] = [] + + for block in blocks: + prospective = current + [block] + prospective_len = rendered_length(prospective) + current_len = rendered_length(current) + + should_close = ( + bool(current) + and current_len >= min_chars + and ( + prospective_len > max_chars + or (current_len >= target_chars and prospective_len > target_chars) + ) + ) + + if should_close: + chunks.append(current) + overlap = current[-overlap_blocks:] if overlap_blocks else [] + current = overlap + [block] + else: + current.append(block) + + if current: + chunks.append(current) + + return rebalance_last_chunk(chunks, min_chars, max_chars) + + +def rebalance_last_chunk( + chunks: list[list[str]], min_chars: int, max_chars: int +) -> list[list[str]]: + """ + Avoid a tiny final chunk by moving complete blocks from the previous chunk. + """ + if len(chunks) < 2 or rendered_length(chunks[-1]) >= min_chars: + return chunks + + previous = chunks[-2] + final = chunks[-1] + + while len(previous) > 1 and rendered_length(final) < min_chars: + candidate = previous[-1] + new_final = [candidate] + final + if rendered_length(new_final) > max_chars: + break + final.insert(0, previous.pop()) + + return chunks + + +def extract_timestamps(text: str) -> list[str]: + return [match.group(0).strip() for match in TIMESTAMP_RE.finditer(text)] + + +def write_chunks( + chunks: list[list[str]], output_dir: Path, input_name: str +) -> list[ChunkInfo]: + output_dir.mkdir(parents=True, exist_ok=True) + infos: list[ChunkInfo] = [] + + width = max(2, len(str(len(chunks)))) + + for number, blocks in enumerate(chunks, start=1): + content = "\n\n".join(blocks).strip() + "\n" + filename = f"chunk_{number:0{width}d}.txt" + path = output_dir / filename + path.write_text(content, encoding="utf-8") + + timestamps = extract_timestamps(content) + infos.append( + ChunkInfo( + number=number, + filename=filename, + chars=len(content), + blocks=len(blocks), + first_timestamp=timestamps[0] if timestamps else None, + last_timestamp=timestamps[-1] if timestamps else None, + ) + ) + + manifest = { + "source_file": input_name, + "chunk_count": len(infos), + "chunks": [asdict(info) for info in infos], + } + (output_dir / "manifest.json").write_text( + json.dumps(manifest, ensure_ascii=False, indent=2) + "\n", + encoding="utf-8", + ) + return infos + + +def main() -> int: + args = parse_args() + + try: + validate_args(args) + + if not args.input_file.is_file(): + raise FileNotFoundError(f"Input file not found: {args.input_file}") + + blocks = read_transcript_blocks(args.input_file) + chunks = build_chunks( + blocks=blocks, + target_chars=args.target_chars, + max_chars=args.max_chars, + min_chars=args.min_chars, + overlap_blocks=args.overlap_blocks, + ) + + output_dir = args.output_dir or args.input_file.with_name( + f"{args.input_file.stem}_chunks" + ) + infos = write_chunks(chunks, output_dir, args.input_file.name) + + except (OSError, UnicodeError, ValueError) as exc: + print(f"Error: {exc}", file=sys.stderr) + return 1 + + print(f"Source: {args.input_file}") + print(f"Blocks: {len(blocks)}") + print(f"Chunks: {len(infos)}") + print(f"Output: {output_dir}") + print() + for info in infos: + time_range = "" + if info.first_timestamp or info.last_timestamp: + time_range = ( + f" | {info.first_timestamp or '?'} to {info.last_timestamp or '?'}" + ) + print( + f"{info.filename}: {info.chars:>6} chars, " + f"{info.blocks:>3} blocks{time_range}" + ) + + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/tests/test_chunking.py b/tests/test_chunking.py new file mode 100644 index 0000000..dde08d9 --- /dev/null +++ b/tests/test_chunking.py @@ -0,0 +1,50 @@ +import json +import tempfile +import unittest +from pathlib import Path + +from src.meeting_lab.chunking.chunk_transcript import ( + build_chunks, + read_transcript_blocks, + rendered_length, +) + + +class ChunkTranscriptTests(unittest.TestCase): + def test_whisper_json_uses_segments_not_aggregate_text(self) -> None: + data = { + "text": "alpha beta gamma", + "segments": [ + {"text": "alpha"}, + {"text": "beta"}, + {"text": "gamma"}, + ], + } + with tempfile.TemporaryDirectory() as directory: + input_file = Path(directory) / "meeting.json" + input_file.write_text(json.dumps(data), encoding="utf-8") + + blocks = read_transcript_blocks(input_file) + + self.assertEqual(blocks, ["alpha", "beta", "gamma"]) + self.assertNotIn("alpha beta gamma", blocks) + + def test_chunks_do_not_share_later_blocks_without_overlap(self) -> None: + blocks = [f"block-{number:02d}-" + ("x" * 20) for number in range(10)] + + chunks = build_chunks( + blocks=blocks, + target_chars=60, + max_chars=80, + min_chars=30, + overlap_blocks=0, + ) + + flattened = [block for chunk in chunks for block in chunk] + self.assertEqual(flattened, blocks) + self.assertEqual(len(flattened), len(set(flattened))) + self.assertTrue(all(rendered_length(chunk) <= 80 for chunk in chunks)) + + +if __name__ == "__main__": + unittest.main()