import unittest from types import SimpleNamespace from decision import candidate_softmax, decide, normalize_case, orders_for, render_prompt class Tokenizer: def apply_chat_template(self, messages, **kwargs): return repr(messages) def decode(self, ids, **kwargs): return "".join({1: "A", 2: "B", 3: "C", 4: "STOP", 5: "\n"}[i] for i in ids) LABELS = {"labels": ["A", "B", "C"], "ids": [1, 2, 3], "stop_id": 4, "newline_id": 5} class Backend: processor = SimpleNamespace(tokenizer=Tokenizer()) def __init__(self, scores): self.scores = iter(scores) self.calls = [] def logits(self, prompt, images, allowed): self.calls.append((prompt, allowed)) return next(self.scores) class DecisionTests(unittest.TestCase): def case(self, multi=False): return {"id": "test", "state": "context", "question": "Select", "options": {"x": "one", "y": "two"}, "kind": "multi_choice" if multi else "choice"} def test_effort2_maps_permuted_token_back_to_original_key(self): backend = Backend([[0, 10], [0, 10]]) result = decide(self.case(), backend, LABELS, effort=2) self.assertEqual(result["choice"], "y") self.assertEqual(backend.calls[0][1], [1, 2]) self.assertEqual(backend.calls[1][1], [2, 1]) self.assertGreater(result["probabilities"]["y"], .99) def test_multiselect_stop_and_no_repeated_selection(self): backend = Backend([[9, 0, -1], [-1, 9]]) result = decide(self.case(True), backend, LABELS) self.assertEqual(result["selected"], ["x"]) self.assertEqual(backend.calls[1][1], [2, 4]) self.assertTrue(backend.calls[1][0].endswith("Answer:\nA\n")) def test_multiselect_can_choose_empty_set(self): self.assertEqual(decide(self.case(True), Backend([[0, 0, 10]]), LABELS)["selected"], []) def test_probabilities_stable_and_invalid_inputs_rejected(self): self.assertAlmostEqual(sum(candidate_softmax([10000, 10001])), 1) for temperature in [0, -1, float("nan")]: with self.assertRaises(ValueError): candidate_softmax([1, 2], temperature) with self.assertRaises(ValueError): normalize_case({"question": "q", "options": [{"key": "a", "text": "1"}, {"key": "a", "text": "2"}]}) def test_public_json_state_is_decoded_once(self): case = self.case() case.update(state='{"customer": "lost card"}', state_format="json") normalized = normalize_case(case) self.assertEqual(normalized["state"], {"customer": "lost card"}) self.assertEqual(normalize_case(normalized)["state"], normalized["state"]) case["state"] = '"plain state"' self.assertEqual(normalize_case(normalize_case(case))["state"], "plain state") def test_invalid_candidate_types_and_task_shapes_are_rejected(self): for options in [{1: "one", 2: "two"}, {"a": {}, "b": "two"}, {"a": "", "b": "two"}]: with self.assertRaises(ValueError): normalize_case(dict(self.case(), options=options)) with self.assertRaises(ValueError): normalize_case(dict(self.case(), kind="noul", options={"a":"one", "b":"two", "c":"three"})) with self.assertRaises(ValueError): normalize_case(dict(self.case(), kind="score", options={"1":"first", "0":"second"})) normalize_case(dict(self.case(), kind="score", options={"0":"low", "1":"high"})) def test_excessive_images_are_rejected_during_input_validation(self): with self.assertRaises(ValueError): normalize_case(dict(self.case(), image_refs=["image.png"]*6)) with self.assertRaises(ValueError): normalize_case(dict(self.case(), image_refs="image.png")) def test_deterministic_orders(self): case = self.case() case["options"]["z"] = "three" self.assertEqual(orders_for(case, 2), orders_for(case, 2)) self.assertNotEqual(*orders_for(case, 2)) def test_effort_three_to_five_add_unique_reproducible_orders(self): case = self.case() case["options"]["z"] = "three" for effort in range(1, 6): orders = orders_for(case, effort) self.assertEqual(len(orders), effort) self.assertEqual(len(set(map(tuple, orders))), effort) self.assertEqual(orders, orders_for(case, effort)) self.assertEqual(orders[:2], orders_for(case, min(effort, 2))) self.assertTrue(all(sorted(order) == [0, 1, 2] for order in orders)) def test_binary_effort_stops_at_two_distinct_orders(self): self.assertEqual(orders_for(self.case(), 5), [[0, 1], [1, 0]]) with self.assertRaises(ValueError): orders_for(self.case(), 6) def test_effort5_averages_five_branches(self): case = self.case() case["options"]["z"] = "three" result = decide(case, Backend([[0, 10, 0]] * 5), LABELS, effort=5) self.assertEqual(len(result["steps"][0]["branches"]), 5) self.assertAlmostEqual(sum(result["probabilities"].values()), 1) def test_open_format_preserves_conversation_and_multiline_candidates(self): case = self.case() case.update(_prompt_format="open-format", state={"messages": [ {"role": "user", "content": "question"}, {"role": "assistant", "content": "answer"}]}) case["options"]["x"] = "first\nsecond" prompt = render_prompt(case, [0, 1], [], Backend.processor, LABELS, []) self.assertIn("'role': 'assistant', 'content': 'answer'", prompt) self.assertIn("first\\n second", prompt) self.assertTrue(prompt.endswith("Answer:\n")) if __name__ == "__main__": unittest.main()