Files
meeting-lab/src/meeting_lab/consolidation/consolidate_facts.py
T

948 lines
33 KiB
Python

#!/usr/bin/env python3
"""LLM-backed V0 semantic consolidation for fact items only."""
from __future__ import annotations
import argparse
import copy
import json
import re
import sys
import time
from collections import Counter
from datetime import datetime
from pathlib import Path
from queue import Empty, Queue
from threading import Thread
from typing import Any
import requests
try:
from meeting_lab.llm.prompts import load_prompt
except ModuleNotFoundError: # pragma: no cover - used by repository-root tests.
from src.meeting_lab.llm.prompts import load_prompt
DEFAULT_MODEL = "qwen3.5:9b"
DEFAULT_ENDPOINT = "http://127.0.0.1:11434/api/generate"
DEFAULT_NUM_CTX = 32768
DEFAULT_MIN_NUM_PREDICT = 4096
DEFAULT_NUM_PREDICT = DEFAULT_MIN_NUM_PREDICT
DEFAULT_PROGRESS_INTERVAL = 30
OUTPUT_CONTEXT_RESERVE_TOKENS = 1024
OUTPUT_TOKEN_ESTIMATE_CHARS = 4
OUTPUT_GROUP_OVERHEAD_CHARS = 320
OUTPUT_SAFETY_MARGIN = 1.35
PROMPT_NAME = "consolidate_facts.md"
REPETITION_LOOP_MIN_CONSECUTIVE_GROUPS = 3
REPETITION_RETRY_INSTRUCTION = """
RETRY SAFETY INSTRUCTION:
Emit each semantic group exactly once. Before returning the JSON, check the
groups array for identical groups using canonical_text, source_item_ids and
merge_reason. If an identical group is already present, do not emit it again.
Finish the complete JSON object without repeating any group.
""".rstrip()
class ConsolidationValidationError(ValueError):
"""Raised when model consolidation output violates strict invariants."""
def parse_args() -> argparse.Namespace:
parser = argparse.ArgumentParser(
description="Consolidate semantically equivalent fact items only."
)
parser.add_argument(
"canonicalized_input",
type=Path,
help="Canonicalizer V1 JSON file.",
)
parser.add_argument(
"-o",
"--output-dir",
type=Path,
required=True,
help="Directory for consolidated_extractions.json, report.md and raw response.",
)
parser.add_argument(
"--model",
default=DEFAULT_MODEL,
help=f"Ollama model name (default: {DEFAULT_MODEL}).",
)
parser.add_argument(
"--endpoint",
default=DEFAULT_ENDPOINT,
help=f"Ollama generate endpoint (default: {DEFAULT_ENDPOINT}).",
)
parser.add_argument(
"--timeout",
type=int,
default=1800,
help="HTTP timeout in seconds (default: 1800).",
)
parser.add_argument(
"--num-ctx",
type=int,
default=DEFAULT_NUM_CTX,
help=f"Context window tokens (default: {DEFAULT_NUM_CTX}).",
)
parser.add_argument(
"--num-predict",
type=int,
default=None,
help=(
"Maximum generated tokens. By default this is estimated from the "
"fact payload size and bounded by the context window. Explicit "
"values preserve the previous fixed-budget behavior."
),
)
thinking = parser.add_mutually_exclusive_group()
thinking.add_argument(
"--think",
dest="think",
action="store_true",
help="Enable Ollama thinking output when the selected model supports it.",
)
thinking.add_argument(
"--no-think",
dest="think",
action="store_false",
help="Disable Ollama thinking output for structured JSON consolidation.",
)
parser.set_defaults(think=False)
parser.add_argument(
"--progress-interval",
type=int,
default=DEFAULT_PROGRESS_INTERVAL,
help=(
"Seconds between waiting-status messages while the non-streaming "
f"Ollama request is in flight (default: {DEFAULT_PROGRESS_INTERVAL})."
),
)
return parser.parse_args()
def load_json_object(path: Path) -> dict[str, Any]:
try:
data = json.loads(path.read_text(encoding="utf-8-sig"))
except json.JSONDecodeError as exc:
raise ValueError(f"Invalid JSON in {path}: {exc}") from exc
if not isinstance(data, dict):
raise ValueError(f"JSON file must contain an object: {path}")
return data
def fact_items(canonicalized: dict[str, Any]) -> list[dict[str, Any]]:
items = canonicalized.get("items")
if not isinstance(items, list):
raise ValueError("Canonicalized input must contain an items list.")
return [item for item in items if item.get("category") == "fact"]
def non_fact_items(canonicalized: dict[str, Any]) -> list[dict[str, Any]]:
items = canonicalized.get("items")
if not isinstance(items, list):
raise ValueError("Canonicalized input must contain an items list.")
return [item for item in items if item.get("category") != "fact"]
def model_fact_payload(facts: list[dict[str, Any]]) -> list[dict[str, Any]]:
payload: list[dict[str, Any]] = []
for item in facts:
payload.append(
{
"item_id": item.get("item_id"),
"text": item.get("text"),
"evidence": item.get("evidence"),
"speaker": item.get("speaker"),
"status": item.get("status"),
"source_file": item.get("source_file"),
"source_index": item.get("source_index"),
}
)
return payload
def build_consolidation_prompt(facts: list[dict[str, Any]]) -> str:
task_prompt = load_prompt(PROMPT_NAME)
payload = json.dumps(
{"fact_items": model_fact_payload(facts)},
ensure_ascii=False,
indent=2,
)
return f"{task_prompt}\n\nFACT ITEMS:\n{payload}\n"
def estimate_response_tokens(facts: list[dict[str, Any]]) -> int:
"""
Estimate the token budget needed for the model's grouping JSON.
Semantic Consolidator V0 asks the model to return one group per source fact
unless it finds a conservative duplicate. The response therefore scales with
the number and text size of fact items. The estimate intentionally includes
per-group JSON overhead and a safety margin; strict validation still decides
whether the actual response is usable.
"""
text_chars = 0
for item in facts:
text_chars += len(str(item.get("text", "")))
text_chars += len(str(item.get("evidence", "")))
estimated_chars = int(
(text_chars + len(facts) * OUTPUT_GROUP_OVERHEAD_CHARS)
* OUTPUT_SAFETY_MARGIN
)
return max(
DEFAULT_MIN_NUM_PREDICT,
(estimated_chars + OUTPUT_TOKEN_ESTIMATE_CHARS - 1)
// OUTPUT_TOKEN_ESTIMATE_CHARS,
)
def resolve_num_predict(
requested_num_predict: int | None,
facts: list[dict[str, Any]],
prompt_token_estimate: int,
num_ctx: int,
) -> int:
if requested_num_predict is not None:
return requested_num_predict
estimated = estimate_response_tokens(facts)
max_available = max(
DEFAULT_MIN_NUM_PREDICT,
num_ctx - prompt_token_estimate - OUTPUT_CONTEXT_RESERVE_TOKENS,
)
return min(estimated, max_available)
def response_text_from_ollama_data(data: dict[str, Any]) -> str | None:
text = data.get("response")
if isinstance(text, str) and text.strip():
return text
message = data.get("message")
if isinstance(message, dict):
content = message.get("content")
if isinstance(content, str) and content.strip():
return content
return text if isinstance(text, str) else None
def build_ollama_payload(
model: str,
prompt: str,
num_ctx: int,
num_predict: int,
think: bool,
) -> dict[str, Any]:
return {
"model": model,
"prompt": prompt,
"think": think,
"stream": False,
"format": "json",
"options": {
"temperature": 0.0,
"num_ctx": num_ctx,
"num_predict": num_predict,
},
}
def print_response_metadata(data: dict[str, Any]) -> None:
fields = [
"total_duration",
"load_duration",
"prompt_eval_count",
"prompt_eval_duration",
"eval_count",
"eval_duration",
]
present = [(field, data.get(field)) for field in fields if field in data]
if not present:
print("Ollama response metadata: unavailable")
return
print("Ollama response metadata:")
for field, value in present:
print(f" {field}: {value}")
def post_with_progress(
endpoint: str,
payload: dict[str, Any],
timeout: int,
progress_interval: int,
) -> tuple[requests.Response, float]:
results: Queue[tuple[str, requests.Response | BaseException, float]] = Queue()
def worker() -> None:
started = time.perf_counter()
try:
response = requests.post(endpoint, json=payload, timeout=timeout)
except BaseException as exc: # noqa: BLE001 - forwarded to main thread.
results.put(("error", exc, time.perf_counter() - started))
else:
results.put(("response", response, time.perf_counter() - started))
started_wait = time.perf_counter()
thread = Thread(target=worker, daemon=True)
thread.start()
while True:
try:
status, result, elapsed = results.get(timeout=max(progress_interval, 1))
except Empty:
print(
"Waiting for Ollama response: "
f"{time.perf_counter() - started_wait:.1f} seconds elapsed",
flush=True,
)
continue
if status == "error":
raise result
return result, elapsed
def call_ollama(
endpoint: str,
model: str,
prompt: str,
timeout: int,
num_ctx: int,
num_predict: int,
think: bool,
progress_interval: int,
) -> tuple[str, dict[str, Any], float]:
payload = build_ollama_payload(
model=model,
prompt=prompt,
num_ctx=num_ctx,
num_predict=num_predict,
think=think,
)
print(f"Ollama request start: {datetime.now().isoformat(timespec='seconds')}")
print(f"Ollama endpoint: {endpoint}")
print(f"Ollama model: {model}")
print(f"Ollama stream: {payload['stream']}")
print(f"Ollama think: {payload['think']}")
print(f"Ollama timeout seconds: {timeout}")
print(f"Ollama options: {json.dumps(payload['options'], sort_keys=True)}")
response, elapsed = post_with_progress(
endpoint=endpoint,
payload=payload,
timeout=timeout,
progress_interval=progress_interval,
)
response.raise_for_status()
data = response.json()
print_response_metadata(data)
text = response_text_from_ollama_data(data)
if not isinstance(text, str) or not text.strip():
raise ValueError("Ollama returned no usable response text.")
return text.strip(), data, elapsed
def parse_model_json(text: str) -> dict[str, Any]:
try:
parsed = json.loads(text)
except json.JSONDecodeError as exc:
raise ConsolidationValidationError(f"Invalid model JSON: {exc}") from exc
if not isinstance(parsed, dict):
raise ConsolidationValidationError("Model JSON must be an object.")
return parsed
def group_repetition_signature(group: dict[str, Any]) -> tuple[str, tuple[str, ...], str] | None:
canonical_text = group.get("canonical_text")
source_item_ids = group.get("source_item_ids")
merge_reason = group.get("merge_reason")
if (
not isinstance(canonical_text, str)
or not isinstance(source_item_ids, list)
or not all(isinstance(item_id, str) for item_id in source_item_ids)
or not isinstance(merge_reason, str)
):
return None
normalize = lambda value: re.sub(r"\s+", " ", value).strip()
return (
normalize(canonical_text),
tuple(source_item_ids),
normalize(merge_reason),
)
def completed_group_objects(text: str) -> list[dict[str, Any]]:
"""Extract complete JSON group objects even when the outer response is truncated."""
groups_key = re.search(r'"groups"\s*:\s*\[', text)
if groups_key is None:
return []
decoder = json.JSONDecoder()
groups: list[dict[str, Any]] = []
position = groups_key.end()
while position < len(text):
while position < len(text) and text[position] in " \t\r\n,":
position += 1
if position >= len(text) or text[position] == "]":
break
if text[position] != "{":
break
try:
value, position = decoder.raw_decode(text, position)
except json.JSONDecodeError:
break
if not isinstance(value, dict):
break
groups.append(value)
return groups
def detect_repetitive_group_loop(
text: str,
minimum_consecutive: int = REPETITION_LOOP_MIN_CONSECUTIVE_GROUPS,
) -> dict[str, Any]:
groups = completed_group_objects(text)
longest_run = 0
longest_signature: tuple[str, tuple[str, ...], str] | None = None
run_start = 0
previous: tuple[str, tuple[str, ...], str] | None = None
current_run = 0
for index, group in enumerate(groups):
signature = group_repetition_signature(group)
if signature is not None and signature == previous:
current_run += 1
else:
current_run = 1 if signature is not None else 0
run_start = index
if current_run > longest_run:
longest_run = current_run
longest_signature = signature
longest_run_start = run_start
previous = signature
detected = longest_signature is not None and longest_run >= minimum_consecutive
signature_data = None
if longest_signature is not None:
signature_data = {
"canonical_text": longest_signature[0],
"source_item_ids": list(longest_signature[1]),
"merge_reason": longest_signature[2],
}
return {
"detected": detected,
"threshold": minimum_consecutive,
"completed_group_count": len(groups),
"longest_consecutive_run": longest_run,
"run_start_group_index": longest_run_start if longest_run else None,
"signature": signature_data,
}
def call_ollama_with_repetition_retry(
*,
endpoint: str,
model: str,
prompt: str,
timeout: int,
num_ctx: int,
num_predict: int,
think: bool,
progress_interval: int,
output_dir: Path,
call_fn: Any = None,
) -> tuple[str, dict[str, Any], float, list[dict[str, Any]]]:
"""Make one call, retrying once only for invalid JSON with a detected group loop."""
if call_fn is None:
call_fn = call_ollama
output_dir.mkdir(parents=True, exist_ok=True)
attempts: list[dict[str, Any]] = []
total_runtime = 0.0
for attempt_number in (1, 2):
attempt_prompt = (
prompt
if attempt_number == 1
else f"{prompt}\n{REPETITION_RETRY_INSTRUCTION}\n"
)
raw_text, response_data, runtime = call_fn(
endpoint=endpoint,
model=model,
prompt=attempt_prompt,
timeout=timeout,
num_ctx=num_ctx,
num_predict=num_predict,
think=think,
progress_interval=progress_interval,
)
total_runtime += runtime
suffix = "" if attempt_number == 1 else "_retry"
(output_dir / f"raw_model_response{suffix}.txt").write_text(
raw_text + "\n", encoding="utf-8"
)
write_json(output_dir / f"raw_ollama_response{suffix}.json", response_data)
repetition = detect_repetitive_group_loop(raw_text)
attempt_metadata = {
"attempt": attempt_number,
"retry_instruction_added": attempt_number == 2,
"runtime_seconds": runtime,
"raw_response_chars": len(raw_text),
"raw_response_bytes": len(raw_text.encode("utf-8")),
"resolved_num_predict": num_predict,
"eval_count": response_data.get("eval_count"),
"done_reason": response_data.get("done_reason"),
"repetition": repetition,
}
attempts.append(attempt_metadata)
try:
parse_model_json(raw_text)
except ConsolidationValidationError as exc:
attempt_metadata["parse_error"] = str(exc)
write_json(
output_dir / "repetition_retry_metadata.json",
{"bug": "BUG-013", "attempts": attempts},
)
if attempt_number == 1 and repetition["detected"]:
continue
raise
write_json(
output_dir / "repetition_retry_metadata.json",
{"bug": "BUG-013", "attempts": attempts},
)
return raw_text, response_data, total_runtime, attempts
raise AssertionError("repetition retry loop exhausted unexpectedly")
def validate_model_groups(
model_output: dict[str, Any],
expected_fact_ids: set[str],
) -> list[dict[str, Any]]:
groups = model_output.get("groups")
if not isinstance(groups, list):
raise ConsolidationValidationError("Model output must contain a groups list.")
seen: list[str] = []
validated: list[dict[str, Any]] = []
for index, group in enumerate(groups):
if not isinstance(group, dict):
raise ConsolidationValidationError(f"Group {index} must be an object.")
source_ids = group.get("source_item_ids")
if not isinstance(source_ids, list) or not source_ids:
raise ConsolidationValidationError(
f"Group {index} must contain source_item_ids."
)
if not all(isinstance(item_id, str) for item_id in source_ids):
raise ConsolidationValidationError(
f"Group {index} source_item_ids must be strings."
)
unknown = sorted(set(source_ids) - expected_fact_ids)
if unknown:
raise ConsolidationValidationError(
f"Group {index} contains unknown source item IDs: {unknown}"
)
duplicates_in_group = [
item_id for item_id, count in Counter(source_ids).items() if count > 1
]
if duplicates_in_group:
raise ConsolidationValidationError(
f"Group {index} repeats source item IDs: {duplicates_in_group}"
)
canonical_text = group.get("canonical_text")
if not isinstance(canonical_text, str) or not canonical_text.strip():
raise ConsolidationValidationError(
f"Group {index} must contain canonical_text."
)
merge_reason = group.get("merge_reason")
if not isinstance(merge_reason, str) or not merge_reason.strip():
raise ConsolidationValidationError(
f"Group {index} must contain merge_reason."
)
seen.extend(source_ids)
validated.append(
{
"canonical_text": canonical_text.strip(),
"source_item_ids": source_ids,
"merge_reason": merge_reason.strip(),
}
)
seen_counts = Counter(seen)
duplicated = sorted(item_id for item_id, count in seen_counts.items() if count > 1)
if duplicated:
raise ConsolidationValidationError(
f"Source item IDs appear in multiple groups: {duplicated}"
)
missing = sorted(expected_fact_ids - set(seen))
if missing:
raise ConsolidationValidationError(f"Missing source item IDs: {missing}")
return validated
def validate_group_shapes(groups: list[dict[str, Any]]) -> None:
for group in groups:
source_count = len(group["source_item_ids"])
if source_count < 1:
raise ConsolidationValidationError("Groups must not be empty.")
if source_count == 1:
continue
if source_count < 2:
raise ConsolidationValidationError("Merged groups need at least two IDs.")
def repair_model_group_coverage(
model_output: dict[str, Any],
facts: list[dict[str, Any]],
) -> tuple[dict[str, Any], list[dict[str, Any]]]:
"""
Apply deterministic source-coverage repairs to model grouping JSON.
The repair is intentionally conservative. Repeated source IDs are removed
after their first occurrence, empty groups created by that removal are
dropped, and missing facts are restored as singleton groups using the
original canonicalized fact text. No existing group text, merge reason or
semantic merge is rewritten.
"""
groups = model_output.get("groups")
if not isinstance(groups, list):
return model_output, []
repaired = copy.deepcopy(model_output)
repaired_groups = repaired["groups"]
facts_by_id = {str(item.get("item_id")): item for item in facts}
expected_ids = set(facts_by_id)
seen: set[str] = set()
changes: list[dict[str, Any]] = []
for group_index, group in enumerate(repaired_groups):
if not isinstance(group, dict):
continue
source_ids = group.get("source_item_ids")
if not isinstance(source_ids, list):
continue
kept_ids: list[str] = []
for id_index, item_id in enumerate(source_ids):
if not isinstance(item_id, str) or item_id not in expected_ids:
kept_ids.append(item_id)
continue
if item_id in seen:
changes.append(
{
"operation": "remove_duplicate_source_id",
"id": item_id,
"group_index": group_index,
"id_index": id_index,
}
)
continue
seen.add(item_id)
kept_ids.append(item_id)
group["source_item_ids"] = kept_ids
non_empty_groups: list[dict[str, Any]] = []
for group_index, group in enumerate(repaired_groups):
if (
isinstance(group, dict)
and isinstance(group.get("source_item_ids"), list)
and len(group["source_item_ids"]) == 0
):
changes.append(
{
"operation": "remove_empty_group",
"group_index": group_index,
"canonical_text": group.get("canonical_text"),
}
)
continue
non_empty_groups.append(group)
repaired["groups"] = non_empty_groups
missing_ids = sorted(expected_ids - seen)
for item_id in missing_ids:
fact = facts_by_id[item_id]
repaired["groups"].append(
{
"canonical_text": str(fact.get("text", "")).strip(),
"source_item_ids": [item_id],
"merge_reason": (
"Deterministic coverage repair: source fact was missing "
"from the model grouping and is preserved as a singleton."
),
}
)
changes.append(
{
"operation": "restore_missing_source_id_as_singleton",
"id": item_id,
}
)
return repaired, changes
def build_consolidated_fact_item(
group: dict[str, Any],
fact_by_id: dict[str, dict[str, Any]],
sequence: int,
) -> dict[str, Any]:
source_ids = group["source_item_ids"]
source_items = [fact_by_id[item_id] for item_id in source_ids]
source_references: list[dict[str, Any]] = []
evidence: list[str] = []
for item in source_items:
source_references.extend(item.get("source_references", []))
item_evidence = item.get("evidence")
if isinstance(item_evidence, str):
evidence.append(item_evidence)
return {
"consolidated_id": f"fact_group_{sequence:04d}",
"category": "fact",
"canonical_text": group["canonical_text"],
"source_item_ids": source_ids,
"source_references": source_references,
"evidence": evidence,
"merge_reason": group["merge_reason"],
}
def build_consolidated_output(
canonicalized: dict[str, Any],
groups: list[dict[str, Any]],
) -> dict[str, Any]:
facts = fact_items(canonicalized)
fact_by_id = {item["item_id"]: item for item in facts}
consolidated_facts = [
build_consolidated_fact_item(group, fact_by_id, index)
for index, group in enumerate(groups, start=1)
]
output_items = consolidated_facts + non_fact_items(canonicalized)
merged_groups = [item for item in consolidated_facts if len(item["source_item_ids"]) > 1]
singleton_groups = [
item for item in consolidated_facts if len(item["source_item_ids"]) == 1
]
output = dict(canonicalized)
output["schema_version"] = "semantic_consolidator_v0"
output["items"] = output_items
output["semantic_consolidation"] = {
"scope": "facts_only",
"merged_fact_group_count": len(merged_groups),
"source_facts_in_merged_groups": sum(
len(item["source_item_ids"]) for item in merged_groups
),
"singleton_fact_group_count": len(singleton_groups),
}
return output
def validate_consolidated_output(
canonicalized: dict[str, Any],
output: dict[str, Any],
) -> None:
original_facts = fact_items(canonicalized)
expected_fact_ids = {item["item_id"] for item in original_facts}
output_items = output.get("items")
if not isinstance(output_items, list):
raise ConsolidationValidationError("Output items must be a list.")
output_fact_groups = [item for item in output_items if item.get("category") == "fact"]
seen: list[str] = []
for item in output_fact_groups:
source_ids = item.get("source_item_ids")
if not isinstance(source_ids, list):
raise ConsolidationValidationError("Fact groups need source_item_ids.")
if len(source_ids) == 0:
raise ConsolidationValidationError("Fact groups must not be empty.")
if len(source_ids) > 1 and not item.get("merge_reason"):
raise ConsolidationValidationError("Merged fact groups need merge_reason.")
seen.extend(source_ids)
source_refs = item.get("source_references")
if not isinstance(source_refs, list) or not source_refs:
raise ConsolidationValidationError("Fact groups need source references.")
seen_counts = Counter(seen)
duplicated = sorted(item_id for item_id, count in seen_counts.items() if count > 1)
if duplicated:
raise ConsolidationValidationError(
f"Output duplicates source fact IDs: {duplicated}"
)
missing = sorted(expected_fact_ids - set(seen))
if missing:
raise ConsolidationValidationError(f"Output misses source fact IDs: {missing}")
unknown = sorted(set(seen) - expected_fact_ids)
if unknown:
raise ConsolidationValidationError(f"Output has unknown source IDs: {unknown}")
original_non_facts = non_fact_items(canonicalized)
output_non_facts = [item for item in output_items if item.get("category") != "fact"]
if output_non_facts != original_non_facts:
raise ConsolidationValidationError("Non-fact categories changed.")
def write_json(path: Path, data: Any) -> None:
path.parent.mkdir(parents=True, exist_ok=True)
path.write_text(
json.dumps(data, ensure_ascii=False, indent=2) + "\n",
encoding="utf-8",
)
def write_report(
path: Path,
model: str,
runtime: float,
prompt_chars: int,
prompt_token_estimate: int,
num_predict: int,
fact_count: int,
groups: list[dict[str, Any]],
output_path: Path,
repair_changes: list[dict[str, Any]] | None = None,
llm_call_count: int = 1,
) -> None:
merged = [group for group in groups if len(group["source_item_ids"]) > 1]
singletons = [group for group in groups if len(group["source_item_ids"]) == 1]
lines = [
"# Semantic Consolidator V0 Report",
"",
"- Scope: facts only",
f"- Model: `{model}`",
f"- LLM call count: {llm_call_count}",
f"- Runtime: {runtime:.3f} seconds",
f"- Fact item count: {fact_count}",
f"- Prompt characters: {prompt_chars}",
f"- Estimated prompt tokens: {prompt_token_estimate}",
f"- num_predict: {num_predict}",
f"- Merged fact groups: {len(merged)}",
f"- Source facts involved in merges: {sum(len(group['source_item_ids']) for group in merged)}",
f"- Singleton fact groups: {len(singletons)}",
f"- Deterministic repair changes: {len(repair_changes or [])}",
f"- Output path: `{output_path}`",
"",
"## Actual Merges",
"",
]
if not merged:
lines.append("- None.")
else:
for group in merged:
lines.extend(
[
f"### {group['canonical_text']}",
"",
f"- Source fact IDs: {', '.join(group['source_item_ids'])}",
f"- Merge reason: {group['merge_reason']}",
"",
]
)
if repair_changes:
lines.extend(["", "## Deterministic Coverage Repairs", ""])
for change in repair_changes:
lines.append(f"- `{change['operation']}`: {json.dumps(change, ensure_ascii=False, sort_keys=True)}")
path.write_text("\n".join(lines).rstrip() + "\n", encoding="utf-8")
def main() -> int:
args = parse_args()
args.output_dir.mkdir(parents=True, exist_ok=True)
raw_response_path = args.output_dir / "raw_model_response.txt"
output_path = args.output_dir / "consolidated_extractions.json"
report_path = args.output_dir / "report.md"
repair_metadata_path = args.output_dir / "repair_metadata.json"
try:
canonicalized = load_json_object(args.canonicalized_input)
facts = fact_items(canonicalized)
prompt = build_consolidation_prompt(facts)
prompt_chars = len(prompt)
prompt_token_estimate = (prompt_chars + 3) // 4
num_predict = resolve_num_predict(
requested_num_predict=args.num_predict,
facts=facts,
prompt_token_estimate=prompt_token_estimate,
num_ctx=args.num_ctx,
)
print(f"Fact item count: {len(facts)}")
print(f"Estimated prompt size chars: {prompt_chars}")
print(f"Estimated prompt tokens: {prompt_token_estimate}")
print(f"Resolved num_predict: {num_predict}")
print("Expected LLM call count: 1; at most 2 only after detected BUG-013 loop")
print("Expected runtime: 5-10 minutes on current local benchmark basis")
raw_text, _response_data, runtime, attempts = call_ollama_with_repetition_retry(
endpoint=args.endpoint,
model=args.model,
prompt=prompt,
timeout=args.timeout,
num_ctx=args.num_ctx,
num_predict=num_predict,
think=args.think,
progress_interval=args.progress_interval,
output_dir=args.output_dir,
)
model_output = parse_model_json(raw_text)
model_output, repair_changes = repair_model_group_coverage(model_output, facts)
if repair_changes:
write_json(
repair_metadata_path,
{
"scope": "semantic_consolidator_v0_source_coverage",
"llm_used": False,
"repair_count": len(repair_changes),
"repairs": repair_changes,
},
)
expected_fact_ids = {item["item_id"] for item in facts}
groups = validate_model_groups(model_output, expected_fact_ids)
validate_group_shapes(groups)
output = build_consolidated_output(canonicalized, groups)
validate_consolidated_output(canonicalized, output)
write_json(output_path, output)
write_report(
path=report_path,
model=args.model,
runtime=runtime,
prompt_chars=prompt_chars,
prompt_token_estimate=prompt_token_estimate,
num_predict=num_predict,
fact_count=len(facts),
groups=groups,
output_path=output_path,
repair_changes=repair_changes,
llm_call_count=len(attempts),
)
except requests.ConnectionError as exc:
print(f"Error: Ollama is not reachable at {args.endpoint}: {exc}", file=sys.stderr)
return 1
except requests.Timeout as exc:
print(
f"Error: Ollama request timed out after {args.timeout} seconds: {exc}",
file=sys.stderr,
)
return 1
except requests.HTTPError as exc:
print(f"Error: Ollama returned an HTTP error: {exc}", file=sys.stderr)
return 1
except (OSError, UnicodeError, ValueError) as exc:
print(f"Error: {exc}", file=sys.stderr)
return 1
merged = [group for group in groups if len(group["source_item_ids"]) > 1]
print(f"Runtime seconds: {runtime:.3f}")
print(f"Merged fact groups: {len(merged)}")
print(
"Source facts involved in merges: "
f"{sum(len(group['source_item_ids']) for group in merged)}"
)
print(f"Singleton fact groups: {len(groups) - len(merged)}")
print("Validation result: passed")
print(f"Output: {output_path}")
print(f"Report: {report_path}")
print(f"Raw model response: {raw_response_path}")
return 0
if __name__ == "__main__":
raise SystemExit(main())