1use std::cmp::Ordering;
9use std::collections::HashMap;
10
11use globset::GlobBuilder;
12use pi_ai::types::{Model, ModelThinkingLevel};
13
14use super::agent_session_services::{
15 DEFAULT_THINKING_LEVEL, FindInitialModelOptions, InitialModelResult, ScopedModel,
16 default_model_per_provider, find_initial_model,
17};
18use super::model_runtime::ModelRuntime;
19
20pub use super::agent_session_services::restore_model_from_session;
22
23#[derive(Clone, Debug, PartialEq)]
25pub struct ParsedModelResult {
26 pub model: Option<Model>,
28 pub thinking_level: Option<ModelThinkingLevel>,
30 pub warning: Option<String>,
32}
33
34#[derive(Clone, Copy, Debug, Default)]
36pub struct ParseModelPatternOptions {
37 pub allow_invalid_thinking_level_fallback: Option<bool>,
40}
41
42#[derive(Clone, Debug, Eq, PartialEq)]
44pub struct ModelScopeDiagnostic {
45 pub kind: ModelScopeDiagnosticKind,
47 pub message: String,
49 pub pattern: String,
51}
52
53#[derive(Clone, Copy, Debug, Eq, PartialEq)]
55pub enum ModelScopeDiagnosticKind {
56 Warning,
58}
59
60#[derive(Clone, Debug, PartialEq)]
62pub struct ResolveModelScopeResult {
63 pub scoped_models: Vec<ScopedModel>,
65 pub diagnostics: Vec<ModelScopeDiagnostic>,
67}
68
69#[derive(Clone, Debug, PartialEq)]
71pub struct ResolveCliModelResult {
72 pub model: Option<Model>,
74 pub thinking_level: Option<ModelThinkingLevel>,
76 pub warning: Option<String>,
78 pub error: Option<String>,
80}
81
82#[derive(Clone, Copy, Debug)]
84pub struct ResolveCliModelOptions<'a> {
85 pub cli_provider: Option<&'a str>,
87 pub cli_model: Option<&'a str>,
89 pub cli_thinking: Option<ModelThinkingLevel>,
91 pub model_runtime: &'a ModelRuntime,
93}
94
95#[derive(Clone, Copy, Debug)]
97pub struct FindInitialModelFullOptions<'a> {
98 pub cli_provider: Option<&'a str>,
100 pub cli_model: Option<&'a str>,
102 pub scoped_models: &'a [ScopedModel],
104 pub is_continuing: bool,
106 pub default_provider: Option<&'a str>,
108 pub default_model_id: Option<&'a str>,
110 pub default_thinking_level: Option<ModelThinkingLevel>,
112 pub model_runtime: &'a ModelRuntime,
114}
115
116#[must_use]
118fn is_alias(id: &str) -> bool {
119 if id.ends_with("-latest") {
120 return true;
121 }
122 if id.len() >= 9 {
124 let bytes = id.as_bytes();
125 let start = bytes.len() - 9;
126 if bytes[start] == b'-' {
127 return !bytes[start + 1..].iter().all(u8::is_ascii_digit);
128 }
129 }
130 true
131}
132
133#[must_use]
135fn models_are_equal(a: &Model, b: &Model) -> bool {
136 a.id == b.id && a.provider == b.provider
137}
138
139#[must_use]
141pub fn is_valid_thinking_level(level: &str) -> bool {
142 parse_thinking_level(level).is_some()
143}
144
145#[must_use]
146fn parse_thinking_level(level: &str) -> Option<ModelThinkingLevel> {
147 match level {
148 "off" => Some(ModelThinkingLevel::Off),
149 "minimal" => Some(ModelThinkingLevel::Minimal),
150 "low" => Some(ModelThinkingLevel::Low),
151 "medium" => Some(ModelThinkingLevel::Medium),
152 "high" => Some(ModelThinkingLevel::High),
153 "xhigh" => Some(ModelThinkingLevel::Xhigh),
154 "max" => Some(ModelThinkingLevel::Max),
155 _ => None,
156 }
157}
158
159#[must_use]
164pub fn find_exact_model_reference_match(
165 model_reference: &str,
166 available_models: &[Model],
167) -> Option<Model> {
168 let trimmed = model_reference.trim();
169 if trimmed.is_empty() {
170 return None;
171 }
172 let normalized = trimmed.to_ascii_lowercase();
173
174 let mut canonical_matches = available_models.iter().filter(|model| {
175 format!("{}/{}", model.provider, model.id).to_ascii_lowercase() == normalized
176 });
177 match (canonical_matches.next(), canonical_matches.next()) {
178 (Some(only), None) => return Some(only.clone()),
179 (Some(_), Some(_)) => return None,
180 (None, _) => {}
181 }
182
183 if let Some(slash_index) = trimmed.find('/') {
184 let provider = trimmed[..slash_index].trim();
185 let model_id = trimmed[slash_index + 1..].trim();
186 if !provider.is_empty() && !model_id.is_empty() {
187 let mut provider_matches = available_models.iter().filter(|model| {
188 model.provider.eq_ignore_ascii_case(provider)
189 && model.id.eq_ignore_ascii_case(model_id)
190 });
191 match (provider_matches.next(), provider_matches.next()) {
192 (Some(only), None) => return Some(only.clone()),
193 (Some(_), Some(_)) => return None,
194 (None, _) => {}
195 }
196 }
197 }
198
199 let mut id_matches = available_models
200 .iter()
201 .filter(|model| model.id.eq_ignore_ascii_case(&normalized));
202 match (id_matches.next(), id_matches.next()) {
203 (Some(only), None) => Some(only.clone()),
204 _ => None,
205 }
206}
207
208#[must_use]
210fn try_match_model(model_pattern: &str, available_models: &[Model]) -> Option<Model> {
211 if let Some(exact) = find_exact_model_reference_match(model_pattern, available_models) {
212 return Some(exact);
213 }
214
215 let needle = model_pattern.to_ascii_lowercase();
216 let mut matches: Vec<&Model> = available_models
217 .iter()
218 .filter(|model| {
219 model.id.to_ascii_lowercase().contains(&needle)
220 || model.name.to_ascii_lowercase().contains(&needle)
221 })
222 .collect();
223
224 if matches.is_empty() {
225 return None;
226 }
227
228 let mut aliases: Vec<&Model> = matches
229 .iter()
230 .copied()
231 .filter(|model| is_alias(&model.id))
232 .collect();
233 if !aliases.is_empty() {
234 aliases.sort_by(|a, b| cmp_id_desc(&a.id, &b.id));
235 return aliases.first().map(|model| (*model).clone());
236 }
237
238 matches.retain(|model| !is_alias(&model.id));
239 matches.sort_by(|a, b| cmp_id_desc(&a.id, &b.id));
240 matches.first().map(|model| (*model).clone())
241}
242
243fn cmp_id_desc(a: &str, b: &str) -> Ordering {
245 b.cmp(a)
246}
247
248fn build_fallback_model(
249 provider: &str,
250 model_id: &str,
251 available_models: &[Model],
252) -> Option<Model> {
253 let provider_models: Vec<&Model> = available_models
254 .iter()
255 .filter(|model| model.provider == provider)
256 .collect();
257 if provider_models.is_empty() {
258 return None;
259 }
260
261 let default_id = default_model_per_provider()
262 .iter()
263 .find(|(name, _)| *name == provider)
264 .map(|(_, id)| *id);
265
266 let base = default_id
267 .and_then(|default_id| {
268 provider_models
269 .iter()
270 .find(|model| model.id == default_id)
271 .copied()
272 })
273 .unwrap_or(provider_models[0]);
274 let mut model = base.clone();
275 model_id.clone_into(&mut model.id);
276 model_id.clone_into(&mut model.name);
277 Some(model)
278}
279
280#[must_use]
284pub fn parse_model_pattern(
285 pattern: &str,
286 available_models: &[Model],
287 options: ParseModelPatternOptions,
288) -> ParsedModelResult {
289 if let Some(exact_match) = try_match_model(pattern, available_models) {
290 return ParsedModelResult {
291 model: Some(exact_match),
292 thinking_level: None,
293 warning: None,
294 };
295 }
296
297 let Some(last_colon_index) = pattern.rfind(':') else {
298 return ParsedModelResult {
299 model: None,
300 thinking_level: None,
301 warning: None,
302 };
303 };
304
305 let prefix = &pattern[..last_colon_index];
306 let suffix = &pattern[last_colon_index + 1..];
307
308 if let Some(level) = parse_thinking_level(suffix) {
309 let result = parse_model_pattern(prefix, available_models, options);
310 if result.model.is_some() {
311 return ParsedModelResult {
312 model: result.model,
313 thinking_level: if result.warning.is_some() {
314 None
315 } else {
316 Some(level)
317 },
318 warning: result.warning,
319 };
320 }
321 return result;
322 }
323
324 let allow_fallback = options
325 .allow_invalid_thinking_level_fallback
326 .unwrap_or(true);
327 if !allow_fallback {
328 return ParsedModelResult {
329 model: None,
330 thinking_level: None,
331 warning: None,
332 };
333 }
334
335 let result = parse_model_pattern(prefix, available_models, options);
336 if result.model.is_some() {
337 return ParsedModelResult {
338 model: result.model,
339 thinking_level: None,
340 warning: Some(format!(
341 "Invalid thinking level \"{suffix}\" in pattern \"{pattern}\". Using default instead."
342 )),
343 };
344 }
345 result
346}
347
348fn pattern_has_glob(pattern: &str) -> bool {
349 pattern.contains('*') || pattern.contains('?') || pattern.contains('[')
350}
351
352fn glob_matches(candidate: &str, pattern: &str) -> bool {
353 let Ok(glob) = GlobBuilder::new(pattern)
354 .case_insensitive(true)
355 .literal_separator(true)
356 .backslash_escape(false)
357 .build()
358 else {
359 return false;
360 };
361 glob.compile_matcher().is_match(candidate)
362}
363
364pub async fn resolve_model_scope_with_diagnostics(
366 patterns: &[String],
367 model_runtime: &ModelRuntime,
368) -> ResolveModelScopeResult {
369 let available = model_runtime.get_available(None).await.unwrap_or_default();
370 resolve_model_scope_from_models(patterns, &available)
371}
372
373#[must_use]
375pub fn resolve_model_scope_from_models(
376 patterns: &[String],
377 available_models: &[Model],
378) -> ResolveModelScopeResult {
379 let mut scoped_models: Vec<ScopedModel> = Vec::new();
380 let mut diagnostics: Vec<ModelScopeDiagnostic> = Vec::new();
381
382 for pattern in patterns {
383 if pattern_has_glob(pattern) {
384 let mut glob_pattern = pattern.as_str();
385 let mut thinking_level = None;
386 if let Some(colon_idx) = pattern.rfind(':') {
387 let suffix = &pattern[colon_idx + 1..];
388 if let Some(level) = parse_thinking_level(suffix) {
389 thinking_level = Some(level);
390 glob_pattern = &pattern[..colon_idx];
391 }
392 }
393
394 let matching: Vec<&Model> = available_models
395 .iter()
396 .filter(|model| {
397 let full_id = format!("{}/{}", model.provider, model.id);
398 glob_matches(&full_id, glob_pattern) || glob_matches(&model.id, glob_pattern)
399 })
400 .collect();
401
402 if matching.is_empty() {
403 diagnostics.push(ModelScopeDiagnostic {
404 kind: ModelScopeDiagnosticKind::Warning,
405 message: format!("No models match pattern \"{pattern}\""),
406 pattern: pattern.clone(),
407 });
408 continue;
409 }
410
411 for model in matching {
412 if !scoped_models
413 .iter()
414 .any(|scoped| models_are_equal(&scoped.model, model))
415 {
416 scoped_models.push(ScopedModel {
417 model: model.clone(),
418 thinking_level,
419 });
420 }
421 }
422 continue;
423 }
424
425 let parsed = parse_model_pattern(
426 pattern,
427 available_models,
428 ParseModelPatternOptions {
429 allow_invalid_thinking_level_fallback: Some(true),
430 },
431 );
432
433 if let Some(warning) = parsed.warning {
434 diagnostics.push(ModelScopeDiagnostic {
435 kind: ModelScopeDiagnosticKind::Warning,
436 message: warning,
437 pattern: pattern.clone(),
438 });
439 }
440
441 let Some(model) = parsed.model else {
442 diagnostics.push(ModelScopeDiagnostic {
443 kind: ModelScopeDiagnosticKind::Warning,
444 message: format!("No models match pattern \"{pattern}\""),
445 pattern: pattern.clone(),
446 });
447 continue;
448 };
449
450 if !scoped_models
451 .iter()
452 .any(|scoped| models_are_equal(&scoped.model, &model))
453 {
454 scoped_models.push(ScopedModel {
455 model,
456 thinking_level: parsed.thinking_level,
457 });
458 }
459 }
460
461 ResolveModelScopeResult {
462 scoped_models,
463 diagnostics,
464 }
465}
466
467pub async fn resolve_model_scope(
472 patterns: &[String],
473 model_runtime: &ModelRuntime,
474) -> Vec<ScopedModel> {
475 resolve_model_scope_with_diagnostics(patterns, model_runtime)
476 .await
477 .scoped_models
478}
479
480#[must_use]
482pub fn resolve_cli_model(options: ResolveCliModelOptions<'_>) -> ResolveCliModelResult {
483 resolve_cli_model_from(
484 options.cli_provider,
485 options.cli_model,
486 options.cli_thinking,
487 &options.model_runtime.get_models(None),
488 |provider| options.model_runtime.has_configured_auth(provider),
489 )
490}
491
492#[must_use]
494pub fn resolve_cli_model_from(
495 cli_provider: Option<&str>,
496 cli_model: Option<&str>,
497 cli_thinking: Option<ModelThinkingLevel>,
498 available_models: &[Model],
499 has_configured_auth: impl Fn(&str) -> bool,
500) -> ResolveCliModelResult {
501 let Some(cli_model) = cli_model else {
502 return empty_cli_result();
503 };
504 if available_models.is_empty() {
505 return no_models_cli_result();
506 }
507
508 let provider_map = build_provider_map(available_models);
509 let (provider, mut pattern, inferred_provider) =
510 infer_provider_and_pattern(cli_provider, cli_model, &provider_map);
511
512 if cli_provider.is_some() && provider.is_none() {
513 return unknown_provider_cli_result(cli_provider);
514 }
515
516 if provider.is_none()
517 && let Some(exact) = find_exact_cli_reference(cli_model, available_models)
518 {
519 return exact_cli_result(exact);
520 }
521
522 strip_explicit_provider_prefix(cli_provider, provider.as_deref(), cli_model, &mut pattern);
523
524 let candidate_owned = filter_candidates(provider.as_deref(), available_models);
525 let parsed = parse_model_pattern(
526 &pattern,
527 &candidate_owned,
528 ParseModelPatternOptions {
529 allow_invalid_thinking_level_fallback: Some(false),
530 },
531 );
532
533 if let Some(model) = parsed.model {
534 if let Some(preferred) = prefer_authenticated_raw_id(
535 inferred_provider,
536 cli_model,
537 &model,
538 available_models,
539 &has_configured_auth,
540 ) {
541 return preferred;
542 }
543 return ResolveCliModelResult {
544 model: Some(model),
545 thinking_level: parsed.thinking_level,
546 warning: parsed.warning,
547 error: None,
548 };
549 }
550
551 if let Some(result) =
552 try_inferred_provider_fallback(inferred_provider, cli_model, available_models)
553 {
554 return result;
555 }
556
557 if let Some(provider_name) = provider.as_deref()
558 && let Some(result) = fallback_custom_model(
559 provider_name,
560 &pattern,
561 cli_thinking,
562 available_models,
563 parsed.warning.as_deref(),
564 )
565 {
566 return result;
567 }
568
569 not_found_cli_result(provider.as_deref(), &pattern, cli_model, parsed.warning)
570}
571
572fn empty_cli_result() -> ResolveCliModelResult {
573 ResolveCliModelResult {
574 model: None,
575 thinking_level: None,
576 warning: None,
577 error: None,
578 }
579}
580
581fn no_models_cli_result() -> ResolveCliModelResult {
582 ResolveCliModelResult {
583 model: None,
584 thinking_level: None,
585 warning: None,
586 error: Some(
587 "No models available. Check your installation or add models to models.json.".to_owned(),
588 ),
589 }
590}
591
592fn unknown_provider_cli_result(cli_provider: Option<&str>) -> ResolveCliModelResult {
593 ResolveCliModelResult {
594 model: None,
595 thinking_level: None,
596 warning: None,
597 error: Some(format!(
598 "Unknown provider \"{}\". Use --list-models to see available providers/models.",
599 cli_provider.unwrap_or_default()
600 )),
601 }
602}
603
604fn exact_cli_result(model: Model) -> ResolveCliModelResult {
605 ResolveCliModelResult {
606 model: Some(model),
607 thinking_level: None,
608 warning: None,
609 error: None,
610 }
611}
612
613fn not_found_cli_result(
614 provider: Option<&str>,
615 pattern: &str,
616 cli_model: &str,
617 warning: Option<String>,
618) -> ResolveCliModelResult {
619 let display = if let Some(provider_name) = provider {
620 format!("{provider_name}/{pattern}")
621 } else {
622 cli_model.to_owned()
623 };
624 ResolveCliModelResult {
625 model: None,
626 thinking_level: None,
627 warning,
628 error: Some(format!(
629 "Model \"{display}\" not found. Use --list-models to see available models."
630 )),
631 }
632}
633
634fn infer_provider_and_pattern(
635 cli_provider: Option<&str>,
636 cli_model: &str,
637 provider_map: &HashMap<String, String>,
638) -> (Option<String>, String, bool) {
639 let mut provider =
640 cli_provider.and_then(|value| provider_map.get(&value.to_ascii_lowercase()).cloned());
641 let mut pattern = cli_model.to_owned();
642 let mut inferred_provider = false;
643 if let Some(slash_index) = cli_model.find('/')
644 && provider.is_none()
645 {
646 let maybe_provider = &cli_model[..slash_index];
647 if let Some(canonical) = provider_map.get(&maybe_provider.to_ascii_lowercase()) {
648 provider = Some(canonical.clone());
649 cli_model[slash_index + 1..].clone_into(&mut pattern);
650 inferred_provider = true;
651 }
652 }
653 (provider, pattern, inferred_provider)
654}
655
656fn strip_explicit_provider_prefix(
657 cli_provider: Option<&str>,
658 provider: Option<&str>,
659 cli_model: &str,
660 pattern: &mut String,
661) {
662 if let (Some(_cli_provider), Some(provider_name)) = (cli_provider, provider) {
663 let prefix = format!("{provider_name}/");
664 if cli_model
665 .to_ascii_lowercase()
666 .starts_with(&prefix.to_ascii_lowercase())
667 {
668 cli_model[prefix.len()..].clone_into(pattern);
669 }
670 }
671}
672
673fn try_inferred_provider_fallback(
674 inferred_provider: bool,
675 cli_model: &str,
676 available_models: &[Model],
677) -> Option<ResolveCliModelResult> {
678 if !inferred_provider {
679 return None;
680 }
681 if let Some(exact) = find_exact_cli_reference(cli_model, available_models) {
682 return Some(exact_cli_result(exact));
683 }
684 let fallback = parse_model_pattern(
685 cli_model,
686 available_models,
687 ParseModelPatternOptions {
688 allow_invalid_thinking_level_fallback: Some(false),
689 },
690 );
691 fallback.model.map(|model| ResolveCliModelResult {
692 model: Some(model),
693 thinking_level: fallback.thinking_level,
694 warning: fallback.warning,
695 error: None,
696 })
697}
698
699fn build_provider_map(available_models: &[Model]) -> HashMap<String, String> {
700 let mut provider_map = HashMap::new();
701 for model in available_models {
702 provider_map
703 .entry(model.provider.to_ascii_lowercase())
704 .or_insert_with(|| model.provider.clone());
705 }
706 provider_map
707}
708
709fn find_exact_cli_reference(cli_model: &str, available_models: &[Model]) -> Option<Model> {
710 let lower = cli_model.to_ascii_lowercase();
711 available_models
712 .iter()
713 .find(|model| {
714 model.id.eq_ignore_ascii_case(&lower)
715 || format!("{}/{}", model.provider, model.id).eq_ignore_ascii_case(&lower)
716 })
717 .cloned()
718}
719
720fn filter_candidates(provider: Option<&str>, available_models: &[Model]) -> Vec<Model> {
721 match provider {
722 Some(provider_name) => available_models
723 .iter()
724 .filter(|model| model.provider == provider_name)
725 .cloned()
726 .collect(),
727 None => available_models.to_vec(),
728 }
729}
730
731fn prefer_authenticated_raw_id(
732 inferred_provider: bool,
733 cli_model: &str,
734 model: &Model,
735 available_models: &[Model],
736 has_configured_auth: &impl Fn(&str) -> bool,
737) -> Option<ResolveCliModelResult> {
738 if !inferred_provider {
739 return None;
740 }
741 let raw_exact: Vec<&Model> = available_models
742 .iter()
743 .filter(|candidate| {
744 candidate.id.eq_ignore_ascii_case(cli_model) && !models_are_equal(candidate, model)
745 })
746 .collect();
747 if raw_exact.is_empty() || has_configured_auth(&model.provider) {
748 return None;
749 }
750 let authenticated_raw: Vec<&Model> = raw_exact
751 .into_iter()
752 .filter(|candidate| has_configured_auth(&candidate.provider))
753 .collect();
754 if authenticated_raw.len() == 1 {
755 return Some(ResolveCliModelResult {
756 model: Some(authenticated_raw[0].clone()),
757 thinking_level: None,
758 warning: None,
759 error: None,
760 });
761 }
762 None
763}
764
765fn fallback_custom_model(
766 provider_name: &str,
767 pattern: &str,
768 cli_thinking: Option<ModelThinkingLevel>,
769 available_models: &[Model],
770 warning: Option<&str>,
771) -> Option<ResolveCliModelResult> {
772 let mut fallback_pattern = pattern;
773 let mut fallback_thinking = None;
774 if cli_thinking.is_none()
775 && let Some(last_colon) = pattern.rfind(':')
776 {
777 let suffix = &pattern[last_colon + 1..];
778 if let Some(level) = parse_thinking_level(suffix) {
779 fallback_pattern = &pattern[..last_colon];
780 fallback_thinking = Some(level);
781 }
782 }
783
784 let mut fallback_model =
785 build_fallback_model(provider_name, fallback_pattern, available_models)?;
786 let requested_thinking = cli_thinking.or(fallback_thinking);
787 if requested_thinking.is_some_and(|level| level != ModelThinkingLevel::Off) {
788 fallback_model.reasoning = true;
789 }
790 let fallback_warning = if let Some(warning) = warning {
791 format!(
792 "{warning} Model \"{fallback_pattern}\" not found for provider \"{provider_name}\". Using custom model id."
793 )
794 } else {
795 format!(
796 "Model \"{fallback_pattern}\" not found for provider \"{provider_name}\". Using custom model id."
797 )
798 };
799 Some(ResolveCliModelResult {
800 model: Some(fallback_model),
801 thinking_level: fallback_thinking,
802 warning: Some(fallback_warning),
803 error: None,
804 })
805}
806
807pub async fn find_initial_model_full(
817 options: FindInitialModelFullOptions<'_>,
818) -> Result<InitialModelResult, String> {
819 if let (Some(cli_provider), Some(cli_model)) = (options.cli_provider, options.cli_model) {
820 let resolved = resolve_cli_model(ResolveCliModelOptions {
821 cli_provider: Some(cli_provider),
822 cli_model: Some(cli_model),
823 cli_thinking: None,
824 model_runtime: options.model_runtime,
825 });
826 if let Some(error) = resolved.error {
827 return Err(error);
828 }
829 if let Some(model) = resolved.model {
830 return Ok(InitialModelResult {
831 model: Some(model),
832 thinking_level: DEFAULT_THINKING_LEVEL,
833 fallback_message: None,
834 });
835 }
836 }
837
838 Ok(find_initial_model(FindInitialModelOptions {
839 cli_model: None,
840 scoped_models: options.scoped_models,
841 is_continuing: options.is_continuing,
842 default_provider: options.default_provider,
843 default_model_id: options.default_model_id,
844 default_thinking_level: options.default_thinking_level,
845 model_runtime: options.model_runtime,
846 })
847 .await)
848}
849
850#[must_use]
852pub fn default_model_per_provider_map() -> &'static [(&'static str, &'static str)] {
853 default_model_per_provider()
854}
855
856#[must_use]
858pub fn default_model_id_for_provider(provider: &str) -> Option<&'static str> {
859 default_model_per_provider()
860 .iter()
861 .find(|(name, _)| *name == provider)
862 .map(|(_, id)| *id)
863}
864
865#[cfg(test)]
866mod tests {
867 use super::*;
868 use pi_ai::types::{ModelCost, ModelInput};
869
870 fn model(id: &str, name: &str, provider: &str, reasoning: bool) -> Model {
871 Model {
872 id: id.to_owned(),
873 name: name.to_owned(),
874 api: "anthropic-messages".to_owned(),
875 provider: provider.to_owned(),
876 base_url: format!("https://{provider}.example"),
877 reasoning,
878 thinking_level_map: None,
879 input: vec![ModelInput::Text],
880 cost: ModelCost {
881 input: 1.0,
882 output: 2.0,
883 cache_read: 0.1,
884 cache_write: 1.0,
885 tiers: None,
886 },
887 context_window: 128_000,
888 max_tokens: 8_192,
889 headers: None,
890 compat: None,
891 extra: std::collections::BTreeMap::default(),
892 }
893 }
894
895 fn all_models() -> Vec<Model> {
896 vec![
897 model("claude-sonnet-4-5", "Claude Sonnet 4.5", "anthropic", true),
898 model("gpt-4o", "GPT-4o", "openai", false),
899 model(
900 "qwen/qwen3-coder:exacto",
901 "Qwen3 Coder Exacto",
902 "openrouter",
903 true,
904 ),
905 model(
906 "openai/gpt-4o:extended",
907 "GPT-4o Extended",
908 "openrouter",
909 false,
910 ),
911 ]
912 }
913
914 #[test]
915 fn parse_exact_and_partial_and_missing() {
916 let models = all_models();
917 let exact = parse_model_pattern(
918 "claude-sonnet-4-5",
919 &models,
920 ParseModelPatternOptions::default(),
921 );
922 assert_eq!(
923 exact.model.as_ref().map(|m| m.id.as_str()),
924 Some("claude-sonnet-4-5")
925 );
926 assert!(exact.thinking_level.is_none());
927 assert!(exact.warning.is_none());
928
929 let partial = parse_model_pattern("sonnet", &models, ParseModelPatternOptions::default());
930 assert_eq!(
931 partial.model.as_ref().map(|m| m.id.as_str()),
932 Some("claude-sonnet-4-5")
933 );
934
935 let missing =
936 parse_model_pattern("nonexistent", &models, ParseModelPatternOptions::default());
937 assert!(missing.model.is_none());
938 }
939
940 #[test]
941 fn parse_valid_thinking_suffixes() {
942 let models = all_models();
943 for level in ["off", "minimal", "low", "medium", "high", "xhigh", "max"] {
944 let result = parse_model_pattern(
945 &format!("sonnet:{level}"),
946 &models,
947 ParseModelPatternOptions::default(),
948 );
949 assert_eq!(
950 result.model.as_ref().map(|m| m.id.as_str()),
951 Some("claude-sonnet-4-5")
952 );
953 assert_eq!(result.thinking_level, parse_thinking_level(level));
954 assert!(result.warning.is_none());
955 }
956 }
957
958 #[test]
959 fn parse_invalid_thinking_suffix_warns_in_scope_mode() {
960 let models = all_models();
961 let result = parse_model_pattern(
962 "sonnet:random",
963 &models,
964 ParseModelPatternOptions::default(),
965 );
966 assert_eq!(
967 result.model.as_ref().map(|m| m.id.as_str()),
968 Some("claude-sonnet-4-5")
969 );
970 assert!(result.thinking_level.is_none());
971 assert!(
972 result
973 .warning
974 .as_deref()
975 .is_some_and(|w| w.contains("Invalid thinking level") && w.contains("random"))
976 );
977 }
978
979 #[test]
980 fn parse_openrouter_colon_ids() {
981 let models = all_models();
982 let exact = parse_model_pattern(
983 "qwen/qwen3-coder:exacto",
984 &models,
985 ParseModelPatternOptions::default(),
986 );
987 assert_eq!(
988 exact.model.as_ref().map(|m| m.id.as_str()),
989 Some("qwen/qwen3-coder:exacto")
990 );
991 assert!(exact.thinking_level.is_none());
992
993 let with_provider = parse_model_pattern(
994 "openrouter/qwen/qwen3-coder:exacto",
995 &models,
996 ParseModelPatternOptions::default(),
997 );
998 assert_eq!(
999 with_provider
1000 .model
1001 .as_ref()
1002 .map(|m| (m.provider.as_str(), m.id.as_str())),
1003 Some(("openrouter", "qwen/qwen3-coder:exacto"))
1004 );
1005
1006 let with_level = parse_model_pattern(
1007 "qwen/qwen3-coder:exacto:high",
1008 &models,
1009 ParseModelPatternOptions::default(),
1010 );
1011 assert_eq!(
1012 with_level.model.as_ref().map(|m| m.id.as_str()),
1013 Some("qwen/qwen3-coder:exacto")
1014 );
1015 assert_eq!(with_level.thinking_level, Some(ModelThinkingLevel::High));
1016
1017 let invalid_tail = parse_model_pattern(
1018 "qwen/qwen3-coder:exacto:random",
1019 &models,
1020 ParseModelPatternOptions::default(),
1021 );
1022 assert_eq!(
1023 invalid_tail.model.as_ref().map(|m| m.id.as_str()),
1024 Some("qwen/qwen3-coder:exacto")
1025 );
1026 assert!(invalid_tail.thinking_level.is_none());
1027 assert!(invalid_tail.warning.is_some());
1028 }
1029
1030 #[test]
1031 fn parse_empty_and_trailing_colon() {
1032 let models = all_models();
1033 let empty = parse_model_pattern("", &models, ParseModelPatternOptions::default());
1034 assert!(empty.model.is_some());
1035
1036 let trailing = parse_model_pattern("sonnet:", &models, ParseModelPatternOptions::default());
1037 assert_eq!(
1038 trailing.model.as_ref().map(|m| m.id.as_str()),
1039 Some("claude-sonnet-4-5")
1040 );
1041 assert!(trailing.warning.is_some());
1042 }
1043
1044 #[test]
1045 fn find_exact_rejects_ambiguous_bare_id() {
1046 let models = vec![
1047 model("shared", "A", "alpha", false),
1048 model("shared", "B", "beta", false),
1049 ];
1050 assert!(find_exact_model_reference_match("shared", &models).is_none());
1051 assert_eq!(
1052 find_exact_model_reference_match("alpha/shared", &models)
1053 .as_ref()
1054 .map(|m| m.provider.as_str()),
1055 Some("alpha")
1056 );
1057 }
1058
1059 #[test]
1060 fn scope_diagnostics_and_duplicate_removal() {
1061 let models = all_models();
1062 let patterns = vec![
1063 "sonnet:high".to_owned(),
1064 "gpt-4o:invalid".to_owned(),
1065 "missing".to_owned(),
1066 "claude-sonnet-4-5".to_owned(),
1067 ];
1068 let result = resolve_model_scope_from_models(&patterns, &models);
1069 assert_eq!(
1070 result
1071 .scoped_models
1072 .iter()
1073 .map(|s| s.model.id.as_str())
1074 .collect::<Vec<_>>(),
1075 vec!["claude-sonnet-4-5", "gpt-4o"]
1076 );
1077 assert_eq!(
1078 result.scoped_models[0].thinking_level,
1079 Some(ModelThinkingLevel::High)
1080 );
1081 assert!(result.scoped_models[1].thinking_level.is_none());
1082 assert_eq!(result.diagnostics.len(), 2);
1083 assert_eq!(
1084 result.diagnostics[0].message,
1085 "Invalid thinking level \"invalid\" in pattern \"gpt-4o:invalid\". Using default instead."
1086 );
1087 assert_eq!(
1088 result.diagnostics[1].message,
1089 "No models match pattern \"missing\""
1090 );
1091 }
1092
1093 #[test]
1094 fn scope_glob_matches_id_or_provider_path() {
1095 let models = all_models();
1096 let patterns = vec!["*sonnet*".to_owned(), "openai/*".to_owned()];
1097 let result = resolve_model_scope_from_models(&patterns, &models);
1098 let ids: Vec<&str> = result
1099 .scoped_models
1100 .iter()
1101 .map(|s| s.model.id.as_str())
1102 .collect();
1103 assert!(ids.contains(&"claude-sonnet-4-5"));
1104 assert!(ids.contains(&"gpt-4o"));
1105 }
1106
1107 #[test]
1108 fn scope_glob_thinking_suffix() {
1109 let models = all_models();
1110 let patterns = vec!["anthropic/*:high".to_owned()];
1111 let result = resolve_model_scope_from_models(&patterns, &models);
1112 assert!(!result.scoped_models.is_empty());
1113 assert!(
1114 result
1115 .scoped_models
1116 .iter()
1117 .all(|s| s.thinking_level == Some(ModelThinkingLevel::High))
1118 );
1119 }
1120
1121 #[test]
1122 fn resolve_cli_provider_slash_and_fuzzy() {
1123 let models = all_models();
1124 let auth = |_p: &str| true;
1125
1126 let by_slash = resolve_cli_model_from(None, Some("openai/gpt-4o"), None, &models, auth);
1127 assert_eq!(
1128 by_slash
1129 .model
1130 .as_ref()
1131 .map(|m| (m.provider.as_str(), m.id.as_str())),
1132 Some(("openai", "gpt-4o"))
1133 );
1134
1135 let fuzzy = resolve_cli_model_from(Some("openai"), Some("4o"), None, &models, auth);
1136 assert_eq!(fuzzy.model.as_ref().map(|m| m.id.as_str()), Some("gpt-4o"));
1137
1138 let thinking = resolve_cli_model_from(None, Some("sonnet:high"), None, &models, auth);
1139 assert_eq!(
1140 thinking.model.as_ref().map(|m| m.id.as_str()),
1141 Some("claude-sonnet-4-5")
1142 );
1143 assert_eq!(thinking.thinking_level, Some(ModelThinkingLevel::High));
1144 }
1145
1146 #[test]
1147 fn resolve_cli_prefers_openrouter_style_raw_id_over_provider_inference() {
1148 let models = all_models();
1149 let result =
1150 resolve_cli_model_from(None, Some("openai/gpt-4o:extended"), None, &models, |_| {
1151 true
1152 });
1153 assert_eq!(
1154 result
1155 .model
1156 .as_ref()
1157 .map(|m| (m.provider.as_str(), m.id.as_str())),
1158 Some(("openrouter", "openai/gpt-4o:extended"))
1159 );
1160 }
1161
1162 #[test]
1163 fn resolve_cli_strict_invalid_suffix_keeps_custom_fallback_id() {
1164 let models = all_models();
1165 let result = resolve_cli_model_from(
1166 Some("openai"),
1167 Some("gpt-4o:extended"),
1168 None,
1169 &models,
1170 |_| true,
1171 );
1172 assert_eq!(
1173 result
1174 .model
1175 .as_ref()
1176 .map(|m| (m.provider.as_str(), m.id.as_str())),
1177 Some(("openai", "gpt-4o:extended"))
1178 );
1179 }
1180
1181 #[test]
1182 fn resolve_cli_custom_model_without_double_prefix() {
1183 let models = all_models();
1184 let result = resolve_cli_model_from(
1185 Some("openrouter"),
1186 Some("openrouter/openai/ghost-model"),
1187 None,
1188 &models,
1189 |_| true,
1190 );
1191 assert_eq!(
1192 result
1193 .model
1194 .as_ref()
1195 .map(|m| (m.provider.as_str(), m.id.as_str())),
1196 Some(("openrouter", "openai/ghost-model"))
1197 );
1198 assert!(result.warning.is_some());
1199 }
1200
1201 #[test]
1202 fn resolve_cli_no_models_error() {
1203 let result = resolve_cli_model_from(Some("openai"), Some("gpt-4o"), None, &[], |_| true);
1204 assert!(result.model.is_none());
1205 assert!(
1206 result
1207 .error
1208 .as_deref()
1209 .is_some_and(|e| e.contains("No models available"))
1210 );
1211 }
1212
1213 #[test]
1214 fn resolve_cli_prefers_provider_split_over_gateway_id() {
1215 let mut models = all_models();
1216 models.push(model("glm-5", "GLM-5", "zai", true));
1217 models.push(model("zai/glm-5", "GLM-5", "vercel-ai-gateway", true));
1218 let result = resolve_cli_model_from(None, Some("zai/glm-5"), None, &models, |_| true);
1219 assert_eq!(
1220 result
1221 .model
1222 .as_ref()
1223 .map(|m| (m.provider.as_str(), m.id.as_str())),
1224 Some(("zai", "glm-5"))
1225 );
1226 }
1227
1228 #[test]
1229 fn resolve_cli_prefers_authenticated_raw_id_over_unauth_inferred_provider() {
1230 let mut models = all_models();
1231 models.push(model(
1232 "xiaomi/mimo-v2.5-pro",
1233 "Xiaomi MiMo via Commandcode",
1234 "commandcode",
1235 false,
1236 ));
1237 models.push(model("mimo-v2.5-pro", "Xiaomi MiMo", "xiaomi", false));
1238 let result = resolve_cli_model_from(
1239 None,
1240 Some("xiaomi/mimo-v2.5-pro"),
1241 None,
1242 &models,
1243 |provider| provider == "commandcode",
1244 );
1245 assert_eq!(
1246 result
1247 .model
1248 .as_ref()
1249 .map(|m| (m.provider.as_str(), m.id.as_str())),
1250 Some(("commandcode", "xiaomi/mimo-v2.5-pro"))
1251 );
1252 }
1253
1254 #[test]
1255 fn resolve_cli_provider_prefixed_fuzzy() {
1256 let models = all_models();
1257 let result = resolve_cli_model_from(None, Some("openrouter/qwen"), None, &models, |_| true);
1258 assert_eq!(
1259 result
1260 .model
1261 .as_ref()
1262 .map(|m| (m.provider.as_str(), m.id.as_str())),
1263 Some(("openrouter", "qwen/qwen3-coder:exacto"))
1264 );
1265 }
1266
1267 #[test]
1268 fn resolve_cli_fallback_strips_thinking_suffix() {
1269 let mut models = all_models();
1270 models.push(model(
1271 "some-base-model",
1272 "Some Base Model",
1273 "neuralwatt",
1274 false,
1275 ));
1276 let result = resolve_cli_model_from(
1277 None,
1278 Some("neuralwatt/zai-org/GLM-5.1-FP8:high"),
1279 None,
1280 &models,
1281 |_| true,
1282 );
1283 assert_eq!(
1284 result
1285 .model
1286 .as_ref()
1287 .map(|m| (m.provider.as_str(), m.id.as_str())),
1288 Some(("neuralwatt", "zai-org/GLM-5.1-FP8"))
1289 );
1290 assert_eq!(result.model.as_ref().map(|m| m.reasoning), Some(true));
1291 assert_eq!(result.thinking_level, Some(ModelThinkingLevel::High));
1292
1293 let invalid = resolve_cli_model_from(
1294 None,
1295 Some("neuralwatt/zai-org/GLM-5.1-FP8:banana"),
1296 None,
1297 &models,
1298 |_| true,
1299 );
1300 assert_eq!(
1301 invalid.model.as_ref().map(|m| m.id.as_str()),
1302 Some("zai-org/GLM-5.1-FP8:banana")
1303 );
1304
1305 let explicit_thinking = resolve_cli_model_from(
1306 None,
1307 Some("neuralwatt/zai-org/GLM-5.1-FP8:high"),
1308 Some(ModelThinkingLevel::Medium),
1309 &models,
1310 |_| true,
1311 );
1312 assert_eq!(
1313 explicit_thinking.model.as_ref().map(|m| m.id.as_str()),
1314 Some("zai-org/GLM-5.1-FP8:high")
1315 );
1316 assert!(explicit_thinking.thinking_level.is_none());
1317 }
1318
1319 #[test]
1320 fn default_model_table_tracks_current_ids() {
1321 assert_eq!(default_model_id_for_provider("openai"), Some("gpt-5.5"));
1322 assert_eq!(
1323 default_model_id_for_provider("openai-codex"),
1324 Some("gpt-5.5")
1325 );
1326 assert_eq!(default_model_id_for_provider("zai"), Some("glm-5.1"));
1327 assert_eq!(
1328 default_model_id_for_provider("minimax"),
1329 Some("MiniMax-M2.7")
1330 );
1331 assert_eq!(
1332 default_model_id_for_provider("vercel-ai-gateway"),
1333 Some("zai/glm-5.1")
1334 );
1335 assert_eq!(
1336 default_model_id_for_provider("ant-ling"),
1337 Some("Ring-2.6-1T")
1338 );
1339 }
1340
1341 #[test]
1342 fn alias_preference_over_dated_versions() {
1343 let models = vec![
1344 model("claude-sonnet-4-5-20241022", "dated old", "anthropic", true),
1345 model("claude-sonnet-4-5-20250929", "dated new", "anthropic", true),
1346 model("claude-sonnet-4-5", "alias", "anthropic", true),
1347 ];
1348 let result = parse_model_pattern("sonnet", &models, ParseModelPatternOptions::default());
1349 assert_eq!(
1350 result.model.as_ref().map(|m| m.id.as_str()),
1351 Some("claude-sonnet-4-5")
1352 );
1353
1354 let dated_only = vec![
1355 model("claude-sonnet-4-5-20241022", "dated old", "anthropic", true),
1356 model("claude-sonnet-4-5-20250929", "dated new", "anthropic", true),
1357 ];
1358 let latest =
1359 parse_model_pattern("sonnet", &dated_only, ParseModelPatternOptions::default());
1360 assert_eq!(
1361 latest.model.as_ref().map(|m| m.id.as_str()),
1362 Some("claude-sonnet-4-5-20250929")
1363 );
1364 }
1365
1366 #[tokio::test]
1367 #[allow(clippy::panic)]
1368 async fn find_initial_model_full_cli_custom_and_available_fallback() {
1369 use std::sync::Arc;
1370
1371 use pi_ai::auth::InMemoryCredentialStore;
1372 use pi_ai::models_store::InMemoryModelsStore;
1373
1374 use crate::core::model_runtime::{
1375 CreateModelRuntimeOptions, ModelsJsonConfig, ProviderConfigInput,
1376 ProviderModelDefinition,
1377 };
1378
1379 let runtime = match ModelRuntime::create(CreateModelRuntimeOptions {
1380 credentials: Some(Arc::new(InMemoryCredentialStore::new())),
1381 models_store: Some(Arc::new(InMemoryModelsStore::new())),
1382 models_config: Some(ModelsJsonConfig::empty()),
1383 allow_model_network: Some(false),
1384 ..CreateModelRuntimeOptions::default()
1385 })
1386 .await
1387 {
1388 Ok(runtime) => runtime,
1389 Err(error) => panic!("runtime: {error}"),
1390 };
1391
1392 if let Err(error) = runtime.register_provider(
1393 "openrouter",
1394 ProviderConfigInput {
1395 base_url: Some("https://openrouter.ai/api/v1".into()),
1396 api: Some("openai-completions".into()),
1397 api_key: Some("sk-test".into()),
1398 models: Some(vec![ProviderModelDefinition {
1399 id: "qwen/qwen3-coder:exacto".into(),
1400 name: Some("Qwen".into()),
1401 api: Some("openai-completions".into()),
1402 base_url: Some("https://openrouter.ai/api/v1".into()),
1403 reasoning: true,
1404 thinking_level_map: None,
1405 input: Some(vec![ModelInput::Text]),
1406 cost: None,
1407 context_window: Some(128_000),
1408 max_tokens: Some(8192),
1409 headers: None,
1410 compat: None,
1411 }]),
1412 ..ProviderConfigInput::default()
1413 },
1414 ) {
1415 panic!("register: {error}");
1416 }
1417
1418 let result = match find_initial_model_full(FindInitialModelFullOptions {
1419 cli_provider: Some("openrouter"),
1420 cli_model: Some("openrouter/openai/ghost-model"),
1421 scoped_models: &[],
1422 is_continuing: false,
1423 default_provider: None,
1424 default_model_id: None,
1425 default_thinking_level: None,
1426 model_runtime: &runtime,
1427 })
1428 .await
1429 {
1430 Ok(result) => result,
1431 Err(error) => panic!("ok: {error}"),
1432 };
1433 assert_eq!(
1434 result
1435 .model
1436 .as_ref()
1437 .map(|m| (m.provider.as_str(), m.id.as_str())),
1438 Some(("openrouter", "openai/ghost-model"))
1439 );
1440
1441 let mut env = pi_ai::auth::ProviderEnv::new();
1443 env.insert("OPENAI_API_KEY".to_owned(), "sk-test".to_owned());
1444 let runtime_auth = match ModelRuntime::create(CreateModelRuntimeOptions {
1445 credentials: Some(Arc::new(InMemoryCredentialStore::new())),
1446 models_store: Some(Arc::new(InMemoryModelsStore::new())),
1447 models_config: Some(ModelsJsonConfig::empty()),
1448 allow_model_network: Some(false),
1449 auth_env: Some(env),
1450 ..CreateModelRuntimeOptions::default()
1451 })
1452 .await
1453 {
1454 Ok(runtime) => runtime,
1455 Err(error) => panic!("runtime: {error}"),
1456 };
1457
1458 let ignored = find_initial_model(FindInitialModelOptions {
1459 cli_model: None,
1460 scoped_models: &[],
1461 is_continuing: false,
1462 default_provider: Some("deepseek"),
1463 default_model_id: Some("deepseek-v4-flash"),
1464 default_thinking_level: None,
1465 model_runtime: &runtime_auth,
1466 })
1467 .await;
1468 assert_eq!(
1469 ignored.model.as_ref().map(|m| m.provider.as_str()),
1470 Some("openai")
1471 );
1472 }
1473}