| import os |
| import tempfile |
| from pathlib import Path |
| import unittest |
|
|
|
|
| os.environ.setdefault("KMP_DUPLICATE_LIB_OK", "TRUE") |
| os.environ["MUSE_SKIP_MODEL_LOAD"] = "1" |
|
|
| from PIL import Image |
|
|
| import app |
|
|
|
|
| class GradioContractTests(unittest.TestCase): |
| def test_click_and_enter_have_validated_serial_generators(self): |
| dependencies = app.demo.get_config_file()["dependencies"] |
| generators = [dependency for dependency in dependencies if dependency["types"]["generator"]] |
| self.assertGreaterEqual(len(generators), 2) |
|
|
| targets = set() |
| api_generator = None |
| for dependency in generators: |
| targets.update(target[1] for target in dependency["targets"]) |
| self.assertTrue(dependency["queue"]) |
| self.assertEqual(dependency.get("api_visibility", "public"), "private") |
| self.assertEqual(len(dependency["outputs"]), 6) |
| self.assertEqual(len(dependency["inputs"]), 16) |
| function = app.demo.fns[dependency["id"]] |
| self.assertEqual(function.concurrency_id, "muse-glimmer-xlarge") |
| self.assertEqual(function.concurrency_limit, 1) |
| self.assertIsNotNone(function.validator) |
| if dependency["api_name"] == "chat": |
| api_generator = dependency |
| if dependency["api_name"] == app.SUBMIT_API_NAME: |
| self.assertEqual(dependency["types"]["generator"], True) |
|
|
| self.assertEqual(targets, {"click", "submit"}) |
| self.assertIsNotNone(api_generator, "Expected one private API-visible generation path named chat.") |
|
|
| generator_functions = [app.demo.fns[dependency["id"]] for dependency in generators] |
| self.assertEqual( |
| {function.concurrency_id for function in generator_functions}, |
| {"muse-glimmer-xlarge"}, |
| ) |
| self.assertTrue(all(function.concurrency_limit == 1 for function in generator_functions)) |
| self.assertTrue(all(function.validator is not None for function in generator_functions)) |
|
|
| api_names = set() |
| for dependency in generators: |
| api_names.add(dependency["api_name"]) |
| self.assertEqual(api_names, {"chat", app.SUBMIT_API_NAME}) |
|
|
| def test_stop_and_clear_cancel_click_and_enter_generation(self): |
| dependencies = app.demo.get_config_file()["dependencies"] |
| generator_ids = { |
| dependency["id"] for dependency in dependencies if dependency["types"]["generator"] |
| } |
| cancellation_edges = [ |
| set(dependency["cancels"]) |
| for dependency in dependencies |
| if dependency["types"]["cancel"] |
| ] |
|
|
| self.assertEqual(len(cancellation_edges), 2) |
| self.assertTrue(all(edge == generator_ids for edge in cancellation_edges)) |
|
|
| def test_invalid_requests_fail_the_queue_free_validator(self): |
| with tempfile.TemporaryDirectory() as tmpdir: |
| tmp_path = Path(tmpdir) |
| original_paths = { |
| model_id: spec["path"] |
| for model_id, spec in app.MODEL_REGISTRY.items() |
| } |
| for spec in app.MODEL_REGISTRY.values(): |
| spec["path"] = tmp_path |
|
|
| try: |
| verdicts = app._validate_generation_request( |
| "", |
| None, |
| app.MODEL_DEFAULT_ID, |
| [], |
| [], |
| "", |
| "high", |
| False, |
| 512, |
| 1.0, |
| 0.95, |
| 64, |
| 1.0, |
| 42, |
| False, |
| True, |
| ) |
| finally: |
| for model_id, original_path in original_paths.items(): |
| app.MODEL_REGISTRY[model_id]["path"] = original_path |
|
|
| self.assertEqual(len(verdicts), 16) |
| self.assertFalse(verdicts[0]["is_valid"]) |
| self.assertIn("prompt", verdicts[0]["message"].lower()) |
|
|
| def test_model_id_payload_formats_from_ui_are_accepted(self): |
| with tempfile.TemporaryDirectory() as tmpdir: |
| tmp_path = Path(tmpdir) |
| original_paths = { |
| model_id: spec["path"] |
| for model_id, spec in app.MODEL_REGISTRY.items() |
| } |
| for spec in app.MODEL_REGISTRY.values(): |
| spec["path"] = tmp_path |
|
|
| try: |
| verdict = app._validate_generation_request( |
| "hi", |
| None, |
| ("Muse Glimmer 30B-assistant", app.ASSISTANT_MODEL_ID), |
| [], |
| [], |
| "", |
| "high", |
| False, |
| 512, |
| 1.0, |
| 0.95, |
| 64, |
| 1.0, |
| 42, |
| False, |
| True, |
| )[0] |
| finally: |
| for model_id, original_path in original_paths.items(): |
| app.MODEL_REGISTRY[model_id]["path"] = original_path |
|
|
| self.assertTrue(verdict["is_valid"]) |
|
|
| def test_legacy_model_labels_are_accepted_by_coercion(self): |
| self.assertEqual(app._coerce_model_id("/Muse-Glimmer 30B"), app.MODEL_ID) |
| self.assertEqual(app._coerce_model_id("/Muse-Glimmer-30B"), app.MODEL_ID) |
| self.assertEqual(app._coerce_model_id("Muse Glimmer 30B (full BF16)"), app.MODEL_ID) |
| self.assertEqual(app._coerce_model_id("Muse Glimmer 30B-assistant (compact)"), app.ASSISTANT_MODEL_ID) |
| self.assertEqual(app._coerce_model_id(""), app.MODEL_ID) |
|
|
| def test_validator_defaults_when_controls_are_missing(self): |
| with tempfile.TemporaryDirectory() as tmpdir: |
| tmp_path = Path(tmpdir) |
| original_paths = { |
| model_id: spec["path"] |
| for model_id, spec in app.MODEL_REGISTRY.items() |
| } |
| for spec in app.MODEL_REGISTRY.values(): |
| spec["path"] = tmp_path |
|
|
| try: |
| verdict = app._validate_generation_request( |
| "hi", |
| None, |
| app.MODEL_DEFAULT_ID, |
| [], |
| [], |
| "", |
| "high", |
| False, |
| None, |
| None, |
| None, |
| None, |
| None, |
| 42, |
| False, |
| True, |
| )[0] |
| finally: |
| for model_id, original_path in original_paths.items(): |
| app.MODEL_REGISTRY[model_id]["path"] = original_path |
|
|
| self.assertTrue(verdict["is_valid"]) |
|
|
| def test_numeric_model_id_payload_formats_from_ui_are_accepted(self): |
| with tempfile.TemporaryDirectory() as tmpdir: |
| tmp_path = Path(tmpdir) |
| original_paths = { |
| model_id: spec["path"] |
| for model_id, spec in app.MODEL_REGISTRY.items() |
| } |
| for spec in app.MODEL_REGISTRY.values(): |
| spec["path"] = tmp_path |
|
|
| try: |
| verdict = app._validate_generation_request( |
| "hi", |
| None, |
| "1", |
| [], |
| [], |
| "", |
| "high", |
| False, |
| 512, |
| 1.0, |
| 0.95, |
| 64, |
| 1.0, |
| 42, |
| False, |
| True, |
| )[0] |
| finally: |
| for model_id, original_path in original_paths.items(): |
| app.MODEL_REGISTRY[model_id]["path"] = original_path |
|
|
| self.assertTrue(verdict["is_valid"]) |
|
|
|
|
| class StateContractTests(unittest.TestCase): |
| def test_only_two_recent_images_remain_in_model_history(self): |
| image = Image.new("RGB", (8, 8), "purple") |
| history = [] |
| for index in range(3): |
| history.extend( |
| [ |
| { |
| "role": "user", |
| "content": [ |
| {"type": "image", "image": image.copy()}, |
| {"type": "text", "text": f"turn {index}"}, |
| ], |
| }, |
| {"role": "assistant", "content": f"answer {index}"}, |
| ] |
| ) |
|
|
| cleaned = app._clean_model_history(history) |
| image_turns = [ |
| message |
| for message in cleaned |
| if isinstance(message.get("content"), list) |
| and any(part.get("type") == "image" for part in message["content"]) |
| ] |
| self.assertEqual(len(image_turns), 2) |
| self.assertEqual(cleaned[0]["content"], "turn 0") |
|
|
| def test_stop_rolls_back_visible_chat_snapshot(self): |
| snapshot = [{"role": "user", "content": "committed"}] |
| chat, status = app._stop_conversation(snapshot) |
|
|
| self.assertEqual(chat, snapshot) |
| self.assertIsNot(chat, snapshot) |
| self.assertIn("not added", status) |
|
|
|
|
| class DurationContractTests(unittest.TestCase): |
| def setUp(self): |
| app.ACTIVE_MODEL_ID = None |
|
|
| def tearDown(self): |
| app.ACTIVE_MODEL_ID = app.MODEL_ID |
|
|
| def test_full_model_cold_start_requests_cold_start_budget(self): |
| duration = app._gpu_duration( |
| "hi", |
| None, |
| app.MODEL_ID, |
| [], |
| [], |
| "", |
| "high", |
| False, |
| 512, |
| 1.0, |
| 0.95, |
| 64, |
| 1.0, |
| 42, |
| False, |
| True, |
| ) |
| self.assertEqual(duration, 120) |
|
|
| def test_full_model_reuse_keeps_estimate_after_load(self): |
| app.ACTIVE_MODEL_ID = app.MODEL_ID |
| duration = app._gpu_duration( |
| "hi", |
| None, |
| app.MODEL_ID, |
| [], |
| [], |
| "", |
| "high", |
| False, |
| 512, |
| 1.0, |
| 0.95, |
| 64, |
| 1.0, |
| 42, |
| False, |
| True, |
| ) |
| self.assertEqual(duration, 121) |
|
|
| def test_assistant_selection_reuses_loaded_full_model(self): |
| app.ACTIVE_MODEL_ID = app.MODEL_ID |
| duration = app._gpu_duration( |
| "hi", |
| None, |
| app.ASSISTANT_MODEL_ID, |
| [], |
| [], |
| "", |
| "high", |
| False, |
| 512, |
| 1.0, |
| 0.95, |
| 64, |
| 1.0, |
| 42, |
| False, |
| True, |
| ) |
| self.assertEqual(duration, 121) |
|
|
|
|
| class RuntimeChatObjectTests(unittest.TestCase): |
| def test_coerce_chat_objects_prefers_nested_tokenizer(self): |
| class FakeResponseToken: |
| def __init__(self, is_processor=True): |
| self.is_processor = is_processor |
|
|
| def apply_chat_template(self, *_args, **_kwargs): |
| return {"input_ids": []} |
|
|
| def get_response_parser(self, *_args, **_kwargs): |
| return "parser" |
|
|
| class FakeProcessor: |
| def __init__(self): |
| self.tokenizer = FakeResponseToken() |
|
|
| original_paths = {model_id: spec["path"] for model_id, spec in app.MODEL_REGISTRY.items()} |
| with tempfile.TemporaryDirectory() as tmpdir: |
| tmp_path = Path(tmpdir) |
| for spec in app.MODEL_REGISTRY.values(): |
| spec["path"] = tmp_path |
| try: |
| processor, tokenizer = app._coerce_chat_objects(FakeProcessor(), app.MODEL_ID) |
| finally: |
| for model_id, original_path in original_paths.items(): |
| app.MODEL_REGISTRY[model_id]["path"] = original_path |
|
|
| self.assertIsInstance(processor, FakeResponseToken) |
| self.assertIs(processor, tokenizer) |
|
|
|
|
| class ModelPathResolutionTests(unittest.TestCase): |
| def test_mount_path_resolution_prefers_nested_checkpoint(self): |
| with tempfile.TemporaryDirectory() as tmpdir: |
| tmp = Path(tmpdir) |
| nested = tmp / "Muse-Glimmer-30B" |
| nested.mkdir() |
| (nested / "config.json").write_text("{}") |
| (nested / "chat_template.jinja").write_text("") |
| resolved = app._resolve_mount_path(tmp) |
|
|
| self.assertEqual(resolved.name, "Muse-Glimmer-30B") |
|
|
|
|
| def test_mount_path_resolution_keeps_direct_checkpoint(self): |
| with tempfile.TemporaryDirectory() as tmpdir: |
| tmp = Path(tmpdir) |
| (tmp / "config.json").write_text("{}") |
| (tmp / "chat_template.jinja").write_text("") |
| resolved = app._resolve_mount_path(tmp) |
|
|
| self.assertEqual(resolved, tmp) |
|
|
|
|
| if __name__ == "__main__": |
| unittest.main() |
|
|