File size: 5,757 Bytes
9d648e6
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
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()