1use async_trait::async_trait;
18use std::marker::PhantomData;
19use std::sync::Arc;
20
21use crate::service::{ServiceError, ServiceResult};
22
23#[derive(Debug, Clone, Copy, PartialEq, Eq)]
27pub enum BulkFailureMode {
28 ContinueOnError,
30 AbortOnFirstError,
32}
33
34impl Default for BulkFailureMode {
35 fn default() -> Self {
36 Self::ContinueOnError
37 }
38}
39
40#[derive(Debug, Clone)]
42pub struct BulkOperationConfig {
43 pub max_batch_size: usize,
45 pub failure_mode: BulkFailureMode,
47}
48
49impl Default for BulkOperationConfig {
50 fn default() -> Self {
51 Self {
52 max_batch_size: 100,
53 failure_mode: BulkFailureMode::ContinueOnError,
54 }
55 }
56}
57
58#[derive(Debug, Clone)]
62pub struct BulkItemResult<E> {
63 pub index: usize,
64 pub result: Result<E, String>,
65}
66
67impl<E> BulkItemResult<E> {
68 pub fn ok(index: usize, entity: E) -> Self {
69 Self {
70 index,
71 result: Ok(entity),
72 }
73 }
74
75 pub fn err(index: usize, error: impl Into<String>) -> Self {
76 Self {
77 index,
78 result: Err(error.into()),
79 }
80 }
81}
82
83#[derive(Debug, Clone)]
85pub struct BulkOperationResult<E> {
86 pub succeeded: Vec<E>,
87 pub failed: Vec<(usize, String)>,
88 pub total: usize,
89}
90
91impl<E: Clone> BulkOperationResult<E> {
92 pub fn new() -> Self {
93 Self {
94 succeeded: Vec::new(),
95 failed: Vec::new(),
96 total: 0,
97 }
98 }
99
100 pub fn success_count(&self) -> usize {
101 self.succeeded.len()
102 }
103
104 pub fn failure_count(&self) -> usize {
105 self.failed.len()
106 }
107
108 pub fn is_fully_successful(&self) -> bool {
109 self.failed.is_empty()
110 }
111
112 pub fn into_item_results(self) -> Vec<BulkItemResult<E>> {
117 let mut items: Vec<BulkItemResult<E>> = self
118 .succeeded
119 .into_iter()
120 .enumerate()
121 .map(|(i, e)| BulkItemResult::ok(i, e))
122 .collect();
123
124 for (index, reason) in self.failed {
125 items.push(BulkItemResult::err(index, reason));
126 }
127
128 items.sort_by_key(|item| item.index);
129 items
130 }
131}
132
133impl<E: Clone> Default for BulkOperationResult<E> {
134 fn default() -> Self {
135 Self::new()
136 }
137}
138
139#[derive(Debug, Clone, Default)]
145pub struct BulkOperationProgress {
146 pub total: usize,
148 pub processed: usize,
150 pub succeeded: usize,
152 pub failed: usize,
154}
155
156impl BulkOperationProgress {
157 pub fn new(total: usize) -> Self {
158 Self {
159 total,
160 processed: 0,
161 succeeded: 0,
162 failed: 0,
163 }
164 }
165
166 pub fn record_success(&mut self) {
167 self.processed += 1;
168 self.succeeded += 1;
169 }
170
171 pub fn record_failure(&mut self) {
172 self.processed += 1;
173 self.failed += 1;
174 }
175
176 pub fn is_complete(&self) -> bool {
177 self.processed >= self.total
178 }
179
180 pub fn success_rate(&self) -> f64 {
181 if self.total == 0 {
182 return 1.0;
183 }
184 self.succeeded as f64 / self.total as f64
185 }
186}
187
188#[async_trait]
192pub trait BulkCapableService<E: Send + Sync + 'static, DTO: Send + Sync + 'static>:
193 Send + Sync
194{
195 async fn create_one(&self, dto: DTO) -> ServiceResult<E>;
196 async fn delete_one(&self, id: &str) -> ServiceResult<bool>;
197}
198
199pub struct GenericBulkService<E, DTO, S> {
207 service: Arc<S>,
208 config: BulkOperationConfig,
209 _phantom: PhantomData<(E, DTO)>,
210}
211
212impl<E, DTO, S> GenericBulkService<E, DTO, S>
213where
214 E: Send + Sync + Clone + 'static,
215 DTO: Send + Sync + 'static,
216 S: BulkCapableService<E, DTO>,
217{
218 pub fn new(service: Arc<S>) -> Self {
219 Self {
220 service,
221 config: BulkOperationConfig::default(),
222 _phantom: PhantomData,
223 }
224 }
225
226 pub fn with_config(service: Arc<S>, config: BulkOperationConfig) -> Self {
227 Self {
228 service,
229 config,
230 _phantom: PhantomData,
231 }
232 }
233
234 pub async fn bulk_create(
239 &self,
240 items: Vec<DTO>,
241 ) -> ServiceResult<(BulkOperationResult<E>, BulkOperationProgress)> {
242 if items.len() > self.config.max_batch_size {
243 return Err(ServiceError::Validation(format!(
244 "bulk create exceeds maximum batch size of {}",
245 self.config.max_batch_size
246 )));
247 }
248
249 let total = items.len();
250 let mut result = BulkOperationResult::new();
251 result.total = total;
252 let mut progress = BulkOperationProgress::new(total);
253
254 for (index, dto) in items.into_iter().enumerate() {
255 match self.service.create_one(dto).await {
256 Ok(entity) => {
257 progress.record_success();
258 result.succeeded.push(entity);
259 }
260 Err(e) => {
261 progress.record_failure();
262 result.failed.push((index, e.to_string()));
263 if self.config.failure_mode == BulkFailureMode::AbortOnFirstError {
264 return Ok((result, progress));
265 }
266 }
267 }
268 }
269
270 Ok((result, progress))
271 }
272
273 pub async fn bulk_delete(
275 &self,
276 ids: Vec<String>,
277 ) -> ServiceResult<(BulkOperationResult<()>, BulkOperationProgress)> {
278 if ids.len() > self.config.max_batch_size {
279 return Err(ServiceError::Validation(format!(
280 "bulk delete exceeds maximum batch size of {}",
281 self.config.max_batch_size
282 )));
283 }
284
285 let total = ids.len();
286 let mut result: BulkOperationResult<()> = BulkOperationResult::new();
287 result.total = total;
288 let mut progress = BulkOperationProgress::new(total);
289
290 for (index, id) in ids.iter().enumerate() {
291 match self.service.delete_one(id).await {
292 Ok(_) => {
293 progress.record_success();
294 result.succeeded.push(());
295 }
296 Err(e) => {
297 progress.record_failure();
298 result.failed.push((index, e.to_string()));
299 if self.config.failure_mode == BulkFailureMode::AbortOnFirstError {
300 return Ok((result, progress));
301 }
302 }
303 }
304 }
305
306 Ok((result, progress))
307 }
308}
309
310#[cfg(test)]
311mod tests {
312 use super::*;
313
314 #[derive(Debug, Clone)]
315 struct Item {
316 id: String,
317 }
318
319 struct CreateItemDto {
320 id: String,
321 }
322
323 struct FakeService {
324 fail_ids: Vec<String>,
325 }
326
327 #[async_trait]
328 impl BulkCapableService<Item, CreateItemDto> for FakeService {
329 async fn create_one(&self, dto: CreateItemDto) -> ServiceResult<Item> {
330 if self.fail_ids.contains(&dto.id) {
331 Err(ServiceError::Internal("injected failure".into()))
332 } else {
333 Ok(Item { id: dto.id })
334 }
335 }
336
337 async fn delete_one(&self, id: &str) -> ServiceResult<bool> {
338 if self.fail_ids.contains(&id.to_string()) {
339 Err(ServiceError::NotFound)
340 } else {
341 Ok(true)
342 }
343 }
344 }
345
346 #[tokio::test]
347 async fn bulk_create_collects_errors_in_continue_mode() {
348 let service = Arc::new(FakeService {
349 fail_ids: vec!["bad".into()],
350 });
351 let bulk = GenericBulkService::new(service);
352
353 let dtos = vec![
354 CreateItemDto { id: "ok1".into() },
355 CreateItemDto { id: "bad".into() },
356 CreateItemDto { id: "ok2".into() },
357 ];
358
359 let (result, progress) = bulk.bulk_create(dtos).await.unwrap();
360 assert_eq!(result.success_count(), 2);
361 assert_eq!(result.failure_count(), 1);
362 assert_eq!(progress.succeeded, 2);
363 assert_eq!(progress.failed, 1);
364 }
365
366 #[tokio::test]
367 async fn bulk_create_aborts_on_first_error() {
368 let service = Arc::new(FakeService {
369 fail_ids: vec!["bad".into()],
370 });
371 let config = BulkOperationConfig {
372 max_batch_size: 100,
373 failure_mode: BulkFailureMode::AbortOnFirstError,
374 };
375 let bulk = GenericBulkService::with_config(service, config);
376
377 let dtos = vec![
378 CreateItemDto { id: "bad".into() },
379 CreateItemDto { id: "ok".into() },
380 ];
381
382 let (result, progress) = bulk.bulk_create(dtos).await.unwrap();
383 assert_eq!(result.success_count(), 0);
385 assert_eq!(result.failure_count(), 1);
386 assert!(progress.is_complete() == false); }
388
389 #[tokio::test]
390 async fn bulk_create_rejects_oversized_batch() {
391 let service = Arc::new(FakeService { fail_ids: vec![] });
392 let config = BulkOperationConfig {
393 max_batch_size: 2,
394 failure_mode: BulkFailureMode::ContinueOnError,
395 };
396 let bulk = GenericBulkService::with_config(service, config);
397
398 let dtos: Vec<_> = (0..3)
399 .map(|i| CreateItemDto {
400 id: i.to_string(),
401 })
402 .collect();
403
404 assert!(bulk.bulk_create(dtos).await.is_err()); }
406}