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 model_ids.insert(request.id.clone(), request.model_id.clone());
178 normalized.push(ModelBatchRequest::Image {
179 id: request.id,
180 model_id: request.model_id,
181 options: ImageBatchRequestOptions {
182 prompt: request.prompt,
183 n: request.n,
184 size: request.size,
185 aspect_ratio: request.aspect_ratio,
186 seed: request.seed,
187 files: request.files,
188 mask: request.mask,
189 provider_options: request.provider_options,
190 },
191 });
192 }
193 #[allow(unreachable_patterns, reason = "BatchRequest is non-exhaustive")]
194 _ => {
195 return Err(Error::invalid_argument(
196 "requests",
197 "unsupported batch request type",
198 ));
199 }
200 }
201 }
202 let result = builder
203 .batch
204 .do_start_batch(BatchStartOptions {
205 requests: normalized,
206 webhook_url: builder.webhook_url,
207 provider_options: base.provider_options.clone(),
208 headers: base.request_headers(),
209 cancellation: token,
210 })
211 .await
212 .map_err(Error::from)?;
213 for warning in &result.warnings {
214 let model_id = warning
215 .request_id
216 .as_ref()
217 .and_then(|id| model_ids.get(id))
218 .cloned()
219 .unwrap_or_else(|| ModelId::new("batch"));
220 spans::log_warnings(
221 std::slice::from_ref(&warning.warning),
222 &ModelIdentity::new(identity.provider.clone(), model_id),
223 );
224 }
225 Ok(result)
226 }
227 .instrument(span)
228 })
229 .await
230}