948 lines
33 KiB
Python
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())
|