Handle repetitive Semantic Consolidator output loops
This commit is contained in:
@@ -1,6 +1,8 @@
|
||||
import unittest
|
||||
import json
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
from unittest.mock import Mock
|
||||
|
||||
from src.meeting_lab.consolidation.consolidate_facts import (
|
||||
ConsolidationValidationError,
|
||||
@@ -8,6 +10,8 @@ from src.meeting_lab.consolidation.consolidate_facts import (
|
||||
DEFAULT_NUM_PREDICT,
|
||||
build_consolidated_output,
|
||||
build_ollama_payload,
|
||||
call_ollama_with_repetition_retry,
|
||||
detect_repetitive_group_loop,
|
||||
estimate_response_tokens,
|
||||
fact_items,
|
||||
parse_model_json,
|
||||
@@ -74,6 +78,14 @@ def canonicalized_fixture():
|
||||
|
||||
|
||||
class ConsolidateFactsTests(unittest.TestCase):
|
||||
@staticmethod
|
||||
def group(text="A", ids=None, reason="Singleton."):
|
||||
return {
|
||||
"canonical_text": text,
|
||||
"source_item_ids": ids or ["fact_0001"],
|
||||
"merge_reason": reason,
|
||||
}
|
||||
|
||||
def test_payload_construction_disables_streaming_and_thinking_by_default(self):
|
||||
payload = build_ollama_payload(
|
||||
model="qwen3.5:9B",
|
||||
@@ -174,6 +186,145 @@ class ConsolidateFactsTests(unittest.TestCase):
|
||||
|
||||
self.assertEqual(len(groups), 2)
|
||||
|
||||
def test_valid_normal_consolidation_has_no_repetition_loop(self):
|
||||
text = json.dumps(
|
||||
{"groups": [self.group("A"), self.group("B", ["fact_0002"])]}
|
||||
)
|
||||
|
||||
detection = detect_repetitive_group_loop(text)
|
||||
|
||||
self.assertFalse(detection["detected"])
|
||||
self.assertEqual(detection["completed_group_count"], 2)
|
||||
|
||||
def test_ordinary_duplicate_group_is_not_a_repetition_loop(self):
|
||||
duplicate = self.group()
|
||||
text = json.dumps({"groups": [duplicate, duplicate]})
|
||||
|
||||
detection = detect_repetitive_group_loop(text)
|
||||
|
||||
self.assertFalse(detection["detected"])
|
||||
self.assertEqual(detection["longest_consecutive_run"], 2)
|
||||
|
||||
def test_repetitive_identical_group_sequence_is_detected_structurally(self):
|
||||
repeated = self.group("Same fact", ["fact_0001", "fact_0002"], "Equivalent.")
|
||||
text = json.dumps({"groups": [repeated, repeated, repeated]})
|
||||
|
||||
detection = detect_repetitive_group_loop(text)
|
||||
|
||||
self.assertTrue(detection["detected"])
|
||||
self.assertEqual(detection["longest_consecutive_run"], 3)
|
||||
self.assertEqual(
|
||||
detection["signature"],
|
||||
{
|
||||
"canonical_text": "Same fact",
|
||||
"source_item_ids": ["fact_0001", "fact_0002"],
|
||||
"merge_reason": "Equivalent.",
|
||||
},
|
||||
)
|
||||
|
||||
def test_truncated_response_after_repetitive_loop_is_detected(self):
|
||||
repeated = json.dumps(self.group(), separators=(",", ":"))
|
||||
text = '{"groups":[' + ",".join([repeated] * 4) + ',{"canonical_text":'
|
||||
|
||||
detection = detect_repetitive_group_loop(text)
|
||||
|
||||
self.assertTrue(detection["detected"])
|
||||
self.assertEqual(detection["completed_group_count"], 4)
|
||||
|
||||
def test_detected_invalid_loop_gets_one_controlled_retry(self):
|
||||
repeated = json.dumps(self.group(), separators=(",", ":"))
|
||||
truncated = '{"groups":[' + ",".join([repeated] * 3) + ',{"canonical_text":'
|
||||
valid = json.dumps({"groups": [self.group()]})
|
||||
call_fn = Mock(
|
||||
side_effect=[
|
||||
(truncated, {"eval_count": 100, "done_reason": "length"}, 1.0),
|
||||
(valid, {"eval_count": 20, "done_reason": "stop"}, 0.5),
|
||||
]
|
||||
)
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
raw_text, _data, runtime, attempts = call_ollama_with_repetition_retry(
|
||||
endpoint="http://example.invalid",
|
||||
model="test-model",
|
||||
prompt="base prompt",
|
||||
timeout=1,
|
||||
num_ctx=1000,
|
||||
num_predict=100,
|
||||
think=False,
|
||||
progress_interval=1,
|
||||
output_dir=Path(directory),
|
||||
call_fn=call_fn,
|
||||
)
|
||||
|
||||
self.assertTrue((Path(directory) / "raw_model_response.txt").exists())
|
||||
self.assertTrue((Path(directory) / "raw_model_response_retry.txt").exists())
|
||||
metadata = json.loads(
|
||||
(Path(directory) / "repetition_retry_metadata.json").read_text()
|
||||
)
|
||||
|
||||
self.assertEqual(raw_text, valid)
|
||||
self.assertEqual(runtime, 1.5)
|
||||
self.assertEqual(len(attempts), 2)
|
||||
self.assertEqual(call_fn.call_count, 2)
|
||||
self.assertEqual(metadata["attempts"][0]["repetition"]["detected"], True)
|
||||
first_prompt = call_fn.call_args_list[0].kwargs["prompt"]
|
||||
retry_prompt = call_fn.call_args_list[1].kwargs["prompt"]
|
||||
self.assertEqual(first_prompt, "base prompt")
|
||||
self.assertIn("Emit each semantic group exactly once", retry_prompt)
|
||||
|
||||
def test_unrelated_malformed_json_is_not_retried(self):
|
||||
call_fn = Mock(return_value=("{invalid", {"done_reason": "length"}, 0.1))
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
with self.assertRaises(ConsolidationValidationError):
|
||||
call_ollama_with_repetition_retry(
|
||||
endpoint="http://example.invalid",
|
||||
model="test-model",
|
||||
prompt="prompt",
|
||||
timeout=1,
|
||||
num_ctx=1000,
|
||||
num_predict=100,
|
||||
think=False,
|
||||
progress_interval=1,
|
||||
output_dir=Path(directory),
|
||||
call_fn=call_fn,
|
||||
)
|
||||
|
||||
self.assertEqual(call_fn.call_count, 1)
|
||||
|
||||
def test_retry_failure_preserves_both_attempts_and_raises(self):
|
||||
repeated = json.dumps(self.group(), separators=(",", ":"))
|
||||
truncated = '{"groups":[' + ",".join([repeated] * 3) + ',{"canonical_text":'
|
||||
call_fn = Mock(
|
||||
side_effect=[
|
||||
(truncated, {"done_reason": "length"}, 1.0),
|
||||
(truncated, {"done_reason": "length"}, 1.0),
|
||||
]
|
||||
)
|
||||
with tempfile.TemporaryDirectory() as directory:
|
||||
root = Path(directory)
|
||||
with self.assertRaises(ConsolidationValidationError):
|
||||
call_ollama_with_repetition_retry(
|
||||
endpoint="http://example.invalid",
|
||||
model="test-model",
|
||||
prompt="prompt",
|
||||
timeout=1,
|
||||
num_ctx=1000,
|
||||
num_predict=100,
|
||||
think=False,
|
||||
progress_interval=1,
|
||||
output_dir=root,
|
||||
call_fn=call_fn,
|
||||
)
|
||||
|
||||
self.assertTrue((root / "raw_model_response.txt").exists())
|
||||
self.assertTrue((root / "raw_ollama_response.json").exists())
|
||||
self.assertTrue((root / "raw_model_response_retry.txt").exists())
|
||||
self.assertTrue((root / "raw_ollama_response_retry.json").exists())
|
||||
metadata = json.loads((root / "repetition_retry_metadata.json").read_text())
|
||||
|
||||
self.assertEqual(call_fn.call_count, 2)
|
||||
self.assertEqual(len(metadata["attempts"]), 2)
|
||||
self.assertIn("parse_error", metadata["attempts"][1])
|
||||
|
||||
def test_no_missing_source_ids(self):
|
||||
with self.assertRaisesRegex(ConsolidationValidationError, "Missing"):
|
||||
validate_model_groups(
|
||||
|
||||
Reference in New Issue
Block a user