test/tests/test_offline_model_selectio...

56 lines
2.0 KiB
Python

import unittest
import sys
import types
from unittest.mock import patch
fastapi_module = types.ModuleType("fastapi")
fastapi_module.Request = object
fastapi_responses_module = types.ModuleType("fastapi.responses")
fastapi_responses_module.JSONResponse = object
sys.modules.setdefault("fastapi", fastapi_module)
sys.modules.setdefault("fastapi.responses", fastapi_responses_module)
manager_module = types.ModuleType("app.services.asr.manager")
manager_module.get_model_manager = lambda: None
model_plan_module = types.ModuleType("app.services.asr.model_plan")
model_plan_module.get_active_qwen_model = lambda: "qwen3-asr-0.6b"
model_plan_module.get_default_model_id = lambda: "qwen3-asr-0.6b"
model_plan_module.get_runtime_model_ids = lambda: ["qwen3-asr-0.6b"]
sys.modules.setdefault("app.services.asr.manager", manager_module)
sys.modules.setdefault("app.services.asr.model_plan", model_plan_module)
from app.core.exceptions import InvalidParameterException
from app.services.asr.model_selection import validate_offline_model_id
class OfflineModelSelectionTest(unittest.TestCase):
def test_empty_model_uses_default_offline_model(self) -> None:
with (
patch(
"app.services.asr.model_selection.get_offline_model_ids",
return_value=["qwen3-asr-0.6b"],
),
patch(
"app.services.asr.model_selection.get_default_offline_model_id",
return_value="qwen3-asr-0.6b",
),
):
self.assertEqual(validate_offline_model_id(None), "qwen3-asr-0.6b")
self.assertEqual(validate_offline_model_id(""), "qwen3-asr-0.6b")
def test_rejects_unavailable_offline_model(self) -> None:
with patch(
"app.services.asr.model_selection.get_offline_model_ids",
return_value=["qwen3-asr-0.6b"],
):
with self.assertRaises(InvalidParameterException):
validate_offline_model_id("qwen3-asr-1.7b")
if __name__ == "__main__":
unittest.main()