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 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}