返回 CodeWhale
configured_models.rs
根目录 / crates / config / tests / configured_models.rs
1 use codewhale_config::catalog::configured::{ConfiguredModel, validate_configured_models};
2 use codewhale_config::route::{
3 CapabilityState, LogicalModelRef, OverrideSource, RouteRequest, RouteResolver,
4 };
5 use codewhale_config::{ConfigStore, ConfigToml, ProviderKind};
6
7 const FIXTURE: &str = include_str!("fixtures/custom_models.toml");
8 const ID: &str = "deepseek-v4.1-flash-expires-on-0910";
9 const BASE: &str = "https://models.example.test/v1";
10
11 fn models() -> Vec<ConfiguredModel> {
12 toml::from_str::<ConfigToml>(FIXTURE)
13 .unwrap()
14 .custom_models
15 .unwrap()
16 }
17
18 fn request(base: &str, id: &str) -> RouteRequest {
19 RouteRequest {
20 explicit_provider: Some(ProviderKind::Deepseek),
21 model_selector: Some(LogicalModelRef::from(id)),
22 base_url_override: Some(base.into()),
23 ..RouteRequest::default()
24 }
25 }
26
27 #[test]
28 fn persisted_model_roundtrip_reload_drives_immutable_route() {
29 let dir = tempfile::tempdir().unwrap();
30 let path = dir.path().join("config.toml");
31 std::fs::write(&path, FIXTURE).unwrap();
32 let mut store = ConfigStore::load(Some(path.clone())).unwrap();
33 let original = store.config.custom_models.clone();
34 store.config.telemetry = Some(true);
35 store.save().unwrap();
36 let mut reloaded = ConfigStore::load(Some(path)).unwrap();
37 assert_eq!(reloaded.config.custom_models, original);
38 let definitions = reloaded.config.custom_models.as_ref().unwrap();
39 assert_eq!(
40 definitions[0].display_name.as_deref(),
41 Some("Temporary preview")
42 );
43 assert_eq!(
44 definitions[0].extras["future_note"].as_str(),
45 Some("preserve this metadata")
46 );
47 let resolver = RouteResolver::new().with_configured_models(
48 definitions,
49 "deepseek",
50 ProviderKind::Deepseek,
51 BASE,
52 );
53 let old = resolver.resolve(&request(BASE, ID)).unwrap();
54 assert_eq!(old.wire_model_id().as_str(), ID);
55 assert_eq!(old.limits().context_tokens, Some(96000));
56 assert_eq!(old.limits().input_tokens, Some(88000));
57 assert_eq!(old.limits().output_tokens, Some(8000));
58 assert_eq!(old.capabilities().image_input, CapabilityState::Unknown);
59 assert_eq!(
60 old.capabilities().native_tool_calls,
61 CapabilityState::Unknown
62 );
63 assert_eq!(
64 old.capabilities().structured_output,
65 CapabilityState::Unsupported
66 );
67 assert!(
68 old.applied_limit_overrides()
69 .iter()
70 .all(|entry| entry.source == OverrideSource::UserModelMetadata)
71 );
72 reloaded.config.custom_models.as_mut().unwrap()[0]
73 .limit
74 .as_mut()
75 .unwrap()
76 .context = Some(128000);
77 reloaded.save().unwrap();
78 reloaded.reload().unwrap();
79 let next = RouteResolver::new()
80 .with_configured_models(
81 reloaded.config.custom_models.as_ref().unwrap(),
82 "deepseek",
83 ProviderKind::Deepseek,
84 BASE,
85 )
86 .resolve(&request(BASE, ID))
87 .unwrap();
88 assert_eq!(next.limits().context_tokens, Some(128000));
89 assert_eq!(old.limits().context_tokens, Some(96000));
90 }
91
92 #[test]
93 fn exact_identity_and_endpoint_do_not_leak_declarations() {
94 let resolver = RouteResolver::new().with_configured_models(
95 &models(),
96 "deepseek",
97 ProviderKind::Deepseek,
98 BASE,
99 );
100 for base in [
101 "https://other.example.test/v1",
102 "http://models.example.test/v1",
103 "https://models.example.test:444/v1",
104 "https://models.example.test/V1",
105 ] {
106 let route = resolver.resolve(&request(base, ID)).unwrap();
107 assert_eq!(route.limits().context_tokens, None, "{base}");
108 }
109 for id in ["deepseek-v4.1-flash", "DEEPSEEK-V4.1-FLASH-EXPIRES-ON-0910"] {
110 assert_eq!(
111 resolver
112 .resolve(&request(BASE, id))
113 .unwrap()
114 .limits()
115 .context_tokens,
116 None,
117 "{id}"
118 );
119 }
120 assert_eq!(
121 resolver
122 .resolve(&request("https://MODELS.example.test:443/v1/", ID))
123 .unwrap()
124 .limits()
125 .context_tokens,
126 Some(96000)
127 );
128 let wrong_identity = RouteResolver::new().with_configured_models(
129 &models(),
130 "another-provider",
131 ProviderKind::Deepseek,
132 BASE,
133 );
134 assert_eq!(
135 wrong_identity
136 .resolve(&request(BASE, ID))
137 .unwrap()
138 .limits()
139 .context_tokens,
140 None
141 );
142 }
143
144 #[test]
145 fn unknown_fields_are_not_filled_from_a_known_sibling() {
146 let mut definitions = models();
147 definitions[0].limit = None;
148 definitions[0].cost = None;
149 definitions[0].reasoning = None;
150 definitions[0].modalities = None;
151 definitions[0].tool_call = None;
152 let offering = definitions[0].to_catalog_offering();
153 assert!(offering.limit.is_none());
154 assert!(offering.cost.is_none());
155 assert!(offering.reasoning.is_none());
156 let route = RouteResolver::new()
157 .with_configured_models(&definitions, "deepseek", ProviderKind::Deepseek, BASE)
158 .resolve(&request(BASE, ID))
159 .unwrap();
160 assert_eq!(route.limits().context_tokens, None);
161 assert_eq!(route.capabilities().reasoning, CapabilityState::Unknown);
162 assert_eq!(
163 route.capabilities().native_tool_calls,
164 CapabilityState::Unknown
165 );
166 }
167
168 #[test]
169 fn declarations_cannot_expand_closed_protocol_rosters() {
170 let mut definitions = models();
171 definitions[0].provider = "opencode-zen".into();
172 definitions[0].id = "unknown-protocol-model".into();
173 let resolver = RouteResolver::new().with_configured_models(
174 &definitions,
175 "opencode-zen",
176 ProviderKind::OpencodeZen,
177 BASE,
178 );
179 let mut req = request(BASE, "unknown-protocol-model");
180 req.explicit_provider = Some(ProviderKind::OpencodeZen);
181 assert!(resolver.resolve(&req).is_err());
182 }
183
184 #[test]
185 fn invalid_limits_prices_units_and_authority_fail_closed() {
186 for (before, after) in [
187 ("context = 96000", "context = 0"),
188 ("context = 96000", "context = 4294967296"),
189 ("output = 8000", "output = 97000"),
190 ("input = 0.4", "input = -1.0"),
191 ("input = 0.4", "input = nan"),
192 ("input = 0.4", "input = inf"),
193 ("input = 0.4", "input = 0.4, currency = 'CNY'"),
194 ("input = 0.4", "input = 0.4, unit = 'per_token'"),
195 ("future_note =", "source ="),
196 ("future_note =", "canonical_model ="),
197 ("future_note =", "api_key ="),
198 (ID, "auto"),
199 (BASE, "https://user:password@models.example.test/v1"),
200 (BASE, "https://models.example.test/v1?key=secret"),
201 ] {
202 assert!(
203 toml::from_str::<ConfigToml>(&FIXTURE.replace(before, after)).is_err(),
204 "{before} -> {after}"
205 );
206 }
207 let mut duplicate = models();
208 let mut other = duplicate[0].clone();
209 other.base_url = "https://MODELS.example.test:443/v1/".into();
210 duplicate.push(other);
211 assert!(validate_configured_models(&duplicate).is_err());
212 }
213
214 #[test]
215 fn declared_wire_ids_are_not_convenience_aliases() {
216 for (kind, provider, base, id) in [
217 (
218 ProviderKind::Together,
219 "together",
220 "https://api.together.xyz/v1",
221 "inkling",
222 ),
223 (
224 ProviderKind::Openrouter,
225 "openrouter",
226 "https://openrouter.ai/api/v1",
227 "qwen3.7-plus",
228 ),
229 (
230 ProviderKind::Concentrate,
231 "concentrate",
232 "https://api.concentrate.ai/v1",
233 "concentrate/example",
234 ),
235 (
236 ProviderKind::Deepseek,
237 "deepseek",
238 "https://api.deepseek.com",
239 "deepseek-v4pro",
240 ),
241 ] {
242 let mut definitions = models();
243 definitions[0].provider = provider.into();
244 definitions[0].base_url = base.into();
245 definitions[0].id = id.into();
246 let resolver =
247 RouteResolver::new().with_configured_models(&definitions, provider, kind, base);
248 let mut req = request(base, id);
249 req.explicit_provider = Some(kind);
250 let candidate = resolver.resolve(&req).unwrap();
251 assert_eq!(candidate.wire_model_id().as_str(), id);
252 assert_eq!(candidate.limits().output_tokens, Some(8000));
253 req.base_url_override = Some("https://elsewhere.example.test/v1".into());
254 assert_eq!(resolver.resolve(&req).unwrap().limits().output_tokens, None);
255 }
256 }
257
258 #[test]
259 fn persisted_metadata_survives_unrelated_save() {
260 let dir = tempfile::tempdir().unwrap();
261 let path = dir.path().join("config.toml");
262 std::fs::write(&path, FIXTURE).unwrap();
263 let mut store = ConfigStore::load(Some(path.clone())).unwrap();
264 store.config.telemetry = Some(true);
265 store.save().unwrap();
266 let persisted: toml::Value = toml::from_str(&std::fs::read_to_string(path).unwrap()).unwrap();
267 assert_eq!(
268 persisted
269 .get("custom_models")
270 .and_then(|models| models.as_array())
271 .and_then(|models| models.first())
272 .and_then(|model| model.get("id"))
273 .and_then(|id| id.as_str()),
274 Some(ID)
275 );
276 }
277
278 #[test]
279 fn declared_wire_id_precedes_real_aggregator_canonical_aliases() {
280 let id = "deepseek-v4-pro";
281 for (provider, kind, base) in [
282 (
283 "openrouter",
284 ProviderKind::Openrouter,
285 "https://openrouter.ai/api/v1",
286 ),
287 (
288 "together",
289 ProviderKind::Together,
290 "https://api.together.xyz/v1",
291 ),
292 ] {
293 let mut definitions = models();
294 definitions[0].provider = provider.into();
295 definitions[0].base_url = base.into();
296 definitions[0].id = id.into();
297 let mut req = request(base, id);
298 req.explicit_provider = Some(kind);
299 let bundled = RouteResolver::new().resolve(&req).unwrap();
300 assert_ne!(
301 bundled.wire_model_id().as_str(),
302 id,
303 "fixture must hit real bundled alias"
304 );
305 let resolver =
306 RouteResolver::new().with_configured_models(&definitions, provider, kind, base);
307 for saved in [false, true] {
308 if saved {
309 req.model_selector = None;
310 req.saved_provider_model = Some(id.into());
311 }
312 let candidate = resolver.resolve(&req).unwrap();
313 assert_eq!(
314 candidate.wire_model_id().as_str(),
315 id,
316 "{provider} saved={saved}"
317 );
318 assert!(candidate.canonical_model().is_none());
319 assert_eq!(
320 candidate.pricing(),
321 Some(&codewhale_config::route::PricingSku::Token {
322 input_per_mtok: Some(0.4),
323 output_per_mtok: Some(1.6),
324 })
325 );
326 assert_eq!(candidate.limits().context_tokens, Some(96000));
327 assert_eq!(candidate.limits().output_tokens, Some(8000));
328 assert!(
329 candidate
330 .applied_limit_overrides()
331 .iter()
332 .all(|entry| entry.source == OverrideSource::UserModelMetadata)
333 );
334 }
335 req.base_url_override = Some("https://unrelated.example.test/v1".into());
336 let unrelated = resolver.resolve(&req).unwrap();
337 assert_eq!(unrelated.limits().context_tokens, None);
338 assert!(unrelated.applied_limit_overrides().is_empty());
339 }
340 }
341
342 #[test]
343 fn all_positive_declared_capabilities_stay_unverified() {
344 let mut definitions = models();
345 for flag in [Some(true), Some(false), None] {
346 let model = &mut definitions[0];
347 model.reasoning = flag;
348 model.tool_call = flag;
349 model.attachment = flag;
350 model.structured_output = flag;
351 model.modalities =
352 flag.map(
353 |supported| codewhale_config::models_dev::ModelsDevModalities {
354 input: if supported {
355 vec!["text".into(), "image".into()]
356 } else {
357 vec!["text".into()]
358 },
359 output: vec!["text".into()],
360 },
361 );
362 assert_eq!(model.to_catalog_offering().reasoning, flag);
363 let candidate = RouteResolver::new()
364 .with_configured_models(&definitions, "deepseek", ProviderKind::Deepseek, BASE)
365 .resolve(&request(BASE, ID))
366 .unwrap();
367 let caps = candidate.capabilities();
368 let expected = if flag == Some(false) {
369 CapabilityState::Unsupported
370 } else {
371 CapabilityState::Unknown
372 };
373 for capability in [
374 caps.reasoning,
375 caps.image_input,
376 caps.attachments,
377 caps.native_tool_calls,
378 caps.structured_output,
379 ] {
380 assert_eq!(capability, expected);
381 }
382 }
383 }
384
384 lines RUST