Clean up build canonical warnings (#6880)
Co-authored-by: Douwe Osinga <douwe@squareup.com>
This commit is contained in:
@@ -19,8 +19,12 @@ use std::collections::{BTreeMap, BTreeSet, HashMap};
|
||||
use std::path::PathBuf;
|
||||
|
||||
const MODELS_DEV_API_URL: &str = "https://models.dev/api.json";
|
||||
const DEFAULT_CONTEXT_LIMIT: usize = 128_000;
|
||||
const SEPARATOR: &str =
|
||||
"================================================================================";
|
||||
const SUBSEPARATOR: &str =
|
||||
"--------------------------------------------------------------------------------";
|
||||
|
||||
// Providers to include in canonical models
|
||||
const ALLOWED_PROVIDERS: &[&str] = &[
|
||||
"anthropic",
|
||||
"google",
|
||||
@@ -37,7 +41,6 @@ const ALLOWED_PROVIDERS: &[&str] = &[
|
||||
"google-vertex",
|
||||
];
|
||||
|
||||
// Normalize provider names from models.dev to our canonical format
|
||||
fn normalize_provider_name(provider: &str) -> &str {
|
||||
match provider {
|
||||
"llama" => "meta-llama",
|
||||
@@ -71,23 +74,11 @@ struct MappingEntry {
|
||||
|
||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||
struct MappingReport {
|
||||
/// Timestamp of this report
|
||||
timestamp: String,
|
||||
|
||||
/// Models that are NOT mapped to canonical models
|
||||
unmapped_models: Vec<ProviderModelPair>,
|
||||
|
||||
/// All mappings: (provider, model) -> canonical model
|
||||
/// Stored per provider for backward compatibility
|
||||
all_mappings: BTreeMap<String, Vec<ModelMapping>>,
|
||||
|
||||
/// Flat list of all mappings for easier comparison (lock file format)
|
||||
mapped_models: Vec<MappingEntry>,
|
||||
|
||||
/// Total models checked per provider
|
||||
model_counts: BTreeMap<String, usize>,
|
||||
|
||||
/// Canonical models referenced
|
||||
canonical_models_used: BTreeSet<String>,
|
||||
}
|
||||
|
||||
@@ -110,16 +101,16 @@ impl MappingReport {
|
||||
mappings: Vec<ModelMapping>,
|
||||
recommended_models: Vec<String>,
|
||||
) {
|
||||
let mapping_map: HashMap<String, String> = mappings
|
||||
let mapping_map: HashMap<&str, &str> = mappings
|
||||
.iter()
|
||||
.map(|m| (m.provider_model.clone(), m.canonical_model.clone()))
|
||||
.map(|m| (m.provider_model.as_str(), m.canonical_model.as_str()))
|
||||
.collect();
|
||||
|
||||
let recommended_set: std::collections::HashSet<String> =
|
||||
recommended_models.into_iter().collect();
|
||||
let recommended_set: std::collections::HashSet<&str> =
|
||||
recommended_models.iter().map(|s| s.as_str()).collect();
|
||||
|
||||
for model in &fetched_models {
|
||||
if !mapping_map.contains_key(model) {
|
||||
if !mapping_map.contains_key(model.as_str()) {
|
||||
self.unmapped_models.push(ProviderModelPair {
|
||||
provider: provider_name.to_string(),
|
||||
model: model.clone(),
|
||||
@@ -127,13 +118,14 @@ impl MappingReport {
|
||||
}
|
||||
}
|
||||
|
||||
for (model, canonical) in &mapping_map {
|
||||
self.canonical_models_used.insert(canonical.clone());
|
||||
for mapping in &mappings {
|
||||
self.canonical_models_used
|
||||
.insert(mapping.canonical_model.clone());
|
||||
self.mapped_models.push(MappingEntry {
|
||||
provider: provider_name.to_string(),
|
||||
model: model.clone(),
|
||||
canonical: canonical.clone(),
|
||||
recommended: recommended_set.contains(model),
|
||||
model: mapping.provider_model.clone(),
|
||||
canonical: mapping.canonical_model.clone(),
|
||||
recommended: recommended_set.contains(mapping.provider_model.as_str()),
|
||||
});
|
||||
}
|
||||
|
||||
@@ -144,13 +136,13 @@ impl MappingReport {
|
||||
}
|
||||
|
||||
fn print_summary(&self) {
|
||||
println!("\n{}", "=".repeat(80));
|
||||
println!("\n{SEPARATOR}");
|
||||
println!("CANONICAL MODEL MAPPING REPORT");
|
||||
println!("{}", "=".repeat(80));
|
||||
println!("{SEPARATOR}");
|
||||
println!("\nGenerated: {}\n", self.timestamp);
|
||||
|
||||
println!("Models Checked Per Provider:");
|
||||
println!("{}", "-".repeat(80));
|
||||
println!("{SUBSEPARATOR}");
|
||||
let mut providers: Vec<_> = self.model_counts.iter().collect();
|
||||
providers.sort_by_key(|(name, _)| *name);
|
||||
for (provider, count) in providers {
|
||||
@@ -166,9 +158,9 @@ impl MappingReport {
|
||||
);
|
||||
}
|
||||
|
||||
println!("\n{}", "=".repeat(80));
|
||||
println!("\n{SEPARATOR}");
|
||||
println!("UNMAPPED MODELS ({})", self.unmapped_models.len());
|
||||
println!("{}", "=".repeat(80));
|
||||
println!("{SEPARATOR}");
|
||||
|
||||
if self.unmapped_models.is_empty() {
|
||||
println!("✓ All models are mapped to canonical models!");
|
||||
@@ -194,12 +186,12 @@ impl MappingReport {
|
||||
}
|
||||
}
|
||||
|
||||
println!("\n{}", "=".repeat(80));
|
||||
println!("\n{SEPARATOR}");
|
||||
println!(
|
||||
"CANONICAL MODELS REFERENCED ({})",
|
||||
self.canonical_models_used.len()
|
||||
);
|
||||
println!("{}", "=".repeat(80));
|
||||
println!("{SEPARATOR}");
|
||||
if self.canonical_models_used.is_empty() {
|
||||
println!(" (none yet)");
|
||||
} else {
|
||||
@@ -210,56 +202,55 @@ impl MappingReport {
|
||||
}
|
||||
}
|
||||
|
||||
println!("\n{}", "=".repeat(80));
|
||||
println!("\n{SEPARATOR}");
|
||||
}
|
||||
|
||||
fn compare_with_previous(&self, previous: &MappingReport) {
|
||||
println!("\n{}", "=".repeat(80));
|
||||
println!("\n{SEPARATOR}");
|
||||
println!("CHANGES SINCE PREVIOUS RUN");
|
||||
println!("{}", "=".repeat(80));
|
||||
println!("{SEPARATOR}");
|
||||
|
||||
let mut prev_map: HashMap<(String, String), String> = HashMap::new();
|
||||
for entry in &previous.mapped_models {
|
||||
prev_map.insert(
|
||||
(entry.provider.clone(), entry.model.clone()),
|
||||
entry.canonical.clone(),
|
||||
);
|
||||
}
|
||||
let prev_map: HashMap<(&str, &str), &str> = previous
|
||||
.mapped_models
|
||||
.iter()
|
||||
.map(|e| {
|
||||
(
|
||||
(e.provider.as_str(), e.model.as_str()),
|
||||
e.canonical.as_str(),
|
||||
)
|
||||
})
|
||||
.collect();
|
||||
|
||||
let mut curr_map: HashMap<(String, String), String> = HashMap::new();
|
||||
for entry in &self.mapped_models {
|
||||
curr_map.insert(
|
||||
(entry.provider.clone(), entry.model.clone()),
|
||||
entry.canonical.clone(),
|
||||
);
|
||||
}
|
||||
let curr_map: HashMap<(&str, &str), &str> = self
|
||||
.mapped_models
|
||||
.iter()
|
||||
.map(|e| {
|
||||
(
|
||||
(e.provider.as_str(), e.model.as_str()),
|
||||
e.canonical.as_str(),
|
||||
)
|
||||
})
|
||||
.collect();
|
||||
|
||||
let mut changed_mappings = Vec::new();
|
||||
let mut added_mappings = Vec::new();
|
||||
let mut removed_mappings = Vec::new();
|
||||
|
||||
for (key @ (provider, model), canonical) in &curr_map {
|
||||
match prev_map.get(key) {
|
||||
Some(prev_canonical) if prev_canonical != canonical => {
|
||||
changed_mappings.push((
|
||||
provider.clone(),
|
||||
model.clone(),
|
||||
prev_canonical.clone(),
|
||||
canonical.clone(),
|
||||
));
|
||||
for (&key @ (provider, model), &canonical) in &curr_map {
|
||||
match prev_map.get(&key) {
|
||||
Some(&prev_canonical) if prev_canonical != canonical => {
|
||||
changed_mappings.push((provider, model, prev_canonical, canonical));
|
||||
}
|
||||
None => {
|
||||
added_mappings.push((provider.clone(), model.clone(), canonical.clone()));
|
||||
}
|
||||
_ => {
|
||||
// No change
|
||||
added_mappings.push((provider, model, canonical));
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
for (key @ (provider, model), canonical) in &prev_map {
|
||||
if !curr_map.contains_key(key) {
|
||||
removed_mappings.push((provider.clone(), model.clone(), canonical.clone()));
|
||||
for (&key @ (provider, model), &canonical) in &prev_map {
|
||||
if !curr_map.contains_key(&key) {
|
||||
removed_mappings.push((provider, model, canonical));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -293,7 +284,7 @@ impl MappingReport {
|
||||
}
|
||||
}
|
||||
|
||||
println!("\n{}", "=".repeat(80));
|
||||
println!("\n{SEPARATOR}");
|
||||
}
|
||||
|
||||
fn save_to_file(&self, path: &PathBuf) -> Result<()> {
|
||||
@@ -328,7 +319,13 @@ impl MappingReport {
|
||||
}
|
||||
}
|
||||
|
||||
async fn build_canonical_models() -> Result<()> {
|
||||
fn data_file_path(filename: &str) -> PathBuf {
|
||||
PathBuf::from(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("src/providers/canonical/data")
|
||||
.join(filename)
|
||||
}
|
||||
|
||||
async fn fetch_models_dev() -> Result<Value> {
|
||||
println!("Fetching models from models.dev API...");
|
||||
|
||||
let client = reqwest::Client::new();
|
||||
@@ -339,10 +336,103 @@ async fn build_canonical_models() -> Result<()> {
|
||||
.await
|
||||
.context("Failed to fetch from models.dev API")?;
|
||||
|
||||
let json: Value = response
|
||||
response
|
||||
.json()
|
||||
.await
|
||||
.context("Failed to parse models.dev response")?;
|
||||
.context("Failed to parse models.dev response")
|
||||
}
|
||||
|
||||
fn get_string(value: &Value, field: &str) -> Option<String> {
|
||||
value.get(field).and_then(|v| v.as_str()).map(String::from)
|
||||
}
|
||||
|
||||
fn parse_modalities(model_data: &Value, field: &str) -> Vec<Modality> {
|
||||
model_data
|
||||
.get("modalities")
|
||||
.and_then(|m| m.get(field))
|
||||
.and_then(|v| v.as_array())
|
||||
.map(|arr| {
|
||||
arr.iter()
|
||||
.filter_map(|v| v.as_str())
|
||||
.filter_map(|s| {
|
||||
serde_json::from_value(serde_json::Value::String(s.to_string())).ok()
|
||||
})
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_else(|| vec![Modality::Text])
|
||||
}
|
||||
|
||||
fn process_model(
|
||||
model_id: &str,
|
||||
model_data: &Value,
|
||||
normalized_provider: &str,
|
||||
) -> Result<Option<(String, CanonicalModel)>> {
|
||||
let cost_data = match model_data.get("cost") {
|
||||
Some(c) if !c.is_null() => c,
|
||||
_ => return Ok(None),
|
||||
};
|
||||
|
||||
let name = model_data["name"]
|
||||
.as_str()
|
||||
.with_context(|| format!("Model {} missing name", model_id))?;
|
||||
|
||||
let canonical_id = canonical_name(normalized_provider, model_id);
|
||||
|
||||
let modalities = Modalities {
|
||||
input: parse_modalities(model_data, "input"),
|
||||
output: parse_modalities(model_data, "output"),
|
||||
};
|
||||
|
||||
let cost = Pricing {
|
||||
input: cost_data.get("input").and_then(|v| v.as_f64()),
|
||||
output: cost_data.get("output").and_then(|v| v.as_f64()),
|
||||
cache_read: cost_data.get("cache_read").and_then(|v| v.as_f64()),
|
||||
cache_write: cost_data.get("cache_write").and_then(|v| v.as_f64()),
|
||||
};
|
||||
|
||||
let limit = Limit {
|
||||
context: model_data
|
||||
.get("limit")
|
||||
.and_then(|l| l.get("context"))
|
||||
.and_then(|v| v.as_u64())
|
||||
.unwrap_or(DEFAULT_CONTEXT_LIMIT as u64) as usize,
|
||||
output: model_data
|
||||
.get("limit")
|
||||
.and_then(|l| l.get("output"))
|
||||
.and_then(|v| v.as_u64())
|
||||
.map(|v| v as usize),
|
||||
};
|
||||
|
||||
let canonical_model = CanonicalModel {
|
||||
id: canonical_id.clone(),
|
||||
name: name.to_string(),
|
||||
family: get_string(model_data, "family"),
|
||||
attachment: model_data.get("attachment").and_then(|v| v.as_bool()),
|
||||
reasoning: model_data.get("reasoning").and_then(|v| v.as_bool()),
|
||||
tool_call: model_data
|
||||
.get("tool_call")
|
||||
.and_then(|v| v.as_bool())
|
||||
.unwrap_or(false),
|
||||
temperature: model_data.get("temperature").and_then(|v| v.as_bool()),
|
||||
knowledge: get_string(model_data, "knowledge"),
|
||||
release_date: get_string(model_data, "release_date"),
|
||||
last_updated: get_string(model_data, "last_updated"),
|
||||
modalities,
|
||||
open_weights: model_data.get("open_weights").and_then(|v| v.as_bool()),
|
||||
cost,
|
||||
limit,
|
||||
};
|
||||
|
||||
let model_name = canonical_id
|
||||
.strip_prefix(&format!("{}/", normalized_provider))
|
||||
.unwrap_or(model_id)
|
||||
.to_string();
|
||||
|
||||
Ok(Some((model_name, canonical_model)))
|
||||
}
|
||||
|
||||
async fn build_canonical_models() -> Result<()> {
|
||||
let json = fetch_models_dev().await?;
|
||||
|
||||
let providers_obj = json
|
||||
.as_object()
|
||||
@@ -355,7 +445,7 @@ async fn build_canonical_models() -> Result<()> {
|
||||
if let Some(provider_data) = providers_obj.get(*provider_key) {
|
||||
let models = provider_data["models"]
|
||||
.as_object()
|
||||
.context(format!("Provider {} missing models object", provider_key))?;
|
||||
.with_context(|| format!("Provider {} missing models object", provider_key))?;
|
||||
|
||||
let normalized_provider = normalize_provider_name(provider_key);
|
||||
|
||||
@@ -366,133 +456,17 @@ async fn build_canonical_models() -> Result<()> {
|
||||
);
|
||||
|
||||
for (model_id, model_data) in models {
|
||||
// Skip models without pricing information
|
||||
let cost_data = match model_data.get("cost") {
|
||||
Some(c) if !c.is_null() => c,
|
||||
_ => continue,
|
||||
};
|
||||
|
||||
let name = model_data["name"]
|
||||
.as_str()
|
||||
.context(format!("Model {} missing name", model_id))?;
|
||||
|
||||
// Use canonical_name to normalize the model ID (strips date stamps, etc.)
|
||||
// This deduplicates different versions of the same model
|
||||
let canonical_id = canonical_name(normalized_provider, model_id);
|
||||
|
||||
let family = model_data
|
||||
.get("family")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string());
|
||||
|
||||
let attachment = model_data.get("attachment").and_then(|v| v.as_bool());
|
||||
|
||||
let reasoning = model_data.get("reasoning").and_then(|v| v.as_bool());
|
||||
|
||||
let tool_call = model_data
|
||||
.get("tool_call")
|
||||
.and_then(|v| v.as_bool())
|
||||
.unwrap_or(false);
|
||||
|
||||
let temperature = model_data.get("temperature").and_then(|v| v.as_bool());
|
||||
|
||||
let knowledge = model_data
|
||||
.get("knowledge")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string());
|
||||
|
||||
let release_date = model_data
|
||||
.get("release_date")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string());
|
||||
|
||||
let last_updated = model_data
|
||||
.get("last_updated")
|
||||
.and_then(|v| v.as_str())
|
||||
.map(|s| s.to_string());
|
||||
|
||||
let modalities = Modalities {
|
||||
input: model_data
|
||||
.get("modalities")
|
||||
.and_then(|m| m.get("input"))
|
||||
.and_then(|v| v.as_array())
|
||||
.map(|arr| {
|
||||
arr.iter()
|
||||
.filter_map(|v| v.as_str())
|
||||
.filter_map(|s| {
|
||||
serde_json::from_value(serde_json::Value::String(s.to_string()))
|
||||
.ok()
|
||||
})
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_else(|| vec![Modality::Text]),
|
||||
output: model_data
|
||||
.get("modalities")
|
||||
.and_then(|m| m.get("output"))
|
||||
.and_then(|v| v.as_array())
|
||||
.map(|arr| {
|
||||
arr.iter()
|
||||
.filter_map(|v| v.as_str())
|
||||
.filter_map(|s| {
|
||||
serde_json::from_value(serde_json::Value::String(s.to_string()))
|
||||
.ok()
|
||||
})
|
||||
.collect()
|
||||
})
|
||||
.unwrap_or_else(|| vec![Modality::Text]),
|
||||
};
|
||||
|
||||
let open_weights = model_data.get("open_weights").and_then(|v| v.as_bool());
|
||||
|
||||
let cost = Pricing {
|
||||
input: cost_data.get("input").and_then(|v| v.as_f64()),
|
||||
output: cost_data.get("output").and_then(|v| v.as_f64()),
|
||||
cache_read: cost_data.get("cache_read").and_then(|v| v.as_f64()),
|
||||
cache_write: cost_data.get("cache_write").and_then(|v| v.as_f64()),
|
||||
};
|
||||
|
||||
let limit = Limit {
|
||||
context: model_data
|
||||
.get("limit")
|
||||
.and_then(|l| l.get("context"))
|
||||
.and_then(|v| v.as_u64())
|
||||
.unwrap_or(128_000) as usize,
|
||||
output: model_data
|
||||
.get("limit")
|
||||
.and_then(|l| l.get("output"))
|
||||
.and_then(|v| v.as_u64())
|
||||
.map(|v| v as usize),
|
||||
};
|
||||
|
||||
let canonical_model = CanonicalModel {
|
||||
id: canonical_id.clone(),
|
||||
name: name.to_string(),
|
||||
family,
|
||||
attachment,
|
||||
reasoning,
|
||||
tool_call,
|
||||
temperature,
|
||||
knowledge,
|
||||
release_date,
|
||||
last_updated,
|
||||
modalities,
|
||||
open_weights,
|
||||
cost,
|
||||
limit,
|
||||
};
|
||||
|
||||
// Extract the normalized model name (everything after "provider/")
|
||||
let model_name = canonical_id
|
||||
.strip_prefix(&format!("{}/", normalized_provider))
|
||||
.unwrap_or(model_id);
|
||||
registry.register(normalized_provider, model_name, canonical_model);
|
||||
total_models += 1;
|
||||
if let Some((model_name, canonical_model)) =
|
||||
process_model(model_id, model_data, normalized_provider)?
|
||||
{
|
||||
registry.register(normalized_provider, &model_name, canonical_model);
|
||||
total_models += 1;
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
let output_path = PathBuf::from(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("src/providers/canonical/data/canonical_models.json");
|
||||
let output_path = data_file_path("canonical_models.json");
|
||||
registry.to_file(&output_path)?;
|
||||
println!(
|
||||
"\n✓ Wrote {} models to {}",
|
||||
@@ -566,11 +540,10 @@ async fn check_provider(
|
||||
}
|
||||
|
||||
async fn check_canonical_mappings() -> Result<()> {
|
||||
println!("\n{}", "=".repeat(80));
|
||||
println!("\n{SEPARATOR}");
|
||||
println!("Canonical Model Checker");
|
||||
println!("Checking model mappings for top providers...\n");
|
||||
|
||||
// Define providers to check with their default models
|
||||
let providers = vec![
|
||||
("anthropic", "claude-3-5-sonnet-20241022"),
|
||||
("openai", "gpt-4"),
|
||||
@@ -595,8 +568,7 @@ async fn check_canonical_mappings() -> Result<()> {
|
||||
|
||||
report.print_summary();
|
||||
|
||||
let output_path = PathBuf::from(env!("CARGO_MANIFEST_DIR"))
|
||||
.join("src/providers/canonical/data/canonical_mapping_report.json");
|
||||
let output_path = data_file_path("canonical_mapping_report.json");
|
||||
|
||||
if output_path.exists() {
|
||||
if let Ok(previous) = MappingReport::load_from_file(&output_path) {
|
||||
@@ -614,10 +586,8 @@ async fn check_canonical_mappings() -> Result<()> {
|
||||
async fn main() -> Result<()> {
|
||||
let args = Args::parse();
|
||||
|
||||
// Build canonical models
|
||||
build_canonical_models().await?;
|
||||
|
||||
// Run the checker unless --no-check is passed
|
||||
if !args.no_check {
|
||||
check_canonical_mappings().await?;
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user