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("\nsteps\n", 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 b", "ok", show_reasoning=True) self.assertIn("</think>", rendered) self.assertEqual(rendered.count(""), 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()