1use std::future::Future;
4
5use chrono::DateTime;
6use chrono::Utc;
7use serde::Deserialize;
8use serde::Serialize;
9use tokio_util::sync::CancellationToken;
10use url::Url;
11
12use crate::dynamic::BoxStream;
13use crate::error::ProviderError;
14use crate::image_model::AspectRatio;
15use crate::image_model::ImageFile;
16use crate::image_model::ImageResult;
17use crate::image_model::ImageSize;
18use crate::language_model::GenerateResult;
19use crate::language_model::Prompt;
20use crate::language_model::ReasoningEffort;
21use crate::language_model::ResponseFormat;
22use crate::language_model::SupportedUrls;
23use crate::language_model::ToolChoice;
24use crate::language_model::ToolDefinition;
25use crate::shared::BatchId;
26use crate::shared::Headers;
27use crate::shared::ModelId;
28use crate::shared::ProviderId;
29use crate::shared::ProviderMetadata;
30use crate::shared::ProviderOptions;
31use crate::shared::Warning;
32
33pub type BatchResultStream = BoxStream<'static, Result<BatchItemResult, ProviderError>>;
35
36pub trait Batch: Send + Sync + 'static {
38 fn provider(&self) -> &ProviderId;
40
41 fn supported_urls(&self) -> impl Future<Output = SupportedUrls> + Send;
43
44 fn do_start_batch(
46 &self,
47 options: BatchStartOptions,
48 ) -> impl Future<Output = Result<BatchStartResult, ProviderError>> + Send;
49
50 fn do_get_batch_status(
52 &self,
53 options: BatchOperationOptions,
54 ) -> impl Future<Output = Result<BatchStatus, ProviderError>> + Send;
55
56 fn do_get_batch_results(
58 &self,
59 options: BatchOperationOptions,
60 ) -> impl Future<Output = Result<BatchResultStream, ProviderError>> + Send;
61
62 fn supports_cancel_batch(&self) -> bool {
64 false
65 }
66
67 fn do_cancel_batch(
69 &self,
70 options: BatchOperationOptions,
71 ) -> impl Future<Output = Result<BatchCancelResult, ProviderError>> + Send {
72 let _ = options;
73 std::future::ready(Err(ProviderError::unsupported("cancel_batch")))
74 }
75
76 fn supports_list_batches(&self) -> bool {
78 false
79 }
80
81 fn do_list_batches(
83 &self,
84 options: BatchListOptions,
85 ) -> impl Future<Output = Result<BatchListResult, ProviderError>> + Send {
86 let _ = options;
87 std::future::ready(Err(ProviderError::unsupported("list_batches")))
88 }
89}
90
91#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
93pub struct TextBatchRequestOptions {
94 pub prompt: Prompt,
96 #[serde(default, skip_serializing_if = "Option::is_none")]
98 pub max_output_tokens: Option<u32>,
99 #[serde(default, skip_serializing_if = "Option::is_none")]
101 pub temperature: Option<f64>,
102 #[serde(default, skip_serializing_if = "Option::is_none")]
104 pub stop_sequences: Option<Vec<String>>,
105 #[serde(default, skip_serializing_if = "Option::is_none")]
107 pub top_p: Option<f64>,
108 #[serde(default, skip_serializing_if = "Option::is_none")]
110 pub top_k: Option<u32>,
111 #[serde(default, skip_serializing_if = "Option::is_none")]
113 pub presence_penalty: Option<f64>,
114 #[serde(default, skip_serializing_if = "Option::is_none")]
116 pub frequency_penalty: Option<f64>,
117 #[serde(default, skip_serializing_if = "Option::is_none")]
119 pub seed: Option<u64>,
120 #[serde(default)]
122 pub reasoning: ReasoningEffort,
123 #[serde(default, skip_serializing_if = "Option::is_none")]
125 pub response_format: Option<ResponseFormat>,
126 #[serde(default, skip_serializing_if = "Option::is_none")]
128 pub tool_choice: Option<ToolChoice>,
129 #[serde(default, skip_serializing_if = "Vec::is_empty")]
131 pub tools: Vec<ToolDefinition>,
132 #[serde(default, skip_serializing_if = "ProviderOptions::is_empty")]
134 pub provider_options: ProviderOptions,
135}
136
137#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
139pub struct ImageBatchRequestOptions {
140 #[serde(default, skip_serializing_if = "Option::is_none")]
142 pub prompt: Option<String>,
143 pub n: u32,
145 #[serde(default, skip_serializing_if = "Option::is_none")]
147 pub size: Option<ImageSize>,
148 #[serde(default, skip_serializing_if = "Option::is_none")]
150 pub aspect_ratio: Option<AspectRatio>,
151 #[serde(default, skip_serializing_if = "Option::is_none")]
153 pub seed: Option<u64>,
154 #[serde(default, skip_serializing_if = "Vec::is_empty")]
156 pub files: Vec<ImageFile>,
157 #[serde(default, skip_serializing_if = "Option::is_none")]
159 pub mask: Option<ImageFile>,
160 #[serde(default, skip_serializing_if = "ProviderOptions::is_empty")]
162 pub provider_options: ProviderOptions,
163}
164
165#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
167#[serde(tag = "type", rename_all = "lowercase")]
168#[non_exhaustive]
169pub enum BatchRequest {
170 Text {
172 id: String,
174 model_id: ModelId,
176 options: TextBatchRequestOptions,
178 },
179 Image {
181 id: String,
183 model_id: ModelId,
185 options: ImageBatchRequestOptions,
187 },
188}
189
190impl BatchRequest {
191 #[must_use]
193 pub fn id(&self) -> &str {
194 match self {
195 Self::Text { id, .. } | Self::Image { id, .. } => id,
196 }
197 }
198}
199
200#[derive(Debug, Clone)]
202pub struct BatchStartOptions {
203 pub requests: Vec<BatchRequest>,
205 pub webhook_url: Option<Url>,
207 pub provider_options: ProviderOptions,
209 pub headers: Headers,
211 pub cancellation: CancellationToken,
213}
214
215#[derive(Debug, Clone)]
217pub struct BatchOperationOptions {
218 pub batch_id: BatchId,
220 pub provider_options: ProviderOptions,
222 pub headers: Headers,
224 pub cancellation: CancellationToken,
226}
227
228impl BatchOperationOptions {
229 #[must_use]
231 pub fn new(batch_id: impl Into<BatchId>) -> Self {
232 Self {
233 batch_id: batch_id.into(),
234 provider_options: ProviderOptions::new(),
235 headers: Headers::new(),
236 cancellation: CancellationToken::new(),
237 }
238 }
239}
240
241#[derive(Debug, Clone, Default)]
243pub struct BatchListOptions {
244 pub limit: Option<usize>,
246 pub cursor: Option<String>,
248 pub provider_options: ProviderOptions,
250 pub headers: Headers,
252 pub cancellation: CancellationToken,
254}
255
256#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
258#[serde(rename_all = "lowercase")]
259#[non_exhaustive]
260pub enum BatchState {
261 Pending,
263 Completed,
265 Failed,
267}
268
269#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
271pub struct BatchError {
272 pub message: String,
274 #[serde(rename = "type", default, skip_serializing_if = "Option::is_none")]
276 pub error_type: Option<String>,
277 #[serde(default, skip_serializing_if = "Option::is_none")]
279 pub code: Option<String>,
280 #[serde(default, skip_serializing_if = "Option::is_none")]
282 pub status_code: Option<u16>,
283}
284
285#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)]
287pub struct BatchRequestCounts {
288 pub total: u64,
290 pub pending: u64,
292 pub completed: u64,
294 pub failed: u64,
296}
297
298#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
300pub struct BatchStatus {
301 pub status: BatchState,
303 #[serde(default, skip_serializing_if = "Option::is_none")]
305 pub raw_status: Option<String>,
306 #[serde(default, skip_serializing_if = "Option::is_none")]
308 pub request_counts: Option<BatchRequestCounts>,
309 #[serde(default, skip_serializing_if = "Option::is_none")]
311 pub error: Option<BatchError>,
312 #[serde(default, skip_serializing_if = "Option::is_none")]
314 pub created_at: Option<DateTime<Utc>>,
315 #[serde(default, skip_serializing_if = "Option::is_none")]
317 pub expires_at: Option<DateTime<Utc>>,
318 #[serde(default, skip_serializing_if = "Option::is_none")]
320 pub provider_metadata: Option<ProviderMetadata>,
321}
322
323impl BatchStatus {
324 #[must_use]
326 pub fn new(status: BatchState) -> Self {
327 Self {
328 status,
329 raw_status: None,
330 request_counts: None,
331 error: None,
332 created_at: None,
333 expires_at: None,
334 provider_metadata: None,
335 }
336 }
337}
338
339#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
341pub struct BatchWarning {
342 #[serde(default, skip_serializing_if = "Option::is_none")]
344 pub request_id: Option<String>,
345 pub warning: Warning,
347}
348
349#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
351pub struct BatchStartResult {
352 pub batch_id: BatchId,
354 #[serde(flatten)]
356 pub status: BatchStatus,
357 #[serde(default)]
359 pub warnings: Vec<BatchWarning>,
360}
361
362#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
364pub struct BatchCancelResult {
365 #[serde(default, skip_serializing_if = "Option::is_none")]
367 pub provider_metadata: Option<ProviderMetadata>,
368}
369
370#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
372pub struct BatchListItem {
373 pub batch_id: BatchId,
375 #[serde(flatten)]
377 pub status: BatchStatus,
378}
379
380#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
382pub struct BatchListResult {
383 pub batches: Vec<BatchListItem>,
385 #[serde(default, skip_serializing_if = "Option::is_none")]
387 pub next_cursor: Option<String>,
388 #[serde(default, skip_serializing_if = "Option::is_none")]
390 pub provider_metadata: Option<ProviderMetadata>,
391}
392
393#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
395#[serde(tag = "status", rename_all = "lowercase")]
396#[non_exhaustive]
397pub enum BatchItem<R> {
398 Succeeded {
400 id: String,
402 result: R,
404 },
405 Failed {
407 id: String,
409 error: BatchError,
411 #[serde(default, skip_serializing_if = "Option::is_none")]
413 provider_metadata: Option<ProviderMetadata>,
414 },
415 Cancelled {
417 id: String,
419 #[serde(default, skip_serializing_if = "Option::is_none")]
421 error: Option<BatchError>,
422 #[serde(default, skip_serializing_if = "Option::is_none")]
424 provider_metadata: Option<ProviderMetadata>,
425 },
426 Expired {
428 id: String,
430 #[serde(default, skip_serializing_if = "Option::is_none")]
432 error: Option<BatchError>,
433 #[serde(default, skip_serializing_if = "Option::is_none")]
435 provider_metadata: Option<ProviderMetadata>,
436 },
437}
438
439impl<R> BatchItem<R> {
440 #[must_use]
442 pub fn id(&self) -> &str {
443 match self {
444 Self::Succeeded { id, .. }
445 | Self::Failed { id, .. }
446 | Self::Cancelled { id, .. }
447 | Self::Expired { id, .. } => id,
448 }
449 }
450}
451
452#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
454#[serde(tag = "type", rename_all = "lowercase")]
455#[non_exhaustive]
456pub enum BatchItemResult {
457 Text(Box<BatchItem<GenerateResult>>),
459 Image(Box<BatchItem<ImageResult>>),
461}
462
463impl BatchItemResult {
464 #[must_use]
466 pub fn id(&self) -> &str {
467 match self {
468 Self::Text(item) => item.id(),
469 Self::Image(item) => item.id(),
470 }
471 }
472}