use std::collections::HashMap;
use std::path::PathBuf;
use std::sync::OnceLock;
use std::time::{Duration, SystemTime};
use serde::Deserialize;
use tokio::sync::RwLock;
use tracing::{debug, warn};
const DEFAULT_MODELS_URL: &str = "https://models.dev/api.json";
const ENV_MODELS_URL: &str = "AGENT_HARNESS_MODELS_URL";
const CACHE_TTL: Duration = Duration::from_secs(5 * 60);
const FETCH_TIMEOUT: Duration = Duration::from_secs(10);
const MAX_RETRIES: u32 = 3;
const BACKOFF_BASE_MS: u64 = 200;
pub const DEFAULT_CONTEXT_TOKENS: u64 = 128_000;
pub const DEFAULT_OUTPUT_TOKENS: u64 = 8_192;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct ModelLimits {
pub context: u64,
pub output: u64,
}
impl ModelLimits {
pub const fn default_fallback() -> Self {
Self {
context: DEFAULT_CONTEXT_TOKENS,
output: DEFAULT_OUTPUT_TOKENS,
}
}
}
#[derive(Debug, Deserialize)]
struct CatalogRoot(HashMap<String, ProviderEntry>);
#[derive(Debug, Deserialize)]
struct ProviderEntry {
#[serde(default)]
models: HashMap<String, ModelEntry>,
}
#[derive(Debug, Deserialize)]
struct ModelEntry {
#[serde(default)]
limit: Option<LimitEntry>,
}
#[derive(Debug, Deserialize)]
struct LimitEntry {
#[serde(default)]
context: Option<u64>,
#[serde(default)]
output: Option<u64>,
}
struct CatalogState {
table: HashMap<String, ModelLimits>,
fetch_attempted: bool,
}
static CATALOG: OnceLock<RwLock<CatalogState>> = OnceLock::new();
fn catalog() -> &'static RwLock<CatalogState> {
CATALOG.get_or_init(|| {
RwLock::new(CatalogState {
#[cfg(not(test))]
table: load_disk_cache_into_memory(),
#[cfg(test)]
table: HashMap::new(),
fetch_attempted: false,
})
})
}
fn load_disk_cache_into_memory() -> HashMap<String, ModelLimits> {
match (|| -> Option<HashMap<String, ModelLimits>> {
let path = cache_path()?;
let bytes = std::fs::read(&path).ok()?;
let parsed: CatalogRoot = serde_json::from_slice(&bytes).ok()?;
Some(extract_table(&parsed))
})() {
Some(t) => {
debug!("loaded {} models from disk cache at {:?}", t.len(), cache_path());
t
}
None => HashMap::new(),
}
}
fn extract_table(root: &CatalogRoot) -> HashMap<String, ModelLimits> {
let mut out = HashMap::new();
for provider in root.0.values() {
for (model_id, model) in &provider.models {
if let Some(limit) = &model.limit {
let context = limit.context.unwrap_or(DEFAULT_CONTEXT_TOKENS);
let output = limit.output.unwrap_or(DEFAULT_OUTPUT_TOKENS);
if context == 0 {
continue;
}
out.insert(model_id.to_ascii_lowercase(), ModelLimits { context, output });
}
}
}
out
}
pub fn resolve_limits(model: &str) -> ModelLimits {
{
let guard = match catalog().try_read() {
Ok(g) => g,
Err(_) => {
return fallback_table_lookup(model);
}
};
if let Some(limits) = guard.table.get(&model.to_ascii_lowercase()) {
let fallback = fallback_table_lookup(model);
if fallback == ModelLimits::default_fallback() {
return *limits;
}
return ModelLimits {
context: limits.context.max(fallback.context),
output: limits.output.max(fallback.output),
};
}
let attempted = guard.fetch_attempted;
drop(guard);
if !attempted {
trigger_background_fetch_if_needed();
}
}
fallback_table_lookup(model)
}
pub fn resolve_context_window_tokens(model: &str) -> u64 {
resolve_limits(model).context
}
pub fn prefetch() {
trigger_background_fetch_if_needed();
}
fn trigger_background_fetch_if_needed() {
let lock = catalog();
let attempted = match lock.try_read() {
Ok(g) => g.fetch_attempted,
Err(_) => return,
};
if attempted {
return;
}
let handle = match tokio::runtime::Handle::try_current() {
Ok(h) => h,
Err(_) => return,
};
let _join_handle: tokio::task::JoinHandle<()> = handle.spawn(async move {
fetch_and_populate().await;
});
}
async fn fetch_and_populate() {
if cache_is_fresh() {
let mut guard = catalog().write().await;
guard.fetch_attempted = true;
return;
}
let url = std::env::var(ENV_MODELS_URL).unwrap_or_else(|_| DEFAULT_MODELS_URL.to_string());
match fetch_with_retry(&url).await {
Ok(table) => {
if let Some(path) = cache_path() {
if let Err(e) = write_disk_cache_atomic(&path, &table).await {
debug!("disk cache write failed (non-fatal): {e}");
}
}
let mut guard = catalog().write().await;
for (k, v) in table {
guard.table.insert(k, v);
}
guard.fetch_attempted = true;
debug!("models.dev fetch ok; {} models in catalog", guard.table.len());
}
Err(e) => {
warn!("models.dev fetch failed after retries, using fallback table: {e}");
let mut guard = catalog().write().await;
guard.fetch_attempted = true;
}
}
}
async fn fetch_with_retry(url: &str) -> Result<HashMap<String, ModelLimits>, String> {
let client = reqwest::Client::builder()
.timeout(FETCH_TIMEOUT)
.build()
.map_err(|e| e.to_string())?;
let mut last_err = String::from("no attempt made");
for attempt in 0..=MAX_RETRIES {
match client.get(url).send().await {
Ok(resp) => match resp.json::<CatalogRoot>().await {
Ok(root) => return Ok(extract_table(&root)),
Err(e) => last_err = format!("decode failed: {e}"),
},
Err(e) => last_err = format!("request failed: {e}"),
}
if attempt < MAX_RETRIES {
let backoff_ms = BACKOFF_BASE_MS * 2u64.pow(attempt);
let jitter = rand_u64_n(backoff_ms / 2 + 1);
tokio::time::sleep(Duration::from_millis(backoff_ms + jitter)).await;
}
}
Err(last_err)
}
fn rand_u64_n(n: u64) -> u64 {
if n == 0 {
return 0;
}
let nanos = SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.map(|d| d.subsec_nanos() as u64)
.unwrap_or(0);
let tid_hash = format!("{:?}", std::thread::current().id())
.bytes()
.fold(0u64, |acc, b| acc.wrapping_mul(31).wrapping_add(b as u64));
(nanos ^ tid_hash) % n
}
fn cache_path() -> Option<PathBuf> {
if let Ok(p) = std::env::var("AGENT_HARNESS_CACHE_PATH") {
if !p.is_empty() {
return Some(PathBuf::from(p));
}
}
let base = dirs::cache_dir()?;
Some(base.join("agent-harness-rs").join("models.json"))
}
fn cache_is_fresh() -> bool {
let path = match cache_path() {
Some(p) => p,
None => return false,
};
let metadata = match std::fs::metadata(&path) {
Ok(m) => m,
Err(_) => return false,
};
let mtime = match metadata.modified() {
Ok(t) => t,
Err(_) => return false,
};
match SystemTime::now().duration_since(mtime) {
Ok(age) => age < CACHE_TTL,
Err(_) => false,
}
}
async fn write_disk_cache_atomic(
path: &PathBuf,
table: &HashMap<String, ModelLimits>,
) -> std::io::Result<()> {
if let Some(parent) = path.parent() {
tokio::fs::create_dir_all(parent).await?;
}
let body = serialize_table(table);
let ts = SystemTime::now()
.duration_since(SystemTime::UNIX_EPOCH)
.map(|d| d.as_millis())
.unwrap_or(0);
let tmp = path.with_extension(format!("json.{}.{}.tmp", std::process::id(), ts));
tokio::fs::write(&tmp, body).await?;
tokio::fs::rename(&tmp, path).await
}
fn serialize_table(table: &HashMap<String, ModelLimits>) -> String {
use serde_json::json;
let mut models = serde_json::Map::new();
for (id, limits) in table {
models.insert(
id.clone(),
json!({
"limit": {
"context": limits.context,
"output": limits.output,
}
}),
);
}
let provider = json!({ "models": serde_json::Value::Object(models) });
let mut providers = serde_json::Map::new();
providers.insert("cached".to_string(), provider);
serde_json::to_string(&serde_json::Value::Object(providers)).unwrap_or_else(|_| "{}".into())
}
fn fallback_table_lookup(model: &str) -> ModelLimits {
let m = model.to_ascii_lowercase();
if m.contains("opus-4-7") || m.contains("opus-4-6") || m.contains("sonnet-4-6") {
return ModelLimits {
context: 1_000_000,
output: 32_000,
};
}
if m.contains("claude") {
return ModelLimits {
context: 200_000,
output: 8_192,
};
}
if m.contains("gpt-4") || m.contains("gpt-4o") || m.contains("gpt-4.1") {
return ModelLimits {
context: 128_000,
output: 16_384,
};
}
if m.starts_with("o1") || m.starts_with("o3") || m.starts_with("o4") {
return ModelLimits {
context: 200_000,
output: 100_000,
};
}
if m.contains("minimax") || m.contains("deepseek") {
return ModelLimits {
context: 1_000_000,
output: 8_192,
};
}
ModelLimits::default_fallback()
}
#[cfg(test)]
async fn inject_test_table(table: HashMap<String, ModelLimits>) {
let mut guard = catalog().write().await;
guard.table = table;
guard.fetch_attempted = true;
}
#[cfg(test)]
fn extract_table_for_test(json: &str) -> HashMap<String, ModelLimits> {
let root: CatalogRoot = serde_json::from_str(json).unwrap();
extract_table(&root)
}
#[cfg(test)]
static TEST_LOCK: tokio::sync::Mutex<()> = tokio::sync::Mutex::const_new(());
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn fallback_table_known_models() {
assert_eq!(fallback_table_lookup("claude-opus-4-7").context, 1_000_000);
assert_eq!(fallback_table_lookup("claude-sonnet-4-6").context, 1_000_000);
assert_eq!(fallback_table_lookup("claude-haiku-4-5").context, 200_000);
assert_eq!(fallback_table_lookup("claude-3-5-sonnet").context, 200_000);
assert_eq!(fallback_table_lookup("gpt-4o").context, 128_000);
assert_eq!(fallback_table_lookup("gpt-4.1-mini").context, 128_000);
assert_eq!(fallback_table_lookup("o3-mini").context, 200_000);
assert_eq!(fallback_table_lookup("MiniMax-M2").context, 1_000_000);
}
#[test]
fn fallback_table_unknown_model_uses_default() {
let limits = fallback_table_lookup("some-unknown-model-xyz");
assert_eq!(limits.context, DEFAULT_CONTEXT_TOKENS);
assert_eq!(limits.output, DEFAULT_OUTPUT_TOKENS);
}
#[test]
fn fallback_table_case_insensitive() {
assert_eq!(
fallback_table_lookup("CLAUDE-OPUS-4-7").context,
fallback_table_lookup("claude-opus-4-7").context
);
}
#[test]
fn extract_table_picks_only_limit_fields() {
let json = r#"{
"anthropic": {
"models": {
"claude-opus-4-7": {
"name": "Claude Opus 4.7",
"limit": { "context": 1000000, "output": 32000 },
"cost": { "input": 15.0, "output": 75.0 }
},
"claude-haiku-4-5": {
"limit": { "context": 200000, "output": 8192 }
},
"claude-broken": {
"limit": { "context": 0, "output": 100 }
}
}
},
"openai": {
"models": {
"gpt-4o": { "limit": { "context": 128000, "output": 16384 } }
}
}
}"#;
let table = extract_table_for_test(json);
assert_eq!(table.get("claude-opus-4-7").unwrap().context, 1_000_000);
assert_eq!(table.get("claude-opus-4-7").unwrap().output, 32_000);
assert_eq!(table.get("claude-haiku-4-5").unwrap().context, 200_000);
assert_eq!(table.get("gpt-4o").unwrap().context, 128_000);
assert!(!table.contains_key("claude-broken"));
assert_eq!(table.len(), 3);
}
#[test]
fn extract_table_missing_limit_is_skipped() {
let json = r#"{
"acme": {
"models": {
"no-limits-here": { "name": "Mystery Model" }
}
}
}"#;
let table = extract_table_for_test(json);
assert!(table.is_empty());
}
#[tokio::test]
async fn resolve_returns_injected_table_value() {
let _guard = TEST_LOCK.lock().await;
let mut table = HashMap::new();
table.insert(
"injected-model".to_string(),
ModelLimits { context: 42_000, output: 4_000 },
);
inject_test_table(table).await;
let limits = resolve_limits("injected-model");
assert_eq!(limits.context, 42_000);
assert_eq!(limits.output, 4_000);
}
#[tokio::test]
async fn resolve_context_window_tokens_shim_matches_resolve_limits() {
let _guard = TEST_LOCK.lock().await;
let mut table = HashMap::new();
table.insert(
"shim-model".to_string(),
ModelLimits { context: 99_999, output: 1_000 },
);
inject_test_table(table).await;
assert_eq!(resolve_context_window_tokens("shim-model"), 99_999);
}
#[tokio::test]
async fn serialize_then_load_roundtrips() {
let _guard = TEST_LOCK.lock().await;
let mut table = HashMap::new();
table.insert(
"roundtrip-model".to_string(),
ModelLimits { context: 123_456, output: 6_543 },
);
let body = serialize_table(&table);
let parsed = extract_table_for_test(&body);
assert_eq!(parsed.get("roundtrip-model").unwrap().context, 123_456);
assert_eq!(parsed.get("roundtrip-model").unwrap().output, 6_543);
}
#[tokio::test]
async fn fetch_and_populate_succeeds_on_mock_server() {
let _guard = TEST_LOCK.lock().await;
let cache_file = std::env::temp_dir().join(format!(
"ahrs-test-cache-{}.json",
std::process::id()
));
let _ = std::fs::remove_file(&cache_file);
std::env::set_var("AGENT_HARNESS_CACHE_PATH", &cache_file);
let body = r#"{
"mock": {
"models": {
"mock-large": { "limit": { "context": 500000, "output": 8000 } }
}
}
}"#;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let url = format!("http://{addr}/api.json");
std::env::set_var("AGENT_HARNESS_MODELS_URL", &url);
let body_clone = body.to_string();
let server = tokio::spawn(async move {
let (mut stream, _) = listener.accept().await.unwrap();
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let mut buf = vec![0u8; 256];
let mut got = String::new();
while !got.contains("\r\n\r\n") {
let n = stream.read(&mut buf).await.unwrap();
if n == 0 {
break;
}
got.push_str(String::from_utf8_lossy(&buf[..n]).as_ref());
}
let resp = format!(
"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: {}\r\nConnection: close\r\n\r\n{}",
body_clone.len(),
body_clone
);
stream.write_all(resp.as_bytes()).await.unwrap();
stream.flush().await.unwrap();
});
{
let mut guard = catalog().write().await;
guard.table.clear();
guard.fetch_attempted = false;
}
fetch_and_populate().await;
let _ = tokio::time::timeout(Duration::from_secs(2), server).await;
let guard = catalog().read().await;
assert_eq!(guard.table.get("mock-large").unwrap().context, 500_000);
assert!(guard.fetch_attempted);
let _ = std::fs::remove_file(&cache_file);
}
#[tokio::test]
async fn fetch_and_populate_falls_back_when_endpoint_unreachable() {
let _guard = TEST_LOCK.lock().await;
let cache_file = std::env::temp_dir().join(format!(
"ahrs-test-cache-fb-{}.json",
std::process::id()
));
let _ = std::fs::remove_file(&cache_file);
std::env::set_var("AGENT_HARNESS_CACHE_PATH", &cache_file);
std::env::set_var("AGENT_HARNESS_MODELS_URL", "http://127.0.0.1:1/api.json");
{
let mut guard = catalog().write().await;
guard.table.clear();
guard.fetch_attempted = false;
}
fetch_and_populate().await;
let guard = catalog().read().await;
assert!(guard.fetch_attempted);
assert!(guard.table.is_empty());
let _ = std::fs::remove_file(&cache_file);
}
}