163 lines
7.1 KiB
Python
163 lines
7.1 KiB
Python
import argparse
|
|
import json
|
|
import tempfile
|
|
import unittest
|
|
from copy import deepcopy
|
|
from pathlib import Path
|
|
from unittest.mock import patch
|
|
|
|
import src.meeting_lab.controlled_semantic_derivation.experiment_negative_act as module
|
|
from src.meeting_lab.controlled_semantic_derivation.experiment_negative_act import (
|
|
DerivationValidationError,
|
|
build_ollama_payload,
|
|
build_prompt,
|
|
evaluate_case,
|
|
load_gold_cases,
|
|
parse_model_json,
|
|
run_experiment,
|
|
validate_classification,
|
|
)
|
|
|
|
|
|
GOLD_PATH = Path("tests/gold/negative_act_form_v0/cases.json")
|
|
|
|
|
|
FORM_TEXT = {
|
|
"NA-01": "Zusammenarbeit mit Dr. Schlummer fortsetzen",
|
|
"NA-02": "externe Lösung weiterverfolgen",
|
|
"NA-03": "reale Anlage für den Versuch nutzen",
|
|
"NA-04": "reale Anlage verwenden",
|
|
"NA-05": "Waschstufe einbauen",
|
|
}
|
|
|
|
|
|
def classification_for(case):
|
|
expected = case["expected"]
|
|
return {
|
|
"observation_id": expected["observation_id"],
|
|
"negative_act_form": expected["negative_act_form"],
|
|
"normalized_action_text": FORM_TEXT.get(case["case_id"]),
|
|
}
|
|
|
|
|
|
class NegativeActFormExperimentTests(unittest.TestCase):
|
|
@classmethod
|
|
def setUpClass(cls):
|
|
cls.cases = load_gold_cases(GOLD_PATH)
|
|
cls.by_id = {case["case_id"]: case for case in cls.cases}
|
|
|
|
def test_fixture_contains_exactly_na_01_through_na_08(self):
|
|
self.assertEqual(list(self.by_id), [f"NA-{number:02d}" for number in range(1, 9)])
|
|
|
|
def test_exact_schema_is_accepted(self):
|
|
case = self.by_id["NA-01"]
|
|
self.assertEqual(validate_classification(classification_for(case), case["observations"]), classification_for(case))
|
|
|
|
def test_unknown_field_is_rejected(self):
|
|
case = self.by_id["NA-01"]
|
|
classification = classification_for(case)
|
|
classification["explanation"] = "extra"
|
|
with self.assertRaisesRegex(DerivationValidationError, "unknown keys"):
|
|
validate_classification(classification, case["observations"])
|
|
|
|
def test_invalid_enum_is_rejected(self):
|
|
case = self.by_id["NA-01"]
|
|
classification = classification_for(case)
|
|
classification["negative_act_form"] = "rejection"
|
|
with self.assertRaisesRegex(DerivationValidationError, "unsupported value"):
|
|
validate_classification(classification, case["observations"])
|
|
|
|
def test_non_none_requires_normalized_action_text(self):
|
|
case = self.by_id["NA-01"]
|
|
for value in (None, ""):
|
|
classification = classification_for(case)
|
|
classification["normalized_action_text"] = value
|
|
with self.subTest(value=value), self.assertRaises(DerivationValidationError):
|
|
validate_classification(classification, case["observations"])
|
|
|
|
def test_none_requires_null_normalized_action_text(self):
|
|
case = self.by_id["NA-06"]
|
|
classification = classification_for(case)
|
|
self.assertIsNone(classification["normalized_action_text"])
|
|
classification["normalized_action_text"] = "Material einsetzen"
|
|
with self.assertRaisesRegex(DerivationValidationError, "requires null"):
|
|
validate_classification(classification, case["observations"])
|
|
|
|
def test_forbidden_normative_fields_are_rejected_recursively(self):
|
|
case = self.by_id["NA-01"]
|
|
fields = (
|
|
"rejection_form", "explicitly_rejected", "status", "decision", "outcome",
|
|
"topic_status", "responsible_person", "responsibility", "owner",
|
|
"requested_actor", "action_item", "protocol_category", "confidence",
|
|
"relation", "relations", "graph", "unresolved_issue",
|
|
)
|
|
for field in fields:
|
|
classification = classification_for(case)
|
|
classification["wrapper"] = {field: "forbidden"}
|
|
with self.subTest(field=field), self.assertRaisesRegex(DerivationValidationError, "forbidden semantic keys"):
|
|
validate_classification(classification, case["observations"])
|
|
|
|
def test_unknown_observation_id_is_rejected(self):
|
|
case = self.by_id["NA-01"]
|
|
classification = classification_for(case)
|
|
classification["observation_id"] = "obs_99"
|
|
with self.assertRaisesRegex(DerivationValidationError, "unknown observation"):
|
|
validate_classification(classification, case["observations"])
|
|
|
|
def test_malformed_json_is_rejected(self):
|
|
with self.assertRaises(json.JSONDecodeError):
|
|
parse_model_json("{bad json")
|
|
|
|
def test_all_expected_classifications_evaluate_as_pass(self):
|
|
for case in self.cases:
|
|
evaluation = evaluate_case(case, classification_for(case))
|
|
with self.subTest(case=case["case_id"]):
|
|
self.assertEqual(evaluation["classification"], "PASS")
|
|
|
|
def test_fixed_prompt_contains_candidate_and_no_gold_expectation(self):
|
|
prompt = build_prompt(self.by_id["NA-03"])
|
|
self.assertIn("candidate observation is obs_2", prompt)
|
|
self.assertNotIn("expected", prompt)
|
|
self.assertNotIn("Who is responsible", prompt)
|
|
|
|
def test_fixed_model_configuration(self):
|
|
payload = build_ollama_payload("qwen3.5:9B", "prompt", 16384, 1024)
|
|
self.assertFalse(payload["think"])
|
|
self.assertFalse(payload["stream"])
|
|
self.assertEqual(payload["options"]["temperature"], 0)
|
|
|
|
def test_no_rejection_or_status_derivation_function_exists(self):
|
|
public_names = {name for name in dir(module) if not name.startswith("_")}
|
|
self.assertNotIn("derive_rejection", public_names)
|
|
self.assertFalse(any(name.startswith("derive_") for name in public_names))
|
|
|
|
def test_artifacts_preserve_semantic_classification_only(self):
|
|
case = self.by_id["NA-01"]
|
|
raw = json.dumps(classification_for(case), ensure_ascii=False)
|
|
with tempfile.TemporaryDirectory() as temporary:
|
|
output = Path(temporary) / "run"
|
|
args = argparse.Namespace(
|
|
cases=GOLD_PATH, output=output, model="qwen3.5:9B",
|
|
endpoint="http://unused", timeout=1, num_ctx=16384, num_predict=1024,
|
|
)
|
|
with patch.object(module, "load_gold_cases", return_value=[deepcopy(case)]), patch.object(
|
|
module, "call_ollama", return_value=(raw, {"model": "qwen3.5:9B"})
|
|
):
|
|
summary = run_experiment(args)
|
|
self.assertEqual(summary["successful_llm_call_count"], 1)
|
|
case_dir = output / "na-01"
|
|
for filename in (
|
|
"v3_style_input_observations.json", "prompt.txt", "raw_model_response.txt",
|
|
"parsed_semantic_classification.json", "structural_validation.json",
|
|
"evaluation.json", "ollama_metadata.json",
|
|
):
|
|
self.assertTrue((case_dir / filename).is_file(), filename)
|
|
self.assertFalse((case_dir / "final_derived_result.json").exists())
|
|
self.assertFalse((case_dir / "deterministic_gate_results.json").exists())
|
|
parsed = json.loads((case_dir / "parsed_semantic_classification.json").read_text())
|
|
self.assertEqual(set(parsed), {"observation_id", "negative_act_form", "normalized_action_text"})
|
|
|
|
|
|
if __name__ == "__main__":
|
|
unittest.main()
|