Guard protocol generation against context truncation

This commit is contained in:
2026-08-24 16:48:06 +02:00
parent d94436af43
commit d77bfedb6e
12 changed files with 524 additions and 12 deletions
+6
View File
@@ -33,6 +33,12 @@ Implemented:
default and may be revisited after empirical comparison without changing the default and may be revisited after empirical comparison without changing the
orchestration API. orchestration API.
- Interim Markdown protocol generation in `src/meeting_lab/protocol/`. - Interim Markdown protocol generation in `src/meeting_lab/protocol/`.
- Direct protocol prompt input protection: diarized transcripts are rendered as
compact adjacent-speaker blocks without per-segment timestamps. Every source
segment remains represented in order. A deterministic heuristic enforces a
configurable safe input budget, falls back to complete plain transcript text
when necessary, and fails before any Ollama request if even that input is too
large. Silent head/tail truncation is prohibited.
- Non-LLM unit tests for chunking, extraction helpers, protocol rendering and - Non-LLM unit tests for chunking, extraction helpers, protocol rendering and
gold-test runner validation. gold-test runner validation.
- Meeting Context V1 scaffold and documentation for manually maintained - Meeting Context V1 scaffold and documentation for manually maintained
+25 -1
View File
@@ -35,4 +35,28 @@ torchcodec file decoder.
Anonymous `SPEAKER_XX` labels are aligned to Whisper segments by maximum Anonymous `SPEAKER_XX` labels are aligned to Whisper segments by maximum
temporal overlap with Community-1 exclusive diarization. The original Whisper temporal overlap with Community-1 exclusive diarization. The original Whisper
transcript is preserved; the derived transcript under `diarization/` is used as transcript is preserved; the derived transcript under `diarization/` is used as
the unchanged direct-protocol generator's input. the direct-protocol generator's source input.
The full diarized JSON and timestamped text remain immutable audit artifacts,
but their per-segment formatting is too verbose for a full-meeting LLM prompt:
timestamps and repeated speaker labels can more than double input size. For
protocol generation, Meeting Lab deterministically groups only adjacent
segments assigned to the same anonymous speaker and omits timestamps. A later
return by the same speaker starts a new block, and unassigned segments remain
under `SPEAKER_UNASSIGNED`. `protocol/transcript_input.txt` preserves the exact
derived representation sent to prompt construction.
Before contacting Ollama, Meeting Lab conservatively estimates prompt tokens
from UTF-8 byte count without adding a model tokenizer dependency. The default safe
budget is 16,200 estimated tokens, below the observed 16,386-token effective
boundary even when a larger `num_ctx` was requested. The estimate is calibrated
against the currently validated German BPD input and is configurable through
`MvpMeetingConfig.protocol_safe_input_token_budget` or
`--protocol-safe-input-token-budget`.
If compact diarized input exceeds the budget, the generator deterministically
uses the complete plain segment transcript and records the fallback. If that
also exceeds the budget, generation fails before model lookup or generation;
it never truncates, chunks, summarizes, retries, or makes multiple protocol
calls implicitly. Full diarization artifacts are never overwritten by this
selection.
+12
View File
@@ -20,6 +20,7 @@ if str(REPO_ROOT) not in sys.path:
from src.meeting_lab.llm.ollama import DEFAULT_ENDPOINT # noqa: E402 from src.meeting_lab.llm.ollama import DEFAULT_ENDPOINT # noqa: E402
from src.meeting_lab.protocol.generate_direct_protocol import ( # noqa: E402 from src.meeting_lab.protocol.generate_direct_protocol import ( # noqa: E402
DEFAULT_MODEL, DEFAULT_MODEL,
DEFAULT_SAFE_INPUT_TOKEN_BUDGET,
DirectProtocolResult, DirectProtocolResult,
generate_direct_protocol, generate_direct_protocol,
) )
@@ -35,6 +36,11 @@ def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
parser.add_argument("--output-root", type=Path, default=DEFAULT_OUTPUT_ROOT) parser.add_argument("--output-root", type=Path, default=DEFAULT_OUTPUT_ROOT)
parser.add_argument("--model", default=DEFAULT_MODEL) parser.add_argument("--model", default=DEFAULT_MODEL)
parser.add_argument("--ollama-endpoint", default=DEFAULT_ENDPOINT) parser.add_argument("--ollama-endpoint", default=DEFAULT_ENDPOINT)
parser.add_argument(
"--safe-input-token-budget",
type=int,
default=DEFAULT_SAFE_INPUT_TOKEN_BUDGET,
)
return parser.parse_args(argv) return parser.parse_args(argv)
@@ -64,6 +70,11 @@ def persist_result(run_dir: Path, result: DirectProtocolResult) -> Path:
(protocol_dir / "exact_prompt.txt").write_text(result.exact_prompt, encoding="utf-8") (protocol_dir / "exact_prompt.txt").write_text(result.exact_prompt, encoding="utf-8")
write_json(protocol_dir / "raw_response.json", result.raw_response) write_json(protocol_dir / "raw_response.json", result.raw_response)
write_json(protocol_dir / "runtime_metadata.json", result.runtime_metadata) write_json(protocol_dir / "runtime_metadata.json", result.runtime_metadata)
transcript_input = getattr(result, "transcript_input", None)
if transcript_input is not None:
(protocol_dir / "transcript_input.txt").write_text(
transcript_input, encoding="utf-8"
)
protocol_path = run_dir / "protocol.md" protocol_path = run_dir / "protocol.md"
protocol_path.write_text(result.protocol_text, encoding="utf-8") protocol_path.write_text(result.protocol_text, encoding="utf-8")
return protocol_path return protocol_path
@@ -116,6 +127,7 @@ def run(args: argparse.Namespace) -> tuple[int, Path, Path | None]:
preserved_context, preserved_context,
model=args.model, model=args.model,
endpoint=args.ollama_endpoint, endpoint=args.ollama_endpoint,
safe_input_token_budget=args.safe_input_token_budget,
) )
protocol_path = persist_result(run_dir, result) protocol_path = persist_result(run_dir, result)
metadata["status"] = "completed" metadata["status"] = "completed"
+8
View File
@@ -19,6 +19,7 @@ from src.meeting_lab.orchestration.mvp import ( # noqa: E402
DEFAULT_DIARIZATION_MODEL, DEFAULT_DIARIZATION_MODEL,
DEFAULT_MODEL, DEFAULT_MODEL,
DEFAULT_OUTPUT_ROOT, DEFAULT_OUTPUT_ROOT,
DEFAULT_SAFE_INPUT_TOKEN_BUDGET,
MvpMeetingConfig, MvpMeetingConfig,
create_unique_run_dir, create_unique_run_dir,
run_mvp_meeting, run_mvp_meeting,
@@ -53,6 +54,12 @@ def parse_args(argv: list[str] | None = None) -> argparse.Namespace:
) )
parser.add_argument("--model", default=DEFAULT_MODEL) parser.add_argument("--model", default=DEFAULT_MODEL)
parser.add_argument("--ollama-endpoint", default=DEFAULT_ENDPOINT) parser.add_argument("--ollama-endpoint", default=DEFAULT_ENDPOINT)
parser.add_argument(
"--protocol-safe-input-token-budget",
type=int,
default=DEFAULT_SAFE_INPUT_TOKEN_BUDGET,
help="Conservative estimated prompt-token limit before any Ollama request.",
)
parser.add_argument( parser.add_argument(
"--diarization", "--diarization",
choices=("auto", "gpu", "cpu", "off"), choices=("auto", "gpu", "cpu", "off"),
@@ -91,6 +98,7 @@ def config_from_args(args: argparse.Namespace) -> MvpMeetingConfig:
threads=args.threads, threads=args.threads,
model=args.model, model=args.model,
ollama_endpoint=args.ollama_endpoint, ollama_endpoint=args.ollama_endpoint,
protocol_safe_input_token_budget=args.protocol_safe_input_token_budget,
diarization=args.diarization, diarization=args.diarization,
diarization_runtime=args.diarization_runtime, diarization_runtime=args.diarization_runtime,
diarization_container_image=args.diarization_container_image, diarization_container_image=args.diarization_container_image,
+10
View File
@@ -30,6 +30,7 @@ from src.meeting_lab.models.meeting_context import (
from src.meeting_lab.progress import ProgressEvent, ProgressSink, ProgressStatus from src.meeting_lab.progress import ProgressEvent, ProgressSink, ProgressStatus
from src.meeting_lab.protocol.generate_direct_protocol import ( from src.meeting_lab.protocol.generate_direct_protocol import (
DEFAULT_MODEL, DEFAULT_MODEL,
DEFAULT_SAFE_INPUT_TOKEN_BUDGET,
DirectProtocolResult, DirectProtocolResult,
generate_direct_protocol, generate_direct_protocol,
load_compact_transcript, load_compact_transcript,
@@ -54,6 +55,7 @@ class MvpMeetingConfig:
threads: str | int = "auto" threads: str | int = "auto"
model: str = DEFAULT_MODEL model: str = DEFAULT_MODEL
ollama_endpoint: str = DEFAULT_ENDPOINT ollama_endpoint: str = DEFAULT_ENDPOINT
protocol_safe_input_token_budget: int = DEFAULT_SAFE_INPUT_TOKEN_BUDGET
diarization: str = "off" diarization: str = "off"
diarization_runtime: str = "native" diarization_runtime: str = "native"
diarization_container_image: str | None = None diarization_container_image: str | None = None
@@ -127,6 +129,8 @@ def _validate_inputs(
and not config.diarization_container_image and not config.diarization_container_image
): ):
raise ValueError("A diarization container image is required.") raise ValueError("A diarization container image is required.")
if config.protocol_safe_input_token_budget <= 0:
raise ValueError("Protocol safe input token budget must be positive.")
def _emit( def _emit(
@@ -154,6 +158,11 @@ def _persist_protocol(run_dir: Path, result: DirectProtocolResult) -> Path:
(protocol_dir / "exact_prompt.txt").write_text(result.exact_prompt, encoding="utf-8") (protocol_dir / "exact_prompt.txt").write_text(result.exact_prompt, encoding="utf-8")
_write_json(protocol_dir / "raw_response.json", result.raw_response) _write_json(protocol_dir / "raw_response.json", result.raw_response)
_write_json(protocol_dir / "runtime_metadata.json", result.runtime_metadata) _write_json(protocol_dir / "runtime_metadata.json", result.runtime_metadata)
transcript_input = getattr(result, "transcript_input", None)
if transcript_input is not None:
(protocol_dir / "transcript_input.txt").write_text(
transcript_input, encoding="utf-8"
)
protocol_path = run_dir / "protocol.md" protocol_path = run_dir / "protocol.md"
protocol_path.write_text(result.protocol_text, encoding="utf-8") protocol_path.write_text(result.protocol_text, encoding="utf-8")
return protocol_path return protocol_path
@@ -354,6 +363,7 @@ def run_mvp_meeting(
preserved_context, preserved_context,
model=config.model, model=config.model,
endpoint=config.ollama_endpoint, endpoint=config.ollama_endpoint,
safe_input_token_budget=config.protocol_safe_input_token_budget,
) )
stage_runtimes["protocol"] = round(time.perf_counter() - stage_started, 3) stage_runtimes["protocol"] = round(time.perf_counter() - stage_started, 3)
protocol_path = _persist_protocol(run_dir, result) protocol_path = _persist_protocol(run_dir, result)
@@ -9,11 +9,22 @@ Das Protokoll muss themenorientiert sein, nicht chronologisch und nicht nach tec
Erzeuge keine reine Wiedergabe des Transkripts und verlängere das Protokoll nicht unnötig durch Wiederholungen. Synthetisiere zusammengehörige Aussagen, entferne Füllwörter und Gesprächsrauschen und erfinde keine Fakten, Entscheidungen, Zustimmungen, Verantwortlichen oder Fristen. Gib kein JSON, keine internen Labels und keine Analyse oder Denkprotokolle aus. Das Ergebnis soll als Markdown-Protokoll nach geringfügiger menschlicher Redaktion intern versendbar sein. Eine kompakte themenübergreifende Maßnahmenliste am Ende ist optional, wenn sie nützlich und vollständig belegt ist.""" Erzeuge keine reine Wiedergabe des Transkripts und verlängere das Protokoll nicht unnötig durch Wiederholungen. Synthetisiere zusammengehörige Aussagen, entferne Füllwörter und Gesprächsrauschen und erfinde keine Fakten, Entscheidungen, Zustimmungen, Verantwortlichen oder Fristen. Gib kein JSON, keine internen Labels und keine Analyse oder Denkprotokolle aus. Das Ergebnis soll als Markdown-Protokoll nach geringfügiger menschlicher Redaktion intern versendbar sein. Eine kompakte themenübergreifende Maßnahmenliste am Ende ist optional, wenn sie nützlich und vollständig belegt ist."""
COMPACT_DIARIZED_PROTOCOL_INSTRUCTION = """Erstelle aus dem vollständigen Transkript und Meeting-Kontext ein vollständiges, professionelles internes Besprechungsprotokoll auf Deutsch. Das Transkript ist in aufeinanderfolgende anonyme Sprecherblöcke gegliedert.
def build_direct_protocol_prompt(transcript: str, meeting_context: str | None = None) -> str: Beginne mit # Meeting Protocol. Gliedere themenorientiert mit ## <Thema> und synthetisiere je Thema den relevanten Diskussionsverlauf, Kontext, unterschiedliche Positionen, Entscheidungsgrundlagen, Einschränkungen und ungelöste Meinungsverschiedenheiten so, dass Dritte ihn nachvollziehen können. Nenne Entscheidungen nur bei Beleg. Nenne Maßnahmen, Verantwortliche und Fristen nur bei expliziter Zuweisung, Annahme oder Bestätigung; Vorschläge sind keine Verpflichtungen.
Entferne nur Wiederholungen, Füllwörter und Gesprächsrauschen. Erfinde keine Fakten oder Identitäten. Gib kein JSON, keine Sprecherlabels und kein Denkprotokoll aus. Eine belegte themenübergreifende Maßnahmenliste am Ende ist optional."""
def build_direct_protocol_prompt(
transcript: str,
meeting_context: str | None = None,
*,
instruction: str = DIRECT_PROTOCOL_INSTRUCTION,
) -> str:
context = meeting_context.strip() if meeting_context else "Kein Meeting-Kontext bereitgestellt." context = meeting_context.strip() if meeting_context else "Kein Meeting-Kontext bereitgestellt."
return ( return (
f"{DIRECT_PROTOCOL_INSTRUCTION}\n\n" f"{instruction}\n\n"
f"MEETING-KONTEXT:\n{context}\n\n" f"MEETING-KONTEXT:\n{context}\n\n"
f"VOLLSTAENDIGES TRANSKRIPT:\n{transcript.strip()}\n" f"VOLLSTAENDIGES TRANSKRIPT:\n{transcript.strip()}\n"
) )
@@ -18,13 +18,23 @@ from src.meeting_lab.models.meeting_context import (
load_meeting_context, load_meeting_context,
render_meeting_context_for_prompt, render_meeting_context_for_prompt,
) )
from src.meeting_lab.protocol.direct_protocol_prompt import build_direct_protocol_prompt from src.meeting_lab.protocol.direct_protocol_prompt import (
COMPACT_DIARIZED_PROTOCOL_INSTRUCTION,
build_direct_protocol_prompt,
)
from src.meeting_lab.protocol.transcript_input import (
TranscriptInputError,
compact_diarized_transcript,
plain_segment_transcript,
)
DEFAULT_MODEL = "qwen3.6:35B-A3B" DEFAULT_MODEL = "qwen3.6:35B-A3B"
DEFAULT_NUM_CTX = 32768 DEFAULT_NUM_CTX = 32768
DEFAULT_NUM_PREDICT = 8192 DEFAULT_NUM_PREDICT = 8192
DEFAULT_TIMEOUT = 1800 DEFAULT_TIMEOUT = 1800
DEFAULT_SAFE_INPUT_TOKEN_BUDGET = 16_200
ESTIMATED_UTF8_BYTES_PER_TOKEN = 4.4
class DirectProtocolError(ValueError): class DirectProtocolError(ValueError):
@@ -38,9 +48,29 @@ class DirectProtocolResult:
model_metadata: dict[str, Any] model_metadata: dict[str, Any]
runtime_metadata: dict[str, Any] runtime_metadata: dict[str, Any]
raw_response: dict[str, Any] raw_response: dict[str, Any]
transcript_input: str | None = None
@dataclass(frozen=True)
class SelectedTranscriptInput:
text: str
prompt: str
representation: str
estimated_input_tokens: int
safe_input_token_budget: int
fallback_used: bool
diarization_enabled: bool
def load_compact_transcript(path: Path) -> str: def load_compact_transcript(path: Path) -> str:
data = _load_transcript_document(path)
text = data.get("text")
if not isinstance(text, str) or not text.strip():
raise DirectProtocolError("Transcript top-level 'text' must be a non-empty string.")
return text
def _load_transcript_document(path: Path) -> dict[str, Any]:
if not path.is_file(): if not path.is_file():
raise DirectProtocolError(f"Transcript file does not exist: {path}") raise DirectProtocolError(f"Transcript file does not exist: {path}")
try: try:
@@ -51,10 +81,73 @@ def load_compact_transcript(path: Path) -> str:
raise DirectProtocolError("Transcript JSON must contain a top-level object.") raise DirectProtocolError("Transcript JSON must contain a top-level object.")
if "text" not in data: if "text" not in data:
raise DirectProtocolError("Transcript JSON must contain top-level 'text'.") raise DirectProtocolError("Transcript JSON must contain top-level 'text'.")
text = data["text"] return data
if not isinstance(text, str) or not text.strip():
raise DirectProtocolError("Transcript top-level 'text' must be a non-empty string.")
return text def estimate_input_tokens(prompt: str) -> int:
"""Estimate tokens without adding a model-specific tokenizer dependency."""
byte_count = len(prompt.encode("utf-8"))
return max(1, int(byte_count / ESTIMATED_UTF8_BYTES_PER_TOKEN + 0.999999))
def select_transcript_input(
transcript: dict[str, Any],
rendered_context: str | None,
*,
safe_input_token_budget: int = DEFAULT_SAFE_INPUT_TOKEN_BUDGET,
) -> SelectedTranscriptInput:
"""Select complete prompt input without allowing silent tail truncation."""
if safe_input_token_budget <= 0:
raise DirectProtocolError("Safe protocol input token budget must be positive.")
diarization_enabled = transcript.get("speaker_labels_anonymous") is True
if diarization_enabled:
try:
compact = compact_diarized_transcript(transcript.get("segments"))
plain_text = plain_segment_transcript(transcript.get("segments"))
except TranscriptInputError as exc:
raise DirectProtocolError(str(exc)) from exc
compact_prompt = build_direct_protocol_prompt(
compact.text,
rendered_context,
instruction=COMPACT_DIARIZED_PROTOCOL_INSTRUCTION,
)
compact_estimate = estimate_input_tokens(compact_prompt)
if compact_estimate <= safe_input_token_budget:
return SelectedTranscriptInput(
text=compact.text,
prompt=compact_prompt,
representation="diarized_compact",
estimated_input_tokens=compact_estimate,
safe_input_token_budget=safe_input_token_budget,
fallback_used=False,
diarization_enabled=True,
)
representation = "plain_transcript_fallback"
fallback_used = True
else:
plain_text = transcript.get("text")
if not isinstance(plain_text, str) or not plain_text.strip():
raise DirectProtocolError("Transcript top-level 'text' must be a non-empty string.")
representation = "plain_transcript"
fallback_used = False
plain_prompt = build_direct_protocol_prompt(plain_text, rendered_context)
plain_estimate = estimate_input_tokens(plain_prompt)
if plain_estimate > safe_input_token_budget:
raise DirectProtocolError(
"Protocol prompt/input is too large for the configured safe input budget "
f"({plain_estimate} estimated tokens > {safe_input_token_budget}). "
"No LLM request was made; silent truncation is not allowed."
)
return SelectedTranscriptInput(
text=plain_text,
prompt=plain_prompt,
representation=representation,
estimated_input_tokens=plain_estimate,
safe_input_token_budget=safe_input_token_budget,
fallback_used=fallback_used,
diarization_enabled=diarization_enabled,
)
def generate_direct_protocol( def generate_direct_protocol(
@@ -66,21 +159,26 @@ def generate_direct_protocol(
timeout: int = DEFAULT_TIMEOUT, timeout: int = DEFAULT_TIMEOUT,
num_ctx: int = DEFAULT_NUM_CTX, num_ctx: int = DEFAULT_NUM_CTX,
num_predict: int = DEFAULT_NUM_PREDICT, num_predict: int = DEFAULT_NUM_PREDICT,
safe_input_token_budget: int = DEFAULT_SAFE_INPUT_TOKEN_BUDGET,
model_check: Callable[[str, str, int], dict[str, Any]] = require_model, model_check: Callable[[str, str, int], dict[str, Any]] = require_model,
generation_call: Callable[..., OllamaGeneration] = generate_once, generation_call: Callable[..., OllamaGeneration] = generate_once,
) -> DirectProtocolResult: ) -> DirectProtocolResult:
transcript = load_compact_transcript(transcript_path) transcript = _load_transcript_document(transcript_path)
context: MeetingContext | None = ( context: MeetingContext | None = (
load_meeting_context(context_path) if context_path is not None else None load_meeting_context(context_path) if context_path is not None else None
) )
rendered_context = render_meeting_context_for_prompt(context) if context else None rendered_context = render_meeting_context_for_prompt(context) if context else None
prompt = build_direct_protocol_prompt(transcript, rendered_context) selected = select_transcript_input(
transcript,
rendered_context,
safe_input_token_budget=safe_input_token_budget,
)
model_metadata = model_check(endpoint, model, 10) model_metadata = model_check(endpoint, model, 10)
generation = generation_call( generation = generation_call(
endpoint, endpoint,
model, model,
prompt, selected.prompt,
timeout=timeout, timeout=timeout,
num_ctx=num_ctx, num_ctx=num_ctx,
num_predict=num_predict, num_predict=num_predict,
@@ -101,11 +199,18 @@ def generate_direct_protocol(
"think": False, "think": False,
"num_ctx": num_ctx, "num_ctx": num_ctx,
"num_predict": num_predict, "num_predict": num_predict,
"selected_transcript_representation": selected.representation,
"estimated_input_tokens": selected.estimated_input_tokens,
"safe_input_token_budget": selected.safe_input_token_budget,
"input_token_estimation_method": "utf8_bytes_divided_by_4.4",
"fallback_used": selected.fallback_used,
"diarization_enabled": selected.diarization_enabled,
} }
return DirectProtocolResult( return DirectProtocolResult(
protocol_text=generation.text, protocol_text=generation.text,
exact_prompt=prompt, exact_prompt=selected.prompt,
model_metadata=model_metadata, model_metadata=model_metadata,
runtime_metadata=runtime_metadata, runtime_metadata=runtime_metadata,
raw_response=data, raw_response=data,
transcript_input=selected.text,
) )
@@ -0,0 +1,97 @@
"""Deterministic transcript representations for one-call protocol prompts."""
from __future__ import annotations
from dataclasses import dataclass
from typing import Any
class TranscriptInputError(ValueError):
"""Raised when a transcript cannot be represented without content loss."""
@dataclass(frozen=True)
class SpeakerBlock:
"""One contiguous run of transcript segments assigned to one speaker."""
speaker_id: str
segment_texts: tuple[str, ...]
@dataclass(frozen=True)
class CompactDiarizedTranscript:
"""Compact prompt text plus structural evidence of segment preservation."""
text: str
blocks: tuple[SpeakerBlock, ...]
source_segment_count: int
@property
def represented_segment_count(self) -> int:
return sum(len(block.segment_texts) for block in self.blocks)
@property
def segment_texts(self) -> tuple[str, ...]:
return tuple(text for block in self.blocks for text in block.segment_texts)
def normalize_segment_text(value: Any, index: int) -> str:
"""Normalize formatting whitespace while retaining all semantic text."""
if not isinstance(value, str):
raise TranscriptInputError(f"Transcript segment {index} text must be a string.")
return " ".join(value.split())
def compact_diarized_transcript(segments: Any) -> CompactDiarizedTranscript:
"""Group only adjacent same-speaker segments and omit repeated timestamps."""
if not isinstance(segments, list) or not segments:
raise TranscriptInputError(
"Diarized transcript must contain a non-empty 'segments' list."
)
mutable_blocks: list[tuple[str, list[str]]] = []
source_texts: list[str] = []
for index, segment in enumerate(segments):
if not isinstance(segment, dict):
raise TranscriptInputError(f"Transcript segment {index} must be an object.")
speaker = segment.get("speaker_id") or "SPEAKER_UNASSIGNED"
if not isinstance(speaker, str) or not speaker.startswith("SPEAKER_"):
raise TranscriptInputError(
f"Transcript segment {index} must use an anonymous SPEAKER_ label."
)
text = normalize_segment_text(segment.get("text"), index)
source_texts.append(text)
if mutable_blocks and mutable_blocks[-1][0] == speaker:
mutable_blocks[-1][1].append(text)
else:
mutable_blocks.append((speaker, [text]))
blocks = tuple(
SpeakerBlock(speaker_id=speaker, segment_texts=tuple(texts))
for speaker, texts in mutable_blocks
)
rendered = "\n".join(
f"{block.speaker_id}: {' '.join(block.segment_texts)}" for block in blocks
)
result = CompactDiarizedTranscript(
text=rendered + "\n",
blocks=blocks,
source_segment_count=len(segments),
)
if result.represented_segment_count != len(segments):
raise TranscriptInputError("Compact diarized transcript lost source segments.")
if result.segment_texts != tuple(source_texts):
raise TranscriptInputError("Compact diarized transcript changed segment order or text.")
return result
def plain_segment_transcript(segments: Any) -> str:
"""Reconstruct plain transcript text from every segment in source order."""
if not isinstance(segments, list) or not segments:
raise TranscriptInputError("Transcript must contain a non-empty 'segments' list.")
texts = []
for index, segment in enumerate(segments):
if not isinstance(segment, dict):
raise TranscriptInputError(f"Transcript segment {index} must be an object.")
texts.append(normalize_segment_text(segment.get("text"), index))
return " ".join(texts)
+5
View File
@@ -232,6 +232,7 @@ class DirectProtocolCliTests(unittest.TestCase):
return_value=type("Result", (), { return_value=type("Result", (), {
"protocol_text": protocol_text, "protocol_text": protocol_text,
"exact_prompt": "exact prompt\n", "exact_prompt": "exact prompt\n",
"transcript_input": "selected transcript\n",
"raw_response": {"response": protocol_text}, "raw_response": {"response": protocol_text},
"runtime_metadata": {"request_count": 1}, "runtime_metadata": {"request_count": 1},
})(), })(),
@@ -245,6 +246,10 @@ class DirectProtocolCliTests(unittest.TestCase):
(run_dir / "protocol/exact_prompt.txt").read_text(encoding="utf-8"), (run_dir / "protocol/exact_prompt.txt").read_text(encoding="utf-8"),
"exact prompt\n", "exact prompt\n",
) )
self.assertEqual(
(run_dir / "protocol/transcript_input.txt").read_text(encoding="utf-8"),
"selected transcript\n",
)
self.assertEqual( self.assertEqual(
json.loads((run_dir / "protocol/raw_response.json").read_text())["response"], json.loads((run_dir / "protocol/raw_response.json").read_text())["response"],
protocol_text, protocol_text,
+1
View File
@@ -180,6 +180,7 @@ class MvpApiTests(unittest.TestCase):
self.assertEqual(delegated.whisper_executable, "whisper-cli") self.assertEqual(delegated.whisper_executable, "whisper-cli")
self.assertEqual(delegated.ffmpeg_executable, "ffmpeg") self.assertEqual(delegated.ffmpeg_executable, "ffmpeg")
self.assertTrue(delegated.audio_normalization) self.assertTrue(delegated.audio_normalization)
self.assertEqual(delegated.protocol_safe_input_token_budget, 16_200)
self.assertEqual(api.call_args.kwargs["meeting_context"], context_data()) self.assertEqual(api.call_args.kwargs["meeting_context"], context_data())
def test_cli_explicit_audio_normalization_values_are_propagated(self): def test_cli_explicit_audio_normalization_values_are_propagated(self):
+2
View File
@@ -39,6 +39,7 @@ def protocol_result(model: str = "chosen:model") -> DirectProtocolResult:
model_metadata={"model": model}, model_metadata={"model": model},
runtime_metadata={"model": model, "request_count": 1, "client_wall_time_seconds": 0.5}, runtime_metadata={"model": model, "request_count": 1, "client_wall_time_seconds": 0.5},
raw_response={"response": text, "done": True}, raw_response={"response": text, "done": True},
transcript_input="selected transcript\n",
) )
@@ -186,6 +187,7 @@ class MvpOrchestratorTests(unittest.TestCase):
"transcript/runtime_metadata.json", "transcript/runtime_metadata.json",
"context/meeting_context.yaml", "context/meeting_context.yaml",
"protocol/exact_prompt.txt", "protocol/exact_prompt.txt",
"protocol/transcript_input.txt",
"protocol/raw_response.json", "protocol/raw_response.json",
"protocol/runtime_metadata.json", "protocol/runtime_metadata.json",
"protocol.md", "protocol.md",
+231
View File
@@ -0,0 +1,231 @@
import json
import tempfile
import unittest
from pathlib import Path
from unittest.mock import Mock
from src.meeting_lab.diarization.alignment import diarized_transcript_text
from src.meeting_lab.llm.ollama import OllamaGeneration
from src.meeting_lab.protocol.direct_protocol_prompt import (
COMPACT_DIARIZED_PROTOCOL_INSTRUCTION,
build_direct_protocol_prompt,
)
from src.meeting_lab.protocol.generate_direct_protocol import (
DirectProtocolError,
estimate_input_tokens,
generate_direct_protocol,
)
from src.meeting_lab.protocol.transcript_input import compact_diarized_transcript
def segments() -> list[dict[str, object]]:
return [
{"id": 0, "start": 0.0, "end": 1.0, "text": "First.", "speaker_id": "SPEAKER_01"},
{"id": 1, "start": 1.0, "end": 2.0, "text": "Second.", "speaker_id": "SPEAKER_01"},
{"id": 2, "start": 2.0, "end": 3.0, "text": "Third.", "speaker_id": "SPEAKER_04"},
{"id": 3, "start": 3.0, "end": 4.0, "text": "Unassigned.", "speaker_id": None},
{"id": 4, "start": 4.0, "end": 5.0, "text": "Last.", "speaker_id": "SPEAKER_01"},
]
def diarized_document(repetitions: int = 1) -> dict[str, object]:
source = segments() * repetitions
return {
"text": diarized_transcript_text(source),
"segments": source,
"speaker_labels_anonymous": True,
"alignment_source": "exclusive_diarization",
}
def completed_generation() -> OllamaGeneration:
text = "# Meeting Protocol\n\nComplete."
return OllamaGeneration(
raw_response={
"response": text,
"done": True,
"done_reason": "stop",
"prompt_eval_count": 100,
"eval_count": 10,
},
text=text,
client_wall_time_seconds=0.1,
)
class CompactDiarizedTranscriptTests(unittest.TestCase):
def test_adjacent_segments_group_and_transitions_remain_separate(self) -> None:
compact = compact_diarized_transcript(segments())
self.assertEqual(
[block.speaker_id for block in compact.blocks],
["SPEAKER_01", "SPEAKER_04", "SPEAKER_UNASSIGNED", "SPEAKER_01"],
)
self.assertEqual(compact.blocks[0].segment_texts, ("First.", "Second."))
self.assertEqual(compact.blocks[-1].segment_texts, ("Last.",))
self.assertEqual(compact.text.count("SPEAKER_01:"), 2)
def test_every_segment_text_and_order_are_preserved(self) -> None:
source = segments()
compact = compact_diarized_transcript(source)
self.assertEqual(compact.source_segment_count, len(source))
self.assertEqual(compact.represented_segment_count, len(source))
self.assertEqual(
compact.segment_texts,
tuple(str(segment["text"]) for segment in source),
)
self.assertEqual(compact.segment_texts[0], "First.")
self.assertEqual(compact.segment_texts[-1], "Last.")
self.assertIn("SPEAKER_UNASSIGNED: Unassigned.", compact.text)
def test_compact_form_is_materially_smaller_than_per_segment_format(self) -> None:
source = [
{
"start": index,
"end": index + 1,
"text": "Repeated transcript content.",
"speaker_id": "SPEAKER_01",
}
for index in range(100)
]
compact = compact_diarized_transcript(source).text
verbose = diarized_transcript_text(source)
self.assertLess(len(compact), len(verbose) * 0.6)
class ProtocolInputBudgetTests(unittest.TestCase):
def _write(self, root: Path, document: dict[str, object]) -> Path:
path = root / "transcript.json"
path.write_text(json.dumps(document), encoding="utf-8")
return path
def test_token_estimate_uses_utf8_bytes_for_non_ascii_safety(self) -> None:
self.assertEqual(estimate_input_tokens("ä" * 44), 20)
def test_compact_diarized_representation_selected_within_budget(self) -> None:
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
document = diarized_document()
transcript = self._write(root, document)
compact = compact_diarized_transcript(document["segments"]).text
budget = estimate_input_tokens(
build_direct_protocol_prompt(
compact,
instruction=COMPACT_DIARIZED_PROTOCOL_INSTRUCTION,
)
)
call = Mock(return_value=completed_generation())
result = generate_direct_protocol(
transcript,
safe_input_token_budget=budget,
model_check=Mock(return_value={}),
generation_call=call,
)
self.assertEqual(
result.runtime_metadata["selected_transcript_representation"],
"diarized_compact",
)
self.assertFalse(result.runtime_metadata["fallback_used"])
self.assertTrue(result.runtime_metadata["diarization_enabled"])
self.assertEqual(result.runtime_metadata["safe_input_token_budget"], budget)
self.assertEqual(result.transcript_input, compact)
self.assertEqual(call.call_count, 1)
def test_plain_fallback_selected_when_diarized_compact_exceeds_budget(self) -> None:
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
alternating_segments = [
{
"id": index,
"start": float(index),
"end": float(index + 1),
"text": "Word.",
"speaker_id": f"SPEAKER_{index % 2:02d}",
}
for index in range(200)
]
document = {
"text": diarized_transcript_text(alternating_segments),
"segments": alternating_segments,
"speaker_labels_anonymous": True,
"alignment_source": "exclusive_diarization",
}
transcript = self._write(root, document)
compact = compact_diarized_transcript(document["segments"]).text
plain = " ".join(str(segment["text"]) for segment in document["segments"])
compact_estimate = estimate_input_tokens(
build_direct_protocol_prompt(
compact,
instruction=COMPACT_DIARIZED_PROTOCOL_INSTRUCTION,
)
)
plain_estimate = estimate_input_tokens(build_direct_protocol_prompt(plain))
self.assertLess(plain_estimate, compact_estimate)
result = generate_direct_protocol(
transcript,
safe_input_token_budget=plain_estimate,
model_check=Mock(return_value={}),
generation_call=Mock(return_value=completed_generation()),
)
self.assertEqual(
result.runtime_metadata["selected_transcript_representation"],
"plain_transcript_fallback",
)
self.assertTrue(result.runtime_metadata["fallback_used"])
self.assertEqual(result.transcript_input, plain)
def test_oversized_plain_transcript_fails_before_any_network_call(self) -> None:
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
transcript = self._write(
root,
{"text": "large input " * 1000, "segments": []},
)
model_check = Mock()
generation_call = Mock()
with self.assertRaisesRegex(
DirectProtocolError,
"No LLM request was made; silent truncation is not allowed",
):
generate_direct_protocol(
transcript,
safe_input_token_budget=1,
model_check=model_check,
generation_call=generation_call,
)
model_check.assert_not_called()
generation_call.assert_not_called()
def test_existing_plain_path_and_metadata_remain_direct(self) -> None:
with tempfile.TemporaryDirectory() as directory:
root = Path(directory)
transcript = self._write(
root,
{"text": "Plain original transcript.", "segments": []},
)
result = generate_direct_protocol(
transcript,
model_check=Mock(return_value={}),
generation_call=Mock(return_value=completed_generation()),
)
self.assertEqual(result.transcript_input, "Plain original transcript.")
self.assertEqual(
result.runtime_metadata["selected_transcript_representation"],
"plain_transcript",
)
self.assertFalse(result.runtime_metadata["fallback_used"])
self.assertFalse(result.runtime_metadata["diarization_enabled"])
if __name__ == "__main__":
unittest.main()