返回 ViMax
test_provider_presets.py
根目录 / tests / test_provider_presets.py
1 """Unit tests for utils.provider_presets."""
2
3 import os
4 import unittest
5 from unittest.mock import patch
6
7 from utils.provider_presets import (
8 PROVIDER_PRESETS,
9 resolve_chat_model_config,
10 detect_provider_from_env,
11 )
12
13
14 class TestProviderPresets(unittest.TestCase):
15 """Tests for the PROVIDER_PRESETS registry."""
16
17 def test_minimax_preset_exists(self):
18 self.assertIn("minimax", PROVIDER_PRESETS)
19
20 def test_minimax_preset_base_url(self):
21 self.assertEqual(
22 PROVIDER_PRESETS["minimax"]["base_url"],
23 "https://api.minimax.io/v1",
24 )
25
26 def test_minimax_preset_env_key(self):
27 self.assertEqual(PROVIDER_PRESETS["minimax"]["env_key"], "MINIMAX_API_KEY")
28
29 def test_minimax_preset_default_model(self):
30 self.assertEqual(PROVIDER_PRESETS["minimax"]["default_model"], "MiniMax-M3")
31
32 def test_minimax_preset_has_models_list(self):
33 models = PROVIDER_PRESETS["minimax"]["models"]
34 self.assertIn("MiniMax-M3", models)
35 self.assertIn("MiniMax-M2.7", models)
36 self.assertIn("MiniMax-M2.7-highspeed", models)
37
38 def test_minimax_preset_temperature_range(self):
39 lo, hi = PROVIDER_PRESETS["minimax"]["temperature_range"]
40 self.assertEqual(lo, 0.0)
41 self.assertEqual(hi, 1.0)
42
43
44 class TestResolveChatModelConfig(unittest.TestCase):
45 """Tests for resolve_chat_model_config()."""
46
47 def test_unknown_provider_passes_through(self):
48 args = {"model_provider": "openai", "model": "gpt-4", "base_url": "https://example.com"}
49 result = resolve_chat_model_config(args)
50 self.assertEqual(result["model_provider"], "openai")
51 self.assertEqual(result["model"], "gpt-4")
52 self.assertEqual(result["base_url"], "https://example.com")
53
54 def test_no_model_provider_passes_through(self):
55 args = {"model": "gpt-4"}
56 result = resolve_chat_model_config(args)
57 self.assertEqual(result["model"], "gpt-4")
58
59 def test_minimax_rewrites_provider_to_openai(self):
60 args = {"model_provider": "minimax", "model": "MiniMax-M3", "api_key": "sk-test"}
61 result = resolve_chat_model_config(args)
62 self.assertEqual(result["model_provider"], "openai")
63
64 def test_minimax_sets_base_url(self):
65 args = {"model_provider": "minimax", "model": "MiniMax-M3", "api_key": "sk-test"}
66 result = resolve_chat_model_config(args)
67 self.assertEqual(result["base_url"], "https://api.minimax.io/v1")
68
69 def test_minimax_preserves_custom_base_url(self):
70 args = {
71 "model_provider": "minimax",
72 "model": "MiniMax-M3",
73 "api_key": "sk-test",
74 "base_url": "https://custom-proxy.example.com/v1",
75 }
76 result = resolve_chat_model_config(args)
77 self.assertEqual(result["base_url"], "https://custom-proxy.example.com/v1")
78
79 def test_minimax_defaults_model(self):
80 args = {"model_provider": "minimax", "api_key": "sk-test"}
81 result = resolve_chat_model_config(args)
82 self.assertEqual(result["model"], "MiniMax-M3")
83
84 def test_minimax_preserves_explicit_model(self):
85 args = {"model_provider": "minimax", "model": "MiniMax-M2.7-highspeed", "api_key": "sk-test"}
86 result = resolve_chat_model_config(args)
87 self.assertEqual(result["model"], "MiniMax-M2.7-highspeed")
88
89 @patch.dict(os.environ, {"MINIMAX_API_KEY": "env-key-123"})
90 def test_minimax_reads_api_key_from_env(self):
91 args = {"model_provider": "minimax", "model": "MiniMax-M3"}
92 result = resolve_chat_model_config(args)
93 self.assertEqual(result["api_key"], "env-key-123")
94
95 def test_minimax_prefers_explicit_api_key_over_env(self):
96 args = {"model_provider": "minimax", "model": "MiniMax-M3", "api_key": "explicit-key"}
97 with patch.dict(os.environ, {"MINIMAX_API_KEY": "env-key"}):
98 result = resolve_chat_model_config(args)
99 self.assertEqual(result["api_key"], "explicit-key")
100
101 def test_minimax_clamps_temperature_above_max(self):
102 args = {"model_provider": "minimax", "model": "MiniMax-M3", "api_key": "sk", "temperature": 1.5}
103 result = resolve_chat_model_config(args)
104 self.assertEqual(result["temperature"], 1.0)
105
106 def test_minimax_clamps_temperature_below_min(self):
107 args = {"model_provider": "minimax", "model": "MiniMax-M3", "api_key": "sk", "temperature": -0.5}
108 result = resolve_chat_model_config(args)
109 self.assertEqual(result["temperature"], 0.0)
110
111 def test_minimax_passes_valid_temperature(self):
112 args = {"model_provider": "minimax", "model": "MiniMax-M3", "api_key": "sk", "temperature": 0.7}
113 result = resolve_chat_model_config(args)
114 self.assertEqual(result["temperature"], 0.7)
115
116 def test_minimax_temperature_zero_allowed(self):
117 args = {"model_provider": "minimax", "model": "MiniMax-M3", "api_key": "sk", "temperature": 0.0}
118 result = resolve_chat_model_config(args)
119 self.assertEqual(result["temperature"], 0.0)
120
121 def test_minimax_no_temperature_key(self):
122 args = {"model_provider": "minimax", "model": "MiniMax-M3", "api_key": "sk"}
123 result = resolve_chat_model_config(args)
124 self.assertNotIn("temperature", result)
125
126 def test_minimax_temperature_none_ignored(self):
127 args = {"model_provider": "minimax", "model": "MiniMax-M3", "api_key": "sk", "temperature": None}
128 result = resolve_chat_model_config(args)
129 self.assertIsNone(result["temperature"])
130
131 def test_original_dict_not_mutated(self):
132 args = {"model_provider": "minimax", "model": "MiniMax-M3", "api_key": "sk"}
133 resolve_chat_model_config(args)
134 self.assertEqual(args["model_provider"], "minimax")
135
136 def test_empty_model_string_gets_default(self):
137 args = {"model_provider": "minimax", "model": "", "api_key": "sk"}
138 result = resolve_chat_model_config(args)
139 self.assertEqual(result["model"], "MiniMax-M3")
140
141
142 class TestDetectProviderFromEnv(unittest.TestCase):
143 """Tests for detect_provider_from_env()."""
144
145 @patch.dict(os.environ, {"MINIMAX_API_KEY": "test-key"}, clear=False)
146 def test_detects_minimax(self):
147 self.assertEqual(detect_provider_from_env(), "minimax")
148
149 @patch.dict(os.environ, {}, clear=True)
150 def test_returns_none_when_no_keys(self):
151 self.assertIsNone(detect_provider_from_env())
152
153
154 class TestConfigYAMLLoading(unittest.TestCase):
155 """Test that MiniMax example config files are valid YAML."""
156
157 def test_idea2video_minimax_yaml(self):
158 import yaml
159 path = os.path.join(os.path.dirname(__file__), "..", "configs", "idea2video_minimax.yaml")
160 with open(path) as f:
161 config = yaml.safe_load(f)
162 self.assertEqual(config["chat_model"]["init_args"]["model_provider"], "minimax")
163 self.assertEqual(config["chat_model"]["init_args"]["model"], "MiniMax-M3")
164
165 def test_script2video_minimax_yaml(self):
166 import yaml
167 path = os.path.join(os.path.dirname(__file__), "..", "configs", "script2video_minimax.yaml")
168 with open(path) as f:
169 config = yaml.safe_load(f)
170 self.assertEqual(config["chat_model"]["init_args"]["model_provider"], "minimax")
171 self.assertEqual(config["chat_model"]["init_args"]["model"], "MiniMax-M3")
172
173
174 if __name__ == "__main__":
175 unittest.main()
176
176 lines PYTHON