use std::time::Duration;
use std::time::Instant;
use bytes::Bytes;
use bytes::BytesMut;
use ferrin_provider_util::http::HttpRequest;
use ferrin_provider_util::http::RequestBody;
use ferrin_provider_util::http::ResponseHandlers;
use ferrin_provider_util::http::delete;
use ferrin_provider_util::http::get;
use ferrin_provider_util::http::json_response_handler;
use ferrin_provider_util::http::post_json;
use ferrin_provider_util::http::send;
use ferrin_provider_util::http::text_response_handler;
use ferrin_provider_util::secure_url::validate_url;
use ferrin_spec::Headers;
use ferrin_spec::JsonObject;
use ferrin_spec::JsonValue;
use ferrin_spec::MediaType;
use ferrin_spec::ProviderId;
use ferrin_spec::ProviderReference;
use ferrin_spec::error::ApiCallError;
use ferrin_spec::error::InvalidArgumentError;
use ferrin_spec::error::InvalidResponseDataError;
use ferrin_spec::error::ProviderError;
use ferrin_spec::files::DeleteFileResult;
use ferrin_spec::files::FileMetadataResult;
use ferrin_spec::files::FileReferenceOptions;
use ferrin_spec::files::Files;
use ferrin_spec::files::UploadData;
use ferrin_spec::files::UploadFileOptions;
use ferrin_spec::files::UploadFileResult;
use futures_util::StreamExt;
use futures_util::future::Either;
use futures_util::future::select;
use serde::Deserialize;
use serde_json::json;
use tokio_util::sync::CancellationToken;
use url::Url;
use crate::api_types::deserialize_count;
use crate::config::CANONICAL_OPTIONS_KEY;
use crate::config::SharedConfig;
use crate::config::UPLOAD_PATH;
use crate::convert_prompt::resolve_reference;
use crate::error::failed_response_handler;
use crate::options::parse_merged;
use crate::output::OutputMapper;
pub const DEFAULT_POLL_INTERVAL_MS: u64 = 2_000;
pub const DEFAULT_POLL_TIMEOUT_MS: u64 = 300_000;
#[derive(Debug, Clone, Default, PartialEq, Eq, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct GoogleFilesOptions {
#[serde(default)]
pub display_name: Option<String>,
#[serde(default)]
pub poll_interval_ms: Option<u64>,
#[serde(default)]
pub poll_timeout_ms: Option<u64>,
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct GoogleFile {
pub name: String,
#[serde(default)]
pub display_name: Option<String>,
#[serde(default)]
pub mime_type: Option<String>,
#[serde(default, deserialize_with = "deserialize_count")]
pub size_bytes: Option<u64>,
#[serde(default)]
pub create_time: Option<String>,
#[serde(default)]
pub update_time: Option<String>,
#[serde(default)]
pub expiration_time: Option<String>,
#[serde(default)]
pub sha256_hash: Option<String>,
#[serde(default)]
pub uri: Option<String>,
#[serde(default)]
pub state: Option<String>,
}
#[derive(Debug, Deserialize)]
#[serde(untagged)]
enum FileResponse {
Envelope { file: GoogleFile },
Bare(GoogleFile),
}
impl FileResponse {
fn into_file(self) -> GoogleFile {
match self {
Self::Envelope { file } | Self::Bare(file) => file,
}
}
}
#[derive(Debug, Clone)]
pub struct UploadRequest {
pub data: Bytes,
pub media_type: String,
pub display_name: Option<String>,
pub poll_interval: Duration,
pub poll_timeout: Duration,
pub headers: Headers,
pub cancellation: CancellationToken,
}
impl UploadRequest {
#[must_use]
pub fn new(data: Bytes, media_type: impl Into<String>) -> Self {
Self {
data,
media_type: media_type.into(),
display_name: None,
poll_interval: Duration::from_millis(DEFAULT_POLL_INTERVAL_MS),
poll_timeout: Duration::from_millis(DEFAULT_POLL_TIMEOUT_MS),
headers: Headers::new(),
cancellation: CancellationToken::new(),
}
}
}
async fn collect(data: UploadData) -> Result<Bytes, ProviderError> {
match data {
UploadData::Bytes(bytes) => Ok(bytes),
UploadData::Text(text) => Ok(Bytes::from(text)),
UploadData::Stream(stream) => {
let mut buffer = BytesMut::new();
let mut stream = stream;
while let Some(chunk) = stream.next().await {
buffer.extend_from_slice(&chunk?);
}
Ok(buffer.freeze())
}
#[allow(unreachable_patterns, reason = "UploadData is non-exhaustive")]
_ => Err(ProviderError::unsupported("upload data type")),
}
}
fn parse_time(value: Option<&str>) -> Option<chrono::DateTime<chrono::Utc>> {
value
.and_then(|time| chrono::DateTime::parse_from_rfc3339(time).ok())
.map(|time| time.with_timezone(&chrono::Utc))
}
fn file_name(reference: &str) -> String {
if let Some((_, rest)) = reference.rsplit_once("/files/") {
return format!("files/{rest}");
}
if reference.starts_with("files/") {
return reference.to_owned();
}
format!("files/{reference}")
}
#[derive(Debug, Clone)]
pub struct GoogleFiles {
config: SharedConfig,
provider: ProviderId,
}
impl GoogleFiles {
#[must_use]
pub fn new(config: SharedConfig) -> Self {
Self {
provider: ProviderId::new(config.name.clone()),
config,
}
}
fn reference(&self, file: &GoogleFile) -> ProviderReference {
let value = file.uri.clone().unwrap_or_else(|| file.name.clone());
let mut reference = ProviderReference::new();
reference.insert(CANONICAL_OPTIONS_KEY.to_owned(), value.clone());
reference.insert(self.config.name.clone(), value);
reference
}
#[must_use]
pub fn to_result(&self, file: &GoogleFile) -> UploadFileResult {
let mut meta = JsonObject::new();
let string =
|value: &Option<String>| value.clone().map_or(JsonValue::Null, JsonValue::from);
meta.insert("name".to_owned(), JsonValue::from(file.name.clone()));
meta.insert("displayName".to_owned(), string(&file.display_name));
meta.insert("mimeType".to_owned(), string(&file.mime_type));
meta.insert(
"sizeBytes".to_owned(),
file.size_bytes
.map_or(JsonValue::Null, |size| JsonValue::from(size.to_string())),
);
meta.insert("state".to_owned(), string(&file.state));
meta.insert("uri".to_owned(), string(&file.uri));
for (key, value) in [
("createTime", &file.create_time),
("updateTime", &file.update_time),
("expirationTime", &file.expiration_time),
("sha256Hash", &file.sha256_hash),
] {
if let Some(value) = value {
meta.insert(key.to_owned(), json!(value));
}
}
let mapper = OutputMapper::new(self.config.clone(), Default::default());
UploadFileResult {
provider_reference: self.reference(file),
media_type: file.mime_type.as_deref().map(MediaType::new),
filename: file.display_name.clone(),
byte_size: file.size_bytes,
created_at: parse_time(file.create_time.as_deref()),
expires_at: parse_time(file.expiration_time.as_deref()),
provider_metadata: Some(mapper.metadata(meta)),
warnings: Vec::new(),
}
}
fn resource_url(&self, name: &str) -> Result<Url, ProviderError> {
let mut url = self.config.base_url.clone();
let mut segments = url.path_segments_mut().map_err(|()| {
InvalidArgumentError::new("base_url", "file base URL cannot contain path segments")
})?;
segments.pop_if_empty();
let id = if let Some(id) = name.strip_prefix("files/").filter(|id| !id.contains('/')) {
segments.push("files");
id
} else {
name
};
segments.push(match id {
"." => "%2E",
".." => "%2E%2E",
id => id,
});
drop(segments);
Ok(url)
}
pub async fn fetch_file(
&self,
name: &str,
headers: &Headers,
cancellation: CancellationToken,
) -> Result<GoogleFile, ProviderError> {
let handlers = ResponseHandlers::new(
json_response_handler::<FileResponse>(),
failed_response_handler(),
);
let response = get(
self.config.transport.as_ref(),
self.resource_url(name)?,
self.config.headers(headers)?,
&handlers,
cancellation,
)
.await?;
Ok(response.value.into_file())
}
#[tracing::instrument(skip_all, fields(media_type = %request.media_type, bytes = request.data.len()))]
pub async fn upload_bytes(&self, request: UploadRequest) -> Result<GoogleFile, ProviderError> {
let start_url = self.config.origin_url(UPLOAD_PATH);
let start_headers = self
.config
.headers(&request.headers)?
.with("x-goog-upload-protocol", "resumable")
.with("x-goog-upload-command", "start")
.with(
"x-goog-upload-header-content-length",
&request.data.len().to_string(),
)
.with("x-goog-upload-header-content-type", &request.media_type);
let mut file = JsonObject::new();
if let Some(name) = &request.display_name {
file.insert("display_name".to_owned(), JsonValue::from(name.as_str()));
}
let start_handlers =
ResponseHandlers::new(text_response_handler(), failed_response_handler());
let started = post_json(
self.config.transport.as_ref(),
start_url,
start_headers,
&json!({"file": file}),
&start_handlers,
request.cancellation.clone(),
)
.await?;
let upload_url = started
.response_headers
.get_str("x-goog-upload-url")
.and_then(|value| Url::parse(value).ok())
.ok_or_else(|| {
ProviderError::InvalidResponseData(Box::new(InvalidResponseDataError::new(
"google did not return a resumable upload URL",
JsonValue::Null,
)))
})?;
let validated = validate_url(&upload_url, &self.config.url_policy)
.await
.map_err(|error| {
InvalidResponseDataError::new(
format!("google returned an unsafe upload URL: {error}"),
JsonValue::Null,
)
})?;
let mut finalize_headers = if upload_url.origin() == self.config.base_url.origin()
|| self.config.url_policy.is_credentialed(&upload_url)
{
self.config.unauthenticated_headers(&request.headers)
} else {
Headers::new().with_user_agent_suffix([crate::config::USER_AGENT])
};
finalize_headers.remove(crate::config::API_KEY_HEADER);
let finalize_headers = finalize_headers
.with("x-goog-upload-offset", "0")
.with("x-goog-upload-command", "upload, finalize");
let handlers = ResponseHandlers::new(
json_response_handler::<FileResponse>()
.with_max_bytes(self.config.url_policy.max_body_bytes),
failed_response_handler().with_max_bytes(self.config.url_policy.max_body_bytes),
);
let finalize = HttpRequest::post(upload_url.clone())
.with_headers(finalize_headers)
.with_body(RequestBody::Bytes {
content_type: request.media_type.clone(),
data: request.data,
})
.with_cancellation(request.cancellation.clone())
.with_pinned_addresses(validated.addresses);
let uploaded = send(self.config.transport.as_ref(), finalize, None, &handlers).await?;
let mut file = uploaded.value.into_file();
let started_at = Instant::now();
while file.state.as_deref() == Some("PROCESSING") {
if started_at.elapsed() > request.poll_timeout {
return Err(ProviderError::ApiCall(Box::new(ApiCallError::new(
format!(
"file processing timed out after {}ms",
request.poll_timeout.as_millis()
),
upload_url,
))));
}
let sleep = Box::pin(tokio::time::sleep(request.poll_interval));
let cancelled = Box::pin(request.cancellation.cancelled());
if let Either::Right(_) = select(sleep, cancelled).await {
return Err(ProviderError::Cancelled);
}
file = self
.fetch_file(&file.name, &request.headers, request.cancellation.clone())
.await?;
}
if file.state.as_deref() == Some("FAILED") {
return Err(ProviderError::ApiCall(Box::new(ApiCallError::new(
format!("file processing failed for {}", file.name),
upload_url,
))));
}
Ok(file)
}
}
impl Files for GoogleFiles {
fn provider(&self) -> &ProviderId {
&self.provider
}
async fn upload_file(
&self,
options: UploadFileOptions,
) -> Result<UploadFileResult, ProviderError> {
let google = parse_merged::<GoogleFilesOptions>(
&self.config,
&options.provider_options,
|mut canonical, custom| {
if custom.display_name.is_some() {
canonical.display_name = custom.display_name;
}
if custom.poll_interval_ms.is_some() {
canonical.poll_interval_ms = custom.poll_interval_ms;
}
if custom.poll_timeout_ms.is_some() {
canonical.poll_timeout_ms = custom.poll_timeout_ms;
}
canonical
},
)?;
if google.poll_interval_ms == Some(0) || google.poll_timeout_ms == Some(0) {
return Err(InvalidArgumentError::new(
"provider_options",
"file polling intervals and timeouts must be positive",
)
.into());
}
let ignored_filename = options.filename.is_some();
let data = match select(
Box::pin(collect(options.data)),
Box::pin(options.cancellation.cancelled()),
)
.await
{
Either::Left((data, _)) => data?,
Either::Right(_) => return Err(ProviderError::Cancelled),
};
let mut request = UploadRequest::new(data, options.media_type.as_str());
request.display_name = google.display_name;
if let Some(interval) = google.poll_interval_ms {
request.poll_interval = Duration::from_millis(interval);
}
if let Some(timeout) = google.poll_timeout_ms {
request.poll_timeout = Duration::from_millis(timeout);
}
request.headers = options.headers;
request.cancellation = options.cancellation;
let file = self.upload_bytes(request).await?;
let mut result = self.to_result(&file);
if ignored_filename {
result
.warnings
.push(ferrin_spec::Warning::unsupported("filename"));
}
if result.media_type.is_none() {
result.media_type = Some(options.media_type);
}
Ok(result)
}
fn supports_get_file_metadata(&self) -> bool {
true
}
async fn get_file_metadata(
&self,
options: FileReferenceOptions,
) -> Result<FileMetadataResult, ProviderError> {
let name = file_name(resolve_reference(&self.config, &options.file)?);
let file = self
.fetch_file(&name, &options.headers, options.cancellation)
.await?;
Ok(self.to_result(&file))
}
fn supports_delete_file(&self) -> bool {
true
}
async fn delete_file(
&self,
options: FileReferenceOptions,
) -> Result<DeleteFileResult, ProviderError> {
let name = file_name(resolve_reference(&self.config, &options.file)?);
let handlers = ResponseHandlers::new(text_response_handler(), failed_response_handler());
delete(
self.config.transport.as_ref(),
self.resource_url(&name)?,
self.config.headers(&options.headers)?,
&handlers,
options.cancellation,
)
.await?;
Ok(DeleteFileResult {
provider_reference: options.file,
deleted: true,
provider_metadata: None,
warnings: Vec::new(),
})
}
}