File size: 4,143 Bytes
4d0d04c
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
import unittest

from muse_core import (
    META_SAMPLING,
    NATIVE_GREEDY,
    choose_seed,
    coerce_parsed_reply,
    estimate_gpu_duration,
    friendly_error,
    generation_kwargs,
    preset_values,
    render_reply,
    validate_controls,
)


class PresetTests(unittest.TestCase):
    def test_native_preset_is_greedy_with_documented_values_ready(self):
        self.assertEqual(preset_values(NATIVE_GREEDY), (False, 1.0, 0.95, 64))

    def test_meta_preset_enables_sampling(self):
        self.assertEqual(preset_values(META_SAMPLING), (True, 1.0, 0.95, 64))

    def test_greedy_kwargs_omit_inert_sampling_parameters(self):
        kwargs = generation_kwargs(
            do_sample=False,
            max_new_tokens=512,
            temperature=1.0,
            top_p=0.95,
            top_k=64,
            repetition_penalty=1.0,
        )
        self.assertNotIn("temperature", kwargs)
        self.assertNotIn("top_p", kwargs)
        self.assertNotIn("top_k", kwargs)
        self.assertEqual(kwargs["eos_token_id"], [200001, 200008])
        self.assertEqual(kwargs["pad_token_id"], 200018)

    def test_sampling_kwargs_match_model_card(self):
        kwargs = generation_kwargs(
            do_sample=True,
            max_new_tokens=512,
            temperature=1.0,
            top_p=0.95,
            top_k=64,
            repetition_penalty=1.0,
        )
        self.assertEqual(kwargs["temperature"], 1.0)
        self.assertEqual(kwargs["top_p"], 0.95)
        self.assertEqual(kwargs["top_k"], 64)


class ValidationTests(unittest.TestCase):
    def test_documented_controls_are_valid(self):
        validate_controls(
            max_new_tokens=512,
            temperature=1.0,
            top_p=0.95,
            top_k=64,
            repetition_penalty=1.0,
            reasoning_strength="high",
        )

    def test_invalid_reasoning_is_rejected(self):
        with self.assertRaises(ValueError):
            validate_controls(
                max_new_tokens=512,
                temperature=1.0,
                top_p=0.95,
                top_k=64,
                repetition_penalty=1.0,
                reasoning_strength="extreme",
            )

    def test_seed_resolution(self):
        self.assertEqual(choose_seed(42, False), 42)
        self.assertEqual(choose_seed(42, True, randbelow=lambda _limit: 123), 123)

    def test_duration_is_bounded_and_image_aware(self):
        text_duration = estimate_gpu_duration(512, False)
        image_duration = estimate_gpu_duration(512, True)
        self.assertGreater(image_duration, text_duration)
        self.assertGreaterEqual(estimate_gpu_duration(1, False), 60)
        self.assertLessEqual(estimate_gpu_duration(100_000, True), 240)


class ResponseTests(unittest.TestCase):
    def test_parser_fields_are_normalized(self):
        reply = coerce_parsed_reply(
            {"reasoning_content": " think ", "content": " answer ", "tool_calls": [{"x": 1}]}
        )
        self.assertEqual(reply.reasoning, "think")
        self.assertEqual(reply.content, "answer")
        self.assertEqual(reply.tool_calls, [{"x": 1}])

    def test_reasoning_renders_in_collapsible_region(self):
        rendered = render_reply("steps", "answer", show_reasoning=True)
        self.assertIn("<think>\nsteps\n</think>", rendered)
        self.assertTrue(rendered.endswith("answer"))

    def test_reasoning_can_be_hidden_without_hiding_answer(self):
        rendered = render_reply("secret chain", "answer", show_reasoning=False)
        self.assertNotIn("secret chain", rendered)
        self.assertEqual(rendered, "answer")

    def test_model_tags_cannot_break_reasoning_wrapper(self):
        rendered = render_reply("a </think> b", "ok", show_reasoning=True)
        self.assertIn("&lt;/think&gt;", rendered)
        self.assertEqual(rendered.count("</think>"), 1)

    def test_errors_do_not_echo_arbitrary_details(self):
        message = friendly_error(RuntimeError("secret prompt at /private/path"))
        self.assertNotIn("secret", message)
        self.assertNotIn("/private/path", message)


if __name__ == "__main__":
    unittest.main()