use std::collections::HashSet;
use std::sync::RwLock;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Availability {
Available,
Unavailable,
Unknown,
}
#[derive(Debug)]
pub struct ModelAvailability {
client: reqwest::Client,
catalog_url: String,
provider_id: String,
ids: RwLock<Option<HashSet<String>>>,
}
impl ModelAvailability {
pub fn new(catalog_url: String, provider_id: String) -> Self {
Self {
client: reqwest::Client::new(),
catalog_url,
provider_id,
ids: RwLock::new(None),
}
}
pub async fn refresh(&self) -> Result<usize, String> {
let body = self
.client
.get(&self.catalog_url)
.send()
.await
.map_err(|e| format!("catalog request failed: {e}"))?;
if !body.status().is_success() {
return Err(format!("catalog returned status {}", body.status()));
}
let text = body.text().await.map_err(|e| e.to_string())?;
let ids = parse_catalog_ids(&text)?;
let count = ids.len();
*self
.ids
.write()
.map_err(|_| "model-availability lock poisoned".to_string())? = Some(ids);
Ok(count)
}
pub fn is_available(&self, provider_id: &str, model_name: &str) -> Availability {
if provider_id != self.provider_id {
return Availability::Unknown;
}
let Ok(guard) = self.ids.read() else {
return Availability::Unknown;
};
match &*guard {
None => Availability::Unknown,
Some(ids) if ids.contains(model_name) => Availability::Available,
Some(_) => Availability::Unavailable,
}
}
}
fn parse_catalog_ids(body: &str) -> Result<HashSet<String>, String> {
let catalog: serde_json::Value =
serde_json::from_str(body).map_err(|e| format!("catalog json parse: {e}"))?;
let entries = catalog
.get("data")
.and_then(|data| data.as_array())
.ok_or("catalog json missing `data` array")?;
Ok(entries
.iter()
.filter_map(|model| {
model
.get("id")
.and_then(|id| id.as_str())
.map(str::to_string)
})
.collect())
}
#[cfg(test)]
mod tests {
use super::*;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
const PROVIDER: &str = "prov";
const CATALOG: &str = r#"{"data":[
{"id":"vendor/model-a"},
{"id":"vendor/model-b"},
{"id":"vendor/model-c"}
]}"#;
async fn serve(status: u16, body: &str) -> (MockServer, String) {
let server = MockServer::start().await;
Mock::given(method("GET"))
.and(path("/v1/models"))
.respond_with(ResponseTemplate::new(status).set_body_string(body.to_string()))
.mount(&server)
.await;
let url = format!("{}/v1/models", server.uri());
(server, url)
}
#[test]
fn before_refresh_everything_is_unknown() {
let probe = ModelAvailability::new("http://unused".to_string(), PROVIDER.to_string());
assert_eq!(
probe.is_available(PROVIDER, "vendor/model-a"),
Availability::Unknown
);
}
#[tokio::test]
async fn refresh_marks_present_available_and_absent_unavailable() {
let (_server, url) = serve(200, CATALOG).await;
let probe = ModelAvailability::new(url, PROVIDER.to_string());
let n = probe.refresh().await.unwrap();
assert_eq!(n, 3);
assert_eq!(
probe.is_available(PROVIDER, "vendor/model-b"),
Availability::Available,
"a live id is Available"
);
assert_eq!(
probe.is_available(PROVIDER, "vendor/withdrawn"),
Availability::Unavailable,
"a withdrawn id is Unavailable"
);
}
#[tokio::test]
async fn failed_first_fetch_stays_unknown_never_unavailable() {
let (_server, url) = serve(500, "upstream down").await;
let probe = ModelAvailability::new(url, PROVIDER.to_string());
assert!(probe.refresh().await.is_err());
assert_eq!(
probe.is_available(PROVIDER, "vendor/some-model"),
Availability::Unknown,
"a failed first fetch stays Unknown (fail-open), never Unavailable"
);
}
#[tokio::test]
async fn failed_refresh_retains_last_good_catalog() {
let (_ok_server, ok_url) = serve(200, CATALOG).await;
let probe = ModelAvailability::new(ok_url, PROVIDER.to_string());
probe.refresh().await.unwrap();
let (_bad_server, bad_url) = serve(500, "upstream down").await;
let probe = ModelAvailability::new(bad_url, PROVIDER.to_string());
{
*probe.ids.write().unwrap() = parse_catalog_ids(CATALOG).ok();
}
assert!(probe.refresh().await.is_err());
assert_eq!(
probe.is_available(PROVIDER, "vendor/model-b"),
Availability::Available,
"a failed refresh retains the previous good catalog"
);
}
#[tokio::test]
async fn other_provider_is_unknown() {
let (_server, url) = serve(200, CATALOG).await;
let probe = ModelAvailability::new(url, PROVIDER.to_string());
probe.refresh().await.unwrap();
assert_eq!(
probe.is_available("some-other-provider", "vendor/model-a"),
Availability::Unknown
);
}
#[test]
fn parse_catalog_ids_extracts_ids_and_rejects_garbage() {
let ids = parse_catalog_ids(CATALOG).unwrap();
assert!(ids.contains("vendor/model-b"));
assert_eq!(ids.len(), 3);
assert!(parse_catalog_ids("not json").is_err());
assert!(parse_catalog_ids(r#"{"no_data":1}"#).is_err());
}
}