pub mod request {
use std::collections::HashMap;
use url::Url;
use crate::{
errors::OapiError,
rest::get::{Get, GetNoStream},
};
#[derive(Debug, Clone, Default)]
pub struct RetrieveRequest<'a> {
pub file_id: &'a str,
pub extra_query: HashMap<&'a str, &'a str>,
}
impl<'a> Get for RetrieveRequest<'a> {
fn build_url(&self, base_url: &str) -> Result<String, OapiError> {
let mut url =
Url::parse(base_url.trim_end_matches('/')).map_err(OapiError::UrlError)?;
url.path_segments_mut()
.map_err(|_| OapiError::UrlError(url::ParseError::RelativeUrlWithoutBase))?
.push("files")
.push(self.file_id);
for (key, value) in &self.extra_query {
url.query_pairs_mut().append_pair(key, value);
}
Ok(url.to_string())
}
}
impl<'a> GetNoStream for RetrieveRequest<'a> {
type Response = crate::files::FileObject;
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::{
files::list::request::ListFilesRequest,
rest::{
RequestOptions, default_client,
get::{Get, GetNoStream},
},
};
use anyhow::bail;
use futures_util::future::{self};
const MODELSCOPE_BASE_URL: &str = "https://dashscope.aliyuncs.com/compatible-mode/v1/";
fn modelstudio_key() -> Option<String> {
std::env::var("QWEN_API_KEY")
.ok()
.map(|key| key.trim().to_string())
.filter(|key| !key.is_empty())
}
#[test]
fn test_build_url() {
let request = request::RetrieveRequest {
file_id: "file_id",
..Default::default()
};
let url = request.build_url("https://api.openai.com/v1/").unwrap();
assert_eq!(url, "https://api.openai.com/v1/files/file_id");
}
#[tokio::test]
async fn test_retrieve_file() -> Result<(), anyhow::Error> {
let Some(api_key) = modelstudio_key() else {
println!("Skipping: set QWEN_API_KEY to run this test");
return Ok(());
};
let list_request = ListFilesRequest {
limit: Some(5), ..Default::default()
};
let list_response = list_request
.get_response(
&default_client(),
MODELSCOPE_BASE_URL,
&RequestOptions::bearer(&api_key),
)
.await?;
let client = default_client();
let futures: Vec<_> = list_response
.data
.iter()
.map(|file_object| {
let client = client.clone();
let file_id = file_object.id.clone();
let base_url = MODELSCOPE_BASE_URL.to_string();
let key = api_key.clone();
async move {
let retrieve_request = request::RetrieveRequest {
file_id: &file_id,
..Default::default()
};
retrieve_request
.get_response(&client, &base_url, &RequestOptions::bearer(&key))
.await
}
})
.collect();
let results = future::join_all(futures).await;
for (i, result) in results.iter().enumerate() {
match result {
Ok(file_object) => {
assert_eq!(&list_response.data[i].id, &file_object.id);
assert_eq!(&list_response.data[i].filename, &file_object.filename);
assert_eq!(&list_response.data[i].purpose, &file_object.purpose);
}
Err(e) => {
bail!(
"Failed to get response: {e:#}. The file is: index {i}, {:?}",
list_response.data[i]
)
}
}
}
Ok(())
}
}