mod generated;
pub use generated::{VISION_MODELS, TEXT_ONLY_MODELS, AUDIO_MODELS, ALL_MODELS};
pub use generated::{ModelInfoEntry, MODEL_INFO};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ModelCapabilities {
pub vision: bool,
pub audio: bool,
pub video: bool,
pub file: bool,
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct ModelRanks {
pub overall: Option<f32>,
pub coding: Option<f32>,
pub math: Option<f32>,
pub hard_prompts: Option<f32>,
pub instruction_following: Option<f32>,
pub vision_rank: Option<f32>,
pub style_control: Option<f32>,
}
#[derive(Debug, Clone, Copy, PartialEq)]
pub struct ModelPricing {
pub input_cost_per_m_tokens: Option<f32>,
pub output_cost_per_m_tokens: Option<f32>,
}
#[derive(Debug, Clone, PartialEq)]
pub struct ModelProfile {
pub capabilities: ModelCapabilities,
pub ranks: ModelRanks,
pub pricing: ModelPricing,
pub max_input_tokens: u32,
pub max_output_tokens: u32,
}
impl ModelCapabilities {
pub const fn vision_only() -> Self {
Self {
vision: true,
audio: false,
video: false,
file: false,
}
}
pub const fn text_only() -> Self {
Self {
vision: false,
audio: false,
video: false,
file: false,
}
}
pub const fn full_multimodal() -> Self {
Self {
vision: true,
audio: true,
video: true,
file: true,
}
}
pub fn lookup(model: &str) -> Option<Self> {
let lower = model.to_lowercase();
let info = lookup_model_info(&lower);
let vision = info.map_or(false, |i| i.supports_vision)
|| is_in_list(&lower, VISION_MODELS)
|| supports_vision_by_pattern(&lower);
let audio = info.map_or(false, |i| i.supports_audio)
|| is_in_list(&lower, AUDIO_MODELS);
let video = info.map_or(false, |i| i.supports_video);
let file = info.map_or(false, |i| i.supports_pdf);
if info.is_some() || vision || audio || is_in_list(&lower, TEXT_ONLY_MODELS) {
Some(Self { vision, audio, video, file })
} else {
None
}
}
}
pub fn arena_rank(model: &str) -> Option<f32> {
let lower = model.to_lowercase();
let info = lookup_model_info(&lower)?;
if info.arena_overall == 0 {
None
} else {
Some(info.arena_overall as f32 / 100.0)
}
}
pub fn model_profile(model: &str) -> Option<ModelProfile> {
let lower = model.to_lowercase();
let info = lookup_model_info(&lower)?;
let capabilities = ModelCapabilities {
vision: info.supports_vision,
audio: info.supports_audio,
video: info.supports_video,
file: info.supports_pdf,
};
let ranks = ModelRanks {
overall: if info.arena_overall > 0 {
Some(info.arena_overall as f32 / 100.0)
} else {
None
},
coding: None,
math: None,
hard_prompts: None,
instruction_following: None,
vision_rank: None,
style_control: None,
};
let pricing = ModelPricing {
input_cost_per_m_tokens: if info.cost_input_x1000 > 0 {
Some(info.cost_input_x1000 as f32 / 1000.0)
} else {
None
},
output_cost_per_m_tokens: if info.cost_output_x1000 > 0 {
Some(info.cost_output_x1000 as f32 / 1000.0)
} else {
None
},
};
Some(ModelProfile {
capabilities,
ranks,
pricing,
max_input_tokens: info.max_input_tokens,
max_output_tokens: info.max_output_tokens,
})
}
pub fn supports_video(model: &str) -> bool {
let lower = model.to_lowercase();
if let Some(info) = lookup_model_info(&lower) {
return info.supports_video;
}
false
}
pub fn supports_pdf(model: &str) -> bool {
let lower = model.to_lowercase();
if let Some(info) = lookup_model_info(&lower) {
return info.supports_pdf;
}
false
}
pub fn supports_vision(model: &str) -> bool {
let lower = model.to_lowercase();
if let Some(info) = lookup_model_info(&lower) {
if info.supports_vision {
return true;
}
}
if is_in_list(&lower, VISION_MODELS) {
return true;
}
supports_vision_by_pattern(&lower)
}
pub fn supports_audio(model: &str) -> bool {
let lower = model.to_lowercase();
if let Some(info) = lookup_model_info(&lower) {
if info.supports_audio {
return true;
}
}
is_in_list(&lower, AUDIO_MODELS)
}
pub fn is_text_only(model: &str) -> bool {
!supports_vision(model) && !supports_audio(model)
}
fn lookup_model_info(model: &str) -> Option<&'static ModelInfoEntry> {
if let Ok(idx) = MODEL_INFO.binary_search_by(|entry| entry.name.cmp(model)) {
return Some(&MODEL_INFO[idx]);
}
if let Some(pos) = model.rfind('/') {
let short = &model[pos + 1..];
if let Ok(idx) = MODEL_INFO.binary_search_by(|entry| entry.name.cmp(short)) {
return Some(&MODEL_INFO[idx]);
}
}
let mut best: Option<&'static ModelInfoEntry> = None;
let mut best_len = 0;
for entry in MODEL_INFO {
if model.contains(entry.name) && entry.name.len() > best_len {
best = Some(entry);
best_len = entry.name.len();
}
if entry.name.contains(model) && model.len() >= 4 {
return Some(entry);
}
}
best
}
fn is_in_list(model: &str, list: &[&str]) -> bool {
for entry in list {
if model == *entry {
return true;
}
if model.contains(entry) {
return true;
}
if entry.contains(model) && model.len() >= 4 {
return true;
}
}
false
}
fn supports_vision_by_pattern(model: &str) -> bool {
const VISION_PATTERNS: &[&str] = &[
"gpt-4o",
"gpt-4-turbo",
"gpt-4-vision",
"o1",
"o3",
"o4",
"claude-3",
"claude-4",
"gemini-1.5",
"gemini-2",
"gemini-flash",
"gemini-pro-vision",
"qwen2-vl",
"qwen2.5-vl",
"qwen-vl",
"qwq",
"llama-3.2-vision",
"-vision",
"-vl-",
"-vl:",
"/vl-",
];
for pattern in VISION_PATTERNS {
if model.contains(pattern) {
return true;
}
}
model.ends_with("-vl") || model.ends_with(":vl") || model.ends_with("/vl")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_supports_vision_openai() {
assert!(supports_vision("gpt-4o"));
assert!(supports_vision("gpt-4o-mini"));
assert!(supports_vision("openai/gpt-4o"));
assert!(!supports_vision("gpt-3.5-turbo"));
}
#[test]
fn test_supports_vision_anthropic() {
assert!(supports_vision("claude-3-sonnet"));
assert!(supports_vision("anthropic/claude-3-opus"));
assert!(!supports_vision("claude-2"));
}
#[test]
fn test_supports_vision_google() {
assert!(supports_vision("gemini-2.0-flash"));
assert!(supports_vision("google/gemini-1.5-pro"));
}
#[test]
fn test_is_text_only() {
assert!(is_text_only("gpt-3.5-turbo"));
assert!(is_text_only("claude-2"));
assert!(!is_text_only("gpt-4o"));
}
#[test]
fn test_model_capabilities_lookup() {
let caps = ModelCapabilities::lookup("gpt-4o");
assert!(caps.is_some());
assert!(caps.unwrap().vision);
}
#[test]
fn test_arena_rank_known_models() {
if let Some(rank) = arena_rank("gpt-4o") {
assert!(rank > 50.0 && rank <= 100.0);
}
}
#[test]
fn test_model_profile_has_pricing() {
if let Some(profile) = model_profile("gpt-4o") {
assert!(profile.pricing.input_cost_per_m_tokens.is_some());
assert!(profile.pricing.output_cost_per_m_tokens.is_some());
}
}
#[test]
fn test_model_profile_has_context_window() {
if let Some(profile) = model_profile("gpt-4o") {
assert!(profile.max_input_tokens > 0);
}
}
#[test]
fn test_supports_video_and_pdf() {
assert!(supports_video("gemini-2.5-pro") || supports_video("gemini-2.0-flash"));
}
#[test]
fn test_lookup_model_info_binary_search() {
let info = lookup_model_info("gpt-4o");
assert!(info.is_some());
}
#[test]
fn test_lookup_model_info_substring() {
let info = lookup_model_info("openai/gpt-4o");
assert!(info.is_some());
}
#[test]
fn test_model_info_sorted() {
for window in MODEL_INFO.windows(2) {
assert!(
window[0].name <= window[1].name,
"MODEL_INFO not sorted: {:?} > {:?}",
window[0].name,
window[1].name
);
}
}
#[test]
fn test_supports_vision_qwen_vl() {
assert!(supports_vision("qwen2-vl-72b"));
assert!(supports_vision("qwen2.5-vl-7b"));
assert!(supports_vision("qwen-vl-max"));
assert!(supports_vision("QWEN2-VL"));
}
#[test]
fn test_supports_vision_case_insensitive() {
assert!(supports_vision("GPT-4O"));
assert!(supports_vision("Claude-3-Sonnet"));
assert!(supports_vision("Gemini-2.0-Flash"));
}
#[test]
fn test_capabilities_merge_sources() {
let caps = ModelCapabilities::lookup("qwen2-vl-72b");
assert!(caps.is_some());
assert!(caps.unwrap().vision);
}
}