1use 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#[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
53pub 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 #[must_use]
77 pub fn webhook_url(mut self, url: Url) -> Self {
78 self.webhook_url = Some(url);
79 self
80 }
81
82 #[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 cache: None,
148 cancellation: &token,
149 },
150 )
151 .await?;
152 let settings = request.settings;
153 model_ids.insert(request.id.clone(), request.model_id.clone());
154 normalized.push(ModelBatchRequest::Text {
155 id: request.id,
156 model_id: request.model_id,
157 options: TextBatchRequestOptions {
158 prompt,
159 max_output_tokens: settings.max_output_tokens,
160 temperature: settings.temperature,
161 stop_sequences: settings.stop_sequences,
162 top_p: settings.top_p,
163 top_k: settings.top_k,
164 presence_penalty: settings.presence_penalty,
165 frequency_penalty: settings.frequency_penalty,
166 seed: settings.seed,
167 reasoning: settings.reasoning,
168 response_format: request.response_format,
169 tool_choice: prepared.tool_choice,
170 tools: prepared.definitions,
171 provider_options: settings.provider_options,
172 },
173 });
174 }
175 BatchRequest::Image(request) => {
176 let request = *request;
177 if request.n == 0 {
178 return Err(Error::invalid_argument("n", "must be at least 1"));
179 }
180 model_ids.insert(request.id.clone(), request.model_id.clone());
181 normalized.push(ModelBatchRequest::Image {
182 id: request.id,
183 model_id: request.model_id,
184 options: ImageBatchRequestOptions {
185 prompt: request.prompt,
186 n: request.n,
187 size: request.size,
188 aspect_ratio: request.aspect_ratio,
189 seed: request.seed,
190 files: request.files,
191 mask: request.mask,
192 provider_options: request.provider_options,
193 },
194 });
195 }
196 #[allow(unreachable_patterns, reason = "BatchRequest is non-exhaustive")]
197 _ => {
198 return Err(Error::invalid_argument(
199 "requests",
200 "unsupported batch request type",
201 ));
202 }
203 }
204 }
205 let result = builder
206 .batch
207 .do_start_batch(BatchStartOptions {
208 requests: normalized,
209 webhook_url: builder.webhook_url,
210 provider_options: base.provider_options.clone(),
211 headers: base.request_headers(),
212 cancellation: token,
213 })
214 .await
215 .map_err(Error::from)?;
216 for warning in &result.warnings {
217 let model_id = warning
218 .request_id
219 .as_ref()
220 .and_then(|id| model_ids.get(id))
221 .cloned()
222 .unwrap_or_else(|| ModelId::new("batch"));
223 spans::log_warnings(
224 std::slice::from_ref(&warning.warning),
225 &ModelIdentity::new(identity.provider.clone(), model_id),
226 );
227 }
228 Ok(result)
229 }
230 .instrument(span)
231 })
232 .await
233}