use std::future::{Future, IntoFuture};
use std::pin::Pin;
use bytes::Bytes;
use futures_util::StreamExt;
use tokio::io::{AsyncWrite, AsyncWriteExt};
use crate::client::Client;
use crate::error::{Result, TransportResultExt};
use crate::http::json;
use crate::models::{ArchitectureVersion, Model, ModelId, Page, ToneId};
impl Client {
pub fn models(&self, tone_id: ToneId) -> ModelList {
ModelList::new(self, tone_id)
}
pub async fn model(&self, id: ModelId) -> Result<Model> {
let req = self.http.get(format!("{}/models/{id}", self.base_url));
let resp = self.send(req).await?;
json(resp).await
}
pub async fn download_model(&self, model: &Model) -> Result<Bytes> {
self.download_url(&model.model_url).await
}
pub async fn download_url(&self, model_url: &str) -> Result<Bytes> {
let req = self.http.get(model_url);
let resp = self.send(req).await?;
resp.bytes().await.transport()
}
pub async fn download_model_json(&self, model: &Model) -> Result<String> {
let bytes = self.download_model(model).await?;
Ok(String::from_utf8(bytes.to_vec())?)
}
pub async fn download_model_to<W>(&self, model: &Model, writer: &mut W) -> Result<u64>
where
W: AsyncWrite + Unpin,
{
self.download_url_to(&model.model_url, writer).await
}
pub async fn download_url_to<W>(&self, model_url: &str, writer: &mut W) -> Result<u64>
where
W: AsyncWrite + Unpin,
{
let req = self.http.get(model_url);
let resp = self.send(req).await?;
let mut stream = resp.bytes_stream();
let mut written: u64 = 0;
while let Some(chunk) = stream.next().await {
let chunk = chunk.transport()?;
writer.write_all(&chunk).await?;
written += chunk.len() as u64;
}
writer.flush().await?;
Ok(written)
}
}
#[must_use = "a request builder does nothing until awaited"]
#[derive(Clone)]
pub struct ModelList {
client: Client,
tone_id: ToneId,
page: Option<u32>,
page_size: Option<u32>,
architecture: Option<ArchitectureVersion>,
}
impl ModelList {
fn new(client: &Client, tone_id: ToneId) -> Self {
Self {
client: client.clone(),
tone_id,
page: None,
page_size: None,
architecture: None,
}
}
pub fn page(mut self, page: u32) -> Self {
self.page = Some(page);
self
}
pub fn page_size(mut self, page_size: u32) -> Self {
self.page_size = Some(page_size);
self
}
pub fn architecture(mut self, architecture: ArchitectureVersion) -> Self {
self.architecture = Some(architecture);
self
}
async fn send(self) -> Result<Page<Model>> {
let mut req = self
.client
.http
.get(format!("{}/models", self.client.base_url))
.query(&[("tone_id", self.tone_id.to_string())]);
if let Some(page) = self.page {
req = req.query(&[("page", page)]);
}
if let Some(page_size) = self.page_size {
req = req.query(&[("page_size", page_size)]);
}
if let Some(arch) = &self.architecture {
req = req.query(&[("architecture", arch.as_str())]);
}
let resp = self.client.send(req).await?;
json(resp).await
}
}
impl IntoFuture for ModelList {
type Output = Result<Page<Model>>;
type IntoFuture = Pin<Box<dyn Future<Output = Self::Output> + Send>>;
fn into_future(self) -> Self::IntoFuture {
Box::pin(self.send())
}
}
impl std::fmt::Debug for ModelList {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ModelList")
.field("tone_id", &self.tone_id)
.field("page", &self.page)
.field("page_size", &self.page_size)
.field("architecture", &self.architecture)
.finish_non_exhaustive()
}
}