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
+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.protocol.generate_direct_protocol import (
DEFAULT_MODEL,
DEFAULT_SAFE_INPUT_TOKEN_BUDGET,
DirectProtocolResult,
generate_direct_protocol,
load_compact_transcript,
@@ -54,6 +55,7 @@ class MvpMeetingConfig:
threads: str | int = "auto"
model: str = DEFAULT_MODEL
ollama_endpoint: str = DEFAULT_ENDPOINT
protocol_safe_input_token_budget: int = DEFAULT_SAFE_INPUT_TOKEN_BUDGET
diarization: str = "off"
diarization_runtime: str = "native"
diarization_container_image: str | None = None
@@ -127,6 +129,8 @@ def _validate_inputs(
and not config.diarization_container_image
):
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(
@@ -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")
_write_json(protocol_dir / "raw_response.json", result.raw_response)
_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.write_text(result.protocol_text, encoding="utf-8")
return protocol_path
@@ -354,6 +363,7 @@ def run_mvp_meeting(
preserved_context,
model=config.model,
endpoint=config.ollama_endpoint,
safe_input_token_budget=config.protocol_safe_input_token_budget,
)
stage_runtimes["protocol"] = round(time.perf_counter() - stage_started, 3)
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."""
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."
return (
f"{DIRECT_PROTOCOL_INSTRUCTION}\n\n"
f"{instruction}\n\n"
f"MEETING-KONTEXT:\n{context}\n\n"
f"VOLLSTAENDIGES TRANSKRIPT:\n{transcript.strip()}\n"
)
@@ -18,13 +18,23 @@ from src.meeting_lab.models.meeting_context import (
load_meeting_context,
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_NUM_CTX = 32768
DEFAULT_NUM_PREDICT = 8192
DEFAULT_TIMEOUT = 1800
DEFAULT_SAFE_INPUT_TOKEN_BUDGET = 16_200
ESTIMATED_UTF8_BYTES_PER_TOKEN = 4.4
class DirectProtocolError(ValueError):
@@ -38,9 +48,29 @@ class DirectProtocolResult:
model_metadata: dict[str, Any]
runtime_metadata: 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:
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():
raise DirectProtocolError(f"Transcript file does not exist: {path}")
try:
@@ -51,10 +81,73 @@ def load_compact_transcript(path: Path) -> str:
raise DirectProtocolError("Transcript JSON must contain a top-level object.")
if "text" not in data:
raise DirectProtocolError("Transcript JSON must contain top-level 'text'.")
text = data["text"]
if not isinstance(text, str) or not text.strip():
raise DirectProtocolError("Transcript top-level 'text' must be a non-empty string.")
return text
return data
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(
@@ -66,21 +159,26 @@ def generate_direct_protocol(
timeout: int = DEFAULT_TIMEOUT,
num_ctx: int = DEFAULT_NUM_CTX,
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,
generation_call: Callable[..., OllamaGeneration] = generate_once,
) -> DirectProtocolResult:
transcript = load_compact_transcript(transcript_path)
transcript = _load_transcript_document(transcript_path)
context: MeetingContext | 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
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)
generation = generation_call(
endpoint,
model,
prompt,
selected.prompt,
timeout=timeout,
num_ctx=num_ctx,
num_predict=num_predict,
@@ -101,11 +199,18 @@ def generate_direct_protocol(
"think": False,
"num_ctx": num_ctx,
"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(
protocol_text=generation.text,
exact_prompt=prompt,
exact_prompt=selected.prompt,
model_metadata=model_metadata,
runtime_metadata=runtime_metadata,
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)