use std::str::FromStr;
use ai::model::Model;
#[test]
fn test_valid_model_names() {
assert_eq!(Model::from_str("gpt-4.1").unwrap(), Model::GPT41);
assert_eq!(Model::from_str("gpt-4.1-mini").unwrap(), Model::GPT41Mini);
assert_eq!(Model::from_str("gpt-4.1-nano").unwrap(), Model::GPT41Nano);
assert_eq!(Model::from_str("gpt-4.5").unwrap(), Model::GPT45);
}
#[test]
fn test_case_insensitive_parsing() {
assert_eq!(Model::from_str("GPT-4.1").unwrap(), Model::GPT41);
assert_eq!(Model::from_str("Gpt-4.1-Mini").unwrap(), Model::GPT41Mini);
assert_eq!(Model::from_str("GPT-4.1-NANO").unwrap(), Model::GPT41Nano);
assert_eq!(Model::from_str("gPt-4.5").unwrap(), Model::GPT45);
}
#[test]
fn test_whitespace_handling() {
assert_eq!(Model::from_str(" gpt-4.1 ").unwrap(), Model::GPT41);
assert_eq!(Model::from_str("\tgpt-4.1-mini\n").unwrap(), Model::GPT41Mini);
}
#[test]
fn test_deprecated_model_backward_compat() {
assert_eq!(Model::from_str("gpt-4").unwrap(), Model::GPT41);
assert_eq!(Model::from_str("gpt-4o").unwrap(), Model::GPT41);
assert_eq!(Model::from_str("gpt-4o-mini").unwrap(), Model::GPT41Mini);
assert_eq!(Model::from_str("gpt-3.5-turbo").unwrap(), Model::GPT41Mini);
}
#[test]
fn test_arbitrary_model_name_accepted() {
let model = Model::from_str("llama3.1:8b").unwrap();
assert_eq!(model, Model::Other("llama3.1:8b".to_string()));
assert_eq!(model.as_str(), "llama3.1:8b");
assert_eq!(model.to_string(), "llama3.1:8b");
let mixed = Model::from_str("MyCustom-Model").unwrap();
assert_eq!(mixed, Model::Other("MyCustom-Model".to_string()));
}
#[test]
fn test_empty_model_name_rejected() {
assert!(Model::from_str("").is_err());
assert!(Model::from_str(" ").is_err());
}
#[test]
fn test_unknown_model_fallback_carries_through() {
let model = Model::from("custom-model");
assert_eq!(model, Model::Other("custom-model".to_string()));
}
#[test]
fn test_model_display() {
assert_eq!(Model::GPT41.to_string(), "gpt-4.1");
assert_eq!(Model::GPT41Mini.to_string(), "gpt-4.1-mini");
assert_eq!(Model::GPT41Nano.to_string(), "gpt-4.1-nano");
assert_eq!(Model::GPT45.to_string(), "gpt-4.5");
}
#[test]
fn test_model_as_str() {
assert_eq!(Model::GPT41.as_str(), "gpt-4.1");
assert_eq!(Model::GPT41Mini.as_str(), "gpt-4.1-mini");
assert_eq!(Model::GPT41Nano.as_str(), "gpt-4.1-nano");
assert_eq!(Model::GPT45.as_str(), "gpt-4.5");
}
#[test]
fn test_model_as_ref() {
fn takes_str_ref<S: AsRef<str>>(s: S) -> String {
s.as_ref().to_string()
}
assert_eq!(takes_str_ref(Model::GPT41), "gpt-4.1");
assert_eq!(takes_str_ref(Model::GPT41Mini), "gpt-4.1-mini");
}
#[test]
fn test_model_from_string() {
let s = String::from("gpt-4.1");
assert_eq!(Model::from(s), Model::GPT41);
let s = String::from("gpt-4.1-mini");
assert_eq!(Model::from(s), Model::GPT41Mini);
}
#[test]
fn test_default_model() {
assert_eq!(Model::default(), Model::GPT41Mini);
}
#[test]
fn test_unknown_model_tokenizer_and_context_fallback() {
let model = Model::from_str("some-unknown-model-xyz").unwrap();
let count = model
.count_tokens("hello world, this is a token count test")
.unwrap();
assert!(count > 0, "unknown model should still count tokens via cl100k_base fallback");
assert_eq!(model.context_size(), 4096);
}