use crate::{
client::{Client, Provider},
http_client::{self, HttpClientExt},
model::{Model, ModelList, ModelListingError},
wasm_compat::{WasmCompatSend, WasmCompatSync},
};
#[derive(Debug, serde::Deserialize)]
pub(crate) struct DataEnvelope<Entry> {
pub(crate) data: Vec<Entry>,
}
#[derive(Debug, serde::Deserialize)]
pub(crate) struct ListModelEntry {
pub(crate) id: String,
pub(crate) name: Option<String>,
pub(crate) created: Option<u64>,
pub(crate) owned_by: Option<String>,
}
impl From<ListModelEntry> for Model {
fn from(value: ListModelEntry) -> Self {
let mut model = Model::from_id(value.id);
model.name = value.name;
model.created_at = value.created;
model.owned_by = value.owned_by;
model
}
}
macro_rules! impl_model_lister {
($(#[$meta:meta])* $name:ident, $client:ty, $entry:ty, $label:literal, $path:literal) => {
$(#[$meta])*
#[derive(Clone)]
pub struct $name<H = reqwest::Client> {
client: $client,
}
impl<H> $crate::client::ModelLister<H> for $name<H>
where
H: $crate::http_client::HttpClientExt
+ $crate::wasm_compat::WasmCompatSend
+ $crate::wasm_compat::WasmCompatSync
+ 'static,
{
type Client = $client;
fn new(client: Self::Client) -> Self {
Self { client }
}
async fn list_all(
&self,
) -> Result<$crate::model::ModelList, $crate::model::ModelListingError> {
$crate::providers::internal::model_listing::list_models::<$entry, _, _>(
&self.client,
$label,
$path,
)
.await
}
}
};
}
pub(crate) use impl_model_lister;
pub(crate) fn map_transport_error(
provider_name: &str,
path: &str,
error: http_client::Error,
) -> ModelListingError {
match error {
http_client::Error::InvalidStatusCodeWithMessage(status, message) => {
ModelListingError::api_error_with_context(
provider_name,
path,
status.as_u16(),
message.as_bytes(),
)
}
http_client::Error::InvalidStatusCodeWithDetails { status, body, .. } => {
ModelListingError::api_error_with_context(
provider_name,
path,
status.as_u16(),
body.as_bytes(),
)
}
other => ModelListingError::from(other),
}
}
pub(crate) async fn get_json<T, Ext, H>(
client: &Client<Ext, H>,
provider_name: &str,
path: &str,
) -> Result<T, ModelListingError>
where
T: serde::de::DeserializeOwned,
Ext: Provider + WasmCompatSend + WasmCompatSync + 'static,
H: HttpClientExt + WasmCompatSend + WasmCompatSync + 'static,
{
let body = get_bytes(client, provider_name, path).await?;
serde_json::from_slice(&body).map_err(|error| {
ModelListingError::parse_error_with_context(provider_name, path, &error, &body)
})
}
pub(crate) async fn decode_json_response<T>(
response: http::Response<http_client::LazyBody<Vec<u8>>>,
provider_name: &str,
path: &str,
) -> Result<T, ModelListingError>
where
T: serde::de::DeserializeOwned,
{
if !response.status().is_success() {
let status_code = response.status().as_u16();
let body = response.into_body().await?;
return Err(ModelListingError::api_error_with_context(
provider_name,
path,
status_code,
&body,
));
}
let body = response.into_body().await?;
serde_json::from_slice(&body).map_err(|error| {
ModelListingError::parse_error_with_context(provider_name, path, &error, &body)
})
}
pub(crate) const MAX_LISTING_PAGES: usize = 1000;
#[derive(Debug)]
pub(crate) struct ListingPage {
pub(crate) models: Vec<Model>,
pub(crate) next_cursor: Option<String>,
}
pub(crate) async fn paginate_models<Ext, H, P, Q>(
client: &Client<Ext, H>,
provider_name: &str,
mut path_for: P,
mut parse_page: Q,
) -> Result<ModelList, ModelListingError>
where
Ext: Provider + WasmCompatSend + WasmCompatSync + 'static,
H: HttpClientExt + WasmCompatSend + WasmCompatSync + 'static,
P: FnMut(Option<&str>) -> String,
Q: FnMut(&[u8], &str) -> Result<ListingPage, ModelListingError>,
{
let mut all_models = Vec::new();
let mut cursor: Option<String> = None;
let mut exhausted_page_budget = true;
for _ in 0..MAX_LISTING_PAGES {
let path = path_for(cursor.as_deref());
let body = get_bytes(client, provider_name, &path).await?;
let page = parse_page(&body, &path)?;
all_models.extend(page.models);
let Some(next) = page.next_cursor else {
exhausted_page_budget = false;
break;
};
if cursor.as_deref() == Some(next.as_str()) {
tracing::warn!(
provider = provider_name,
models = all_models.len(),
"model listing repeated its pagination cursor; returning the pages fetched \
so far"
);
exhausted_page_budget = false;
break;
}
cursor = Some(next);
continue;
}
if exhausted_page_budget {
tracing::warn!(
provider = provider_name,
models = all_models.len(),
pages = MAX_LISTING_PAGES,
"model listing hit its page ceiling with a cursor still advancing; returning \
the pages fetched so far"
);
}
Ok(ModelList::new(all_models))
}
pub(crate) fn with_query_pairs(path: &str, pairs: &[(&str, &str)]) -> String {
let mut serializer = url::form_urlencoded::Serializer::new(String::new());
for (name, value) in pairs {
serializer.append_pair(name, value);
}
format!("{path}?{}", serializer.finish())
}
pub(crate) async fn get_bytes<Ext, H>(
client: &Client<Ext, H>,
provider_name: &str,
path: &str,
) -> Result<Vec<u8>, ModelListingError>
where
Ext: Provider + WasmCompatSend + WasmCompatSync + 'static,
H: HttpClientExt + WasmCompatSend + WasmCompatSync + 'static,
{
let req = client.get(path)?.body(http_client::NoBody)?;
let response = client
.send::<_, Vec<u8>>(req)
.await
.map_err(|error| map_transport_error(provider_name, path, error))?;
if !response.status().is_success() {
let status_code = response.status().as_u16();
let body = response.into_body().await?;
return Err(ModelListingError::api_error_with_context(
provider_name,
path,
status_code,
&body,
));
}
Ok(response.into_body().await?)
}
pub(crate) async fn list_models<Entry, Ext, H>(
client: &Client<Ext, H>,
provider_name: &str,
path: &str,
) -> Result<ModelList, ModelListingError>
where
Entry: serde::de::DeserializeOwned + Into<Model>,
Ext: Provider + WasmCompatSend + WasmCompatSync + 'static,
H: HttpClientExt + WasmCompatSend + WasmCompatSync + 'static,
{
let envelope: DataEnvelope<Entry> = get_json(client, provider_name, path).await?;
let models = envelope.data.into_iter().map(Into::into).collect();
Ok(ModelList::new(models))
}
#[cfg(test)]
mod tests {
use super::ListModelEntry;
use crate::model::Model;
#[test]
fn minimal_entry_decodes_with_id_alone() {
let entry: ListModelEntry =
serde_json::from_str(r#"{"id":"gpt-test"}"#).expect("minimal entry should decode");
let model = Model::from(entry);
assert_eq!(model.id, "gpt-test");
assert_eq!(model.name, None);
assert_eq!(model.created_at, None);
assert_eq!(model.owned_by, None);
}
}
#[cfg(test)]
mod transport_error_tests {
use super::*;
#[test]
fn details_variant_maps_to_api_error_with_context() {
let error = map_transport_error(
"test-provider",
"/models",
http_client::Error::InvalidStatusCodeWithDetails {
status: http::StatusCode::UNAUTHORIZED,
body: r#"{"error":"no"}"#.to_string(),
headers: Box::new(http::HeaderMap::new()),
},
);
let with_message = map_transport_error(
"test-provider",
"/models",
http_client::Error::InvalidStatusCodeWithMessage(
http::StatusCode::UNAUTHORIZED,
r#"{"error":"no"}"#.to_string(),
),
);
assert_eq!(format!("{error}"), format!("{with_message}"));
assert!(
matches!(error, ModelListingError::ApiError { .. }),
"got {error:?}"
);
}
}