Skip to main content

ferrin_core/batch/
start.rs

1//! `start_batch`.
2
3use std::collections::HashMap;
4use std::fmt;
5use std::future::IntoFuture;
6use std::sync::Arc;
7
8use ferrin_spec::BatchRef;
9use ferrin_spec::BoxFuture;
10use ferrin_spec::ModelId;
11use ferrin_spec::ToolDefinition;
12use ferrin_spec::ToolName;
13use ferrin_spec::batch::BatchRequest as ModelBatchRequest;
14use ferrin_spec::batch::BatchStartOptions;
15use ferrin_spec::batch::BatchStartResult;
16use ferrin_spec::batch::ImageBatchRequestOptions;
17use ferrin_spec::batch::TextBatchRequestOptions;
18use tracing::Instrument;
19use url::Url;
20
21use crate::error::Error;
22use crate::modality::ModalityOptions;
23use crate::modality::impl_modality_builder;
24use crate::prompt::DownloadFn;
25use crate::prompt::convert::ConvertContext;
26use crate::prompt::convert::convert_to_prompt;
27use crate::prompt::prepare_tools::PrepareToolsInput;
28use crate::prompt::prepare_tools::prepare_tools;
29use crate::prompt::standardize::standardize;
30use crate::telemetry::ModelIdentity;
31
32use super::request::BatchRequest;
33use super::request::validate_compatible_tools;
34use super::request::validate_requests;
35use super::service_identity;
36use crate::telemetry::spans;
37
38/// Starts a batch.
39#[must_use]
40pub fn start_batch(
41    batch: impl Into<BatchRef>,
42    requests: impl IntoIterator<Item = impl Into<BatchRequest>>,
43) -> StartBatch {
44    StartBatch {
45        batch: batch.into(),
46        requests: requests.into_iter().map(Into::into).collect(),
47        webhook_url: None,
48        download: None,
49        base: ModalityOptions::default(),
50    }
51}
52
53/// Builder returned by [`start_batch`]; `.await` submits the batch.
54pub struct StartBatch {
55    batch: BatchRef,
56    requests: Vec<BatchRequest>,
57    webhook_url: Option<Url>,
58    download: Option<Arc<dyn DownloadFn>>,
59    base: ModalityOptions,
60}
61
62impl fmt::Debug for StartBatch {
63    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
64        f.debug_struct("StartBatch")
65            .field("batch", &self.batch)
66            .field("requests", &self.requests.len())
67            .field("webhook_url", &self.webhook_url)
68            .field("has_download", &self.download.is_some())
69            .field("base", &self.base)
70            .finish()
71    }
72}
73
74impl StartBatch {
75    /// URL the provider calls when the batch completes.
76    #[must_use]
77    pub fn webhook_url(mut self, url: Url) -> Self {
78        self.webhook_url = Some(url);
79        self
80    }
81
82    /// Sets the function used to fetch prompt URLs the provider cannot.
83    #[must_use]
84    pub fn download(mut self, download: Arc<dyn DownloadFn>) -> Self {
85        self.download = Some(download);
86        self
87    }
88}
89
90impl_modality_builder!(@no_retry StartBatch);
91
92impl IntoFuture for StartBatch {
93    type Output = Result<BatchStartResult, Error>;
94    type IntoFuture = BoxFuture<'static, Self::Output>;
95
96    fn into_future(self) -> Self::IntoFuture {
97        Box::pin(run_start(self))
98    }
99}
100
101async fn run_start(builder: StartBatch) -> Result<BatchStartResult, Error> {
102    validate_requests(&builder.requests)?;
103    let identity = service_identity(&builder.batch);
104    let span = spans::modality_span("start_batch", &identity);
105    let base = builder.base.clone();
106    base.run(|base, token| {
107        async move {
108            let supported_urls = builder.batch.supported_urls().await;
109            let mut normalized: Vec<ModelBatchRequest> = Vec::with_capacity(builder.requests.len());
110            let mut model_ids: HashMap<String, ModelId> = HashMap::new();
111            let mut seen_tools: HashMap<ToolName, ToolDefinition> = HashMap::new();
112            for request in builder.requests {
113                if token.is_cancelled() {
114                    return Err(Error::Cancelled);
115                }
116                match request {
117                    BatchRequest::Text(request) => {
118                        let request = *request;
119                        let standardized = standardize(
120                            request.system,
121                            request.prompt,
122                            request.messages,
123                            request.allow_system_in_messages,
124                        )?;
125                        request.settings.validate()?;
126                        let prepared = prepare_tools(PrepareToolsInput {
127                            tools: &request.tools,
128                            active_tools: request.active_tools.as_deref(),
129                            tool_order: &request.tool_order,
130                            tool_choice: request.tool_choice,
131                            tools_context: request.tools_context.as_ref(),
132                            #[cfg(feature = "sandbox")]
133                            sandbox: None,
134                        })
135                        .await?;
136                        validate_compatible_tools(
137                            &request.id,
138                            &prepared.definitions,
139                            &mut seen_tools,
140                        )?;
141                        let prompt = convert_to_prompt(
142                            standardized.system.as_ref(),
143                            &standardized.messages,
144                            ConvertContext {
145                                supported_urls: &supported_urls,
146                                download: builder.download.as_deref(),
147                                cancellation: &token,
148                            },
149                        )
150                        .await?;
151                        let settings = request.settings;
152                        model_ids.insert(request.id.clone(), request.model_id.clone());
153                        normalized.push(ModelBatchRequest::Text {
154                            id: request.id,
155                            model_id: request.model_id,
156                            options: TextBatchRequestOptions {
157                                prompt,
158                                max_output_tokens: settings.max_output_tokens,
159                                temperature: settings.temperature,
160                                stop_sequences: settings.stop_sequences,
161                                top_p: settings.top_p,
162                                top_k: settings.top_k,
163                                presence_penalty: settings.presence_penalty,
164                                frequency_penalty: settings.frequency_penalty,
165                                seed: settings.seed,
166                                reasoning: settings.reasoning,
167                                response_format: request.response_format,
168                                tool_choice: prepared.tool_choice,
169                                tools: prepared.definitions,
170                                provider_options: settings.provider_options,
171                            },
172                        });
173                    }
174                    BatchRequest::Image(request) => {
175                        let request = *request;
176                        if request.n == 0 {
177                            return Err(Error::invalid_argument("n", "must be at least 1"));
178                        }
179                        model_ids.insert(request.id.clone(), request.model_id.clone());
180                        normalized.push(ModelBatchRequest::Image {
181                            id: request.id,
182                            model_id: request.model_id,
183                            options: ImageBatchRequestOptions {
184                                prompt: request.prompt,
185                                n: request.n,
186                                size: request.size,
187                                aspect_ratio: request.aspect_ratio,
188                                seed: request.seed,
189                                files: request.files,
190                                mask: request.mask,
191                                provider_options: request.provider_options,
192                            },
193                        });
194                    }
195                    #[allow(unreachable_patterns, reason = "BatchRequest is non-exhaustive")]
196                    _ => {
197                        return Err(Error::invalid_argument(
198                            "requests",
199                            "unsupported batch request type",
200                        ));
201                    }
202                }
203            }
204            let result = builder
205                .batch
206                .do_start_batch(BatchStartOptions {
207                    requests: normalized,
208                    webhook_url: builder.webhook_url,
209                    provider_options: base.provider_options.clone(),
210                    headers: base.request_headers(),
211                    cancellation: token,
212                })
213                .await
214                .map_err(Error::from)?;
215            for warning in &result.warnings {
216                let model_id = warning
217                    .request_id
218                    .as_ref()
219                    .and_then(|id| model_ids.get(id))
220                    .cloned()
221                    .unwrap_or_else(|| ModelId::new("batch"));
222                spans::log_warnings(
223                    std::slice::from_ref(&warning.warning),
224                    &ModelIdentity::new(identity.provider.clone(), model_id),
225                );
226            }
227            Ok(result)
228        }
229        .instrument(span)
230    })
231    .await
232}