use std::{borrow::Cow, fmt, time::Duration};
use bytes::Bytes;
use http::{HeaderMap, Method, StatusCode, Uri};
use serde::{
Deserialize, Deserializer, Serialize, Serializer,
de::{self, IgnoredAny, MapAccess, Visitor},
ser::SerializeStruct,
};
use crate::{
client::Client,
codec,
de::{KeyIn, invalid_response},
error::Error,
request::{CallHeaders, Deadline},
response::ResponseMeta,
retry::{self, RetryPolicy},
transport::{self, Exchange, HttpService},
};
pub struct Models<'a, S> {
client: &'a Client<S>,
}
impl<'a, S> Models<'a, S> {
pub(crate) fn new(client: &'a Client<S>) -> Self {
Self { client }
}
pub fn list(&self) -> ListModels<'a, S> {
ListModels {
client: self.client,
deadline: Deadline::Client,
headers: CallHeaders::default(),
retry: None,
}
}
}
impl<S> fmt::Debug for Models<'_, S> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.debug_struct("Models").finish_non_exhaustive()
}
}
#[must_use = "a request does nothing until it is sent"]
pub struct ListModels<'a, S> {
client: &'a Client<S>,
deadline: Deadline,
headers: CallHeaders<'a>,
retry: Option<RetryPolicy>,
}
impl<'a, S> ListModels<'a, S> {
pub fn timeout(mut self, timeout: Duration) -> Self {
self.deadline = Deadline::After(timeout);
self
}
pub fn no_timeout(mut self) -> Self {
self.deadline = Deadline::Never;
self
}
pub fn header(mut self, name: impl Into<Cow<'a, str>>, value: impl Into<Cow<'a, str>>) -> Self {
self.headers.push(name.into(), value.into());
self
}
pub fn retry(mut self, policy: RetryPolicy) -> Self {
self.retry = Some(policy);
self
}
}
impl<S> ListModels<'_, S>
where
S: HttpService,
{
pub async fn send(self) -> Result<ListModelsResponse, Error> {
let shared = self.client.shared();
let deadline = self.deadline.resolve(shared.config.timeout())?;
let headers = self.headers.parse(false)?;
let uri = shared.config.endpoints().models();
let exchange = Exchange {
method: &Method::GET,
uri,
base_headers: &shared.get_headers,
call_headers: &headers,
deadline,
max_response_bytes: shared.config.max_response_bytes(),
};
let policy = self.retry.as_ref().unwrap_or(&shared.retry);
retry::run(policy, &Method::GET, uri, |retry| async move {
let (status, headers, body) =
transport::attempt(&shared.service, exchange, retry, None).await?;
decode_list_models(body, status, headers, Some((&Method::GET, uri)))
})
.await
}
}
impl<S> fmt::Debug for ListModels<'_, S> {
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
let mut shown = formatter.debug_struct("ListModels");
shown.field("deadline", &self.deadline).field("headers", &self.headers);
if let Some(retry) = &self.retry {
shown.field("retry", retry);
}
shown.finish_non_exhaustive()
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize)]
pub struct ModelMetadata {
name: String,
description: String,
release_date: String,
}
impl ModelMetadata {
#[must_use]
pub fn name(&self) -> &str {
&self.name
}
#[must_use]
pub fn description(&self) -> &str {
&self.description
}
#[must_use]
pub fn release_date(&self) -> &str {
&self.release_date
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct ListModelsResponse {
models: Vec<ModelMetadata>,
meta: ResponseMeta,
}
impl ListModelsResponse {
#[must_use]
pub fn models(&self) -> &[ModelMetadata] {
&self.models
}
#[must_use]
pub fn meta(&self) -> &ResponseMeta {
&self.meta
}
#[must_use]
pub fn into_models(self) -> Vec<ModelMetadata> {
self.models
}
}
impl Serialize for ListModelsResponse {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
let mut out = serializer.serialize_struct("ListModelsResponse", 1)?;
out.serialize_field("models", &self.models)?;
out.end()
}
}
impl<'de> Deserialize<'de> for ModelMetadata {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
deserializer.deserialize_map(CardVisitor)
}
}
struct CardVisitor;
impl<'de> Visitor<'de> for CardVisitor {
type Value = ModelMetadata;
fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str("a model card")
}
fn visit_map<M>(self, mut map: M) -> Result<ModelMetadata, M::Error>
where
M: MapAccess<'de>,
{
let (mut name, mut description, mut release_date) = (None, None, None);
while let Some(index) =
map.next_key_seed(KeyIn(&["name", "description", "release_date"]))?
{
match index {
Some(0) => name = Some(map.next_value()?),
Some(1) => description = Some(map.next_value()?),
Some(2) => release_date = Some(map.next_value()?),
_ => {
map.next_value::<IgnoredAny>()?;
}
}
}
Ok(ModelMetadata {
name: name.ok_or_else(|| de::Error::missing_field("name"))?,
description: description.ok_or_else(|| de::Error::missing_field("description"))?,
release_date: release_date.ok_or_else(|| de::Error::missing_field("release_date"))?,
})
}
}
struct ModelList {
models: Vec<ModelMetadata>,
}
impl<'de> Deserialize<'de> for ModelList {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
deserializer.deserialize_map(ModelListVisitor)
}
}
struct ModelListVisitor;
impl<'de> Visitor<'de> for ModelListVisitor {
type Value = ModelList;
fn expecting(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
formatter.write_str("a list of models")
}
fn visit_map<M>(self, mut map: M) -> Result<ModelList, M::Error>
where
M: MapAccess<'de>,
{
let mut models = None;
while let Some(index) = map.next_key_seed(KeyIn(&["models"]))? {
match index {
Some(_) => models = Some(map.next_value()?),
None => {
map.next_value::<IgnoredAny>()?;
}
}
}
Ok(ModelList { models: models.ok_or_else(|| de::Error::missing_field("models"))? })
}
}
pub(crate) fn decode_list_models(
body: Bytes,
status: StatusCode,
headers: HeaderMap,
endpoint: Option<(&Method, &Uri)>,
) -> Result<ListModelsResponse, Error> {
let meta = ResponseMeta::new(status, headers, body);
let decoded = codec::decode::<ModelList>(meta.raw_body());
match decoded {
Ok(ModelList { models }) => Ok(ListModelsResponse { models, meta }),
Err(source) => Err(invalid_response(meta, endpoint, source)),
}
}
#[cfg(test)]
#[path = "models_tests.rs"]
mod tests;