use std::future::Future;
use chrono::DateTime;
use chrono::Utc;
use serde::Deserialize;
use serde::Serialize;
use tokio_util::sync::CancellationToken;
use url::Url;
use crate::dynamic::BoxStream;
use crate::error::ProviderError;
use crate::image_model::AspectRatio;
use crate::image_model::ImageFile;
use crate::image_model::ImageResult;
use crate::image_model::ImageSize;
use crate::language_model::GenerateResult;
use crate::language_model::Prompt;
use crate::language_model::ReasoningEffort;
use crate::language_model::ResponseFormat;
use crate::language_model::SupportedUrls;
use crate::language_model::ToolChoice;
use crate::language_model::ToolDefinition;
use crate::shared::BatchId;
use crate::shared::Headers;
use crate::shared::ModelId;
use crate::shared::ProviderId;
use crate::shared::ProviderMetadata;
use crate::shared::ProviderOptions;
use crate::shared::Warning;
pub type BatchResultStream = BoxStream<'static, Result<BatchItemResult, ProviderError>>;
pub trait Batch: Send + Sync + 'static {
fn provider(&self) -> &ProviderId;
fn supported_urls(&self) -> impl Future<Output = SupportedUrls> + Send;
fn do_start_batch(
&self,
options: BatchStartOptions,
) -> impl Future<Output = Result<BatchStartResult, ProviderError>> + Send;
fn do_get_batch_status(
&self,
options: BatchOperationOptions,
) -> impl Future<Output = Result<BatchStatus, ProviderError>> + Send;
fn do_get_batch_results(
&self,
options: BatchOperationOptions,
) -> impl Future<Output = Result<BatchResultStream, ProviderError>> + Send;
fn supports_cancel_batch(&self) -> bool {
false
}
fn do_cancel_batch(
&self,
options: BatchOperationOptions,
) -> impl Future<Output = Result<BatchCancelResult, ProviderError>> + Send {
let _ = options;
std::future::ready(Err(ProviderError::unsupported("cancel_batch")))
}
fn supports_list_batches(&self) -> bool {
false
}
fn do_list_batches(
&self,
options: BatchListOptions,
) -> impl Future<Output = Result<BatchListResult, ProviderError>> + Send {
let _ = options;
std::future::ready(Err(ProviderError::unsupported("list_batches")))
}
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct TextBatchRequestOptions {
pub prompt: Prompt,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub max_output_tokens: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub temperature: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub stop_sequences: Option<Vec<String>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub top_p: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub top_k: Option<u32>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub presence_penalty: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub frequency_penalty: Option<f64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub seed: Option<u64>,
#[serde(default)]
pub reasoning: ReasoningEffort,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub response_format: Option<ResponseFormat>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub tool_choice: Option<ToolChoice>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub tools: Vec<ToolDefinition>,
#[serde(default, skip_serializing_if = "ProviderOptions::is_empty")]
pub provider_options: ProviderOptions,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct ImageBatchRequestOptions {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub prompt: Option<String>,
pub n: u32,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub size: Option<ImageSize>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub aspect_ratio: Option<AspectRatio>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub seed: Option<u64>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub files: Vec<ImageFile>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub mask: Option<ImageFile>,
#[serde(default, skip_serializing_if = "ProviderOptions::is_empty")]
pub provider_options: ProviderOptions,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "lowercase")]
#[non_exhaustive]
pub enum BatchRequest {
Text {
id: String,
model_id: ModelId,
options: TextBatchRequestOptions,
},
Image {
id: String,
model_id: ModelId,
options: ImageBatchRequestOptions,
},
}
impl BatchRequest {
#[must_use]
pub fn id(&self) -> &str {
match self {
Self::Text { id, .. } | Self::Image { id, .. } => id,
}
}
}
#[derive(Debug, Clone)]
pub struct BatchStartOptions {
pub requests: Vec<BatchRequest>,
pub webhook_url: Option<Url>,
pub provider_options: ProviderOptions,
pub headers: Headers,
pub cancellation: CancellationToken,
}
#[derive(Debug, Clone)]
pub struct BatchOperationOptions {
pub batch_id: BatchId,
pub provider_options: ProviderOptions,
pub headers: Headers,
pub cancellation: CancellationToken,
}
impl BatchOperationOptions {
#[must_use]
pub fn new(batch_id: impl Into<BatchId>) -> Self {
Self {
batch_id: batch_id.into(),
provider_options: ProviderOptions::new(),
headers: Headers::new(),
cancellation: CancellationToken::new(),
}
}
}
#[derive(Debug, Clone, Default)]
pub struct BatchListOptions {
pub limit: Option<usize>,
pub cursor: Option<String>,
pub provider_options: ProviderOptions,
pub headers: Headers,
pub cancellation: CancellationToken,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
#[non_exhaustive]
pub enum BatchState {
Pending,
Completed,
Failed,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct BatchError {
pub message: String,
#[serde(rename = "type", default, skip_serializing_if = "Option::is_none")]
pub error_type: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub code: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub status_code: Option<u16>,
}
#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct BatchRequestCounts {
pub total: u64,
pub pending: u64,
pub completed: u64,
pub failed: u64,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct BatchStatus {
pub status: BatchState,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub raw_status: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub request_counts: Option<BatchRequestCounts>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub error: Option<BatchError>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub created_at: Option<DateTime<Utc>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub expires_at: Option<DateTime<Utc>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_metadata: Option<ProviderMetadata>,
}
impl BatchStatus {
#[must_use]
pub fn new(status: BatchState) -> Self {
Self {
status,
raw_status: None,
request_counts: None,
error: None,
created_at: None,
expires_at: None,
provider_metadata: None,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct BatchWarning {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub request_id: Option<String>,
pub warning: Warning,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct BatchStartResult {
pub batch_id: BatchId,
#[serde(flatten)]
pub status: BatchStatus,
#[serde(default)]
pub warnings: Vec<BatchWarning>,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct BatchCancelResult {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_metadata: Option<ProviderMetadata>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct BatchListItem {
pub batch_id: BatchId,
#[serde(flatten)]
pub status: BatchStatus,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct BatchListResult {
pub batches: Vec<BatchListItem>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub next_cursor: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub provider_metadata: Option<ProviderMetadata>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "status", rename_all = "lowercase")]
#[non_exhaustive]
pub enum BatchItem<R> {
Succeeded {
id: String,
result: R,
},
Failed {
id: String,
error: BatchError,
#[serde(default, skip_serializing_if = "Option::is_none")]
provider_metadata: Option<ProviderMetadata>,
},
Cancelled {
id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
error: Option<BatchError>,
#[serde(default, skip_serializing_if = "Option::is_none")]
provider_metadata: Option<ProviderMetadata>,
},
Expired {
id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
error: Option<BatchError>,
#[serde(default, skip_serializing_if = "Option::is_none")]
provider_metadata: Option<ProviderMetadata>,
},
}
impl<R> BatchItem<R> {
#[must_use]
pub fn id(&self) -> &str {
match self {
Self::Succeeded { id, .. }
| Self::Failed { id, .. }
| Self::Cancelled { id, .. }
| Self::Expired { id, .. } => id,
}
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type", rename_all = "lowercase")]
#[non_exhaustive]
pub enum BatchItemResult {
Text(Box<BatchItem<GenerateResult>>),
Image(Box<BatchItem<ImageResult>>),
}
impl BatchItemResult {
#[must_use]
pub fn id(&self) -> &str {
match self {
Self::Text(item) => item.id(),
Self::Image(item) => item.id(),
}
}
}