use serde::Serialize;
use std::path::PathBuf;
use url::Url;
use crate::errors::OapiError;
use crate::rest::post::{Post, PostNoStream};
#[derive(Debug, Serialize, Clone, Default)]
pub struct CreateFileRequest {
#[serde(skip_serializing)]
pub file: PathBuf,
pub purpose: FilePurpose,
#[serde(skip_serializing_if = "Option::is_none")]
pub expires_after: Option<ExpiresAfter>,
#[serde(skip_serializing_if = "Option::is_none")]
pub extra_body: Option<serde_json::Map<String, serde_json::Value>>,
}
#[derive(Debug, Serialize, Clone, Default)]
pub enum FilePurpose {
#[serde(rename = "assistants")]
Assistants,
#[serde(rename = "batch")]
#[default]
Batch,
#[serde(rename = "fine-tune")]
FineTune,
#[serde(rename = "vision")]
Vision,
#[serde(rename = "user_data")]
UserData,
#[serde(rename = "evals")]
Evals,
#[serde(untagged)]
Other(String),
}
#[derive(Debug, Serialize, Clone)]
#[serde(tag = "anchor", rename_all = "snake_case")]
pub enum ExpiresAfter {
CreatedAt {
seconds: usize,
},
}
impl Post for CreateFileRequest {
#[inline]
fn is_streaming(&self) -> bool {
false
}
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::UrlCannotBeBase(base_url.to_string()))?
.push("files");
Ok(url.to_string())
}
}
impl PostNoStream for CreateFileRequest {
type Response = crate::files::FileObject;
async fn get_response_string(
&self,
client: &reqwest::Client,
url: &str,
key: &str,
) -> Result<String, OapiError> {
let file_content = tokio::fs::read(&self.file).await?;
let file_name = self
.file
.file_name()
.and_then(|name| name.to_str())
.ok_or_else(|| OapiError::ResponseError("Invalid file name".to_string()))?
.to_string();
let file_part = reqwest::multipart::Part::bytes(file_content).file_name(file_name);
let mut form = reqwest::multipart::Form::new().part("file", file_part);
let purpose_str = serde_json::to_string(&self.purpose)
.map_err(|e| OapiError::ResponseError(format!("Failed to serialize purpose: {}", e)))?;
let trimmed_purpose = purpose_str.trim_matches('"').to_string();
form = form.text("purpose", trimmed_purpose);
if let Some(expires_after) = &self.expires_after {
let (anchor, seconds) = match expires_after {
ExpiresAfter::CreatedAt { seconds } => ("created_at", *seconds),
};
form = form
.text("expires_after[anchor]", anchor)
.text("expires_after[seconds]", seconds.to_string());
}
if let Some(extra_body) = &self.extra_body {
for (key, value) in extra_body {
let text = match value {
serde_json::Value::String(s) => s.clone(),
other => other.to_string(),
};
form = form.text(key.clone(), text);
}
}
let response = client
.post(url)
.header("Accept", "application/json")
.bearer_auth(key)
.multipart(form)
.send()
.await?;
crate::rest::response_text_checked(response).await
}
}