taquba-bulk 0.4.0

Bulk multi-step processing on the Taquba durable task queue: run one pipeline over many inputs in parallel, with per-item memoization and cost rollup.
Documentation
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
348
349
350
351
352
353
354
355
356
357
358
359
360
361
362
363
364
365
366
367
368
369
370
371
372
373
374
375
376
377
378
379
380
381
382
383
384
385
386
387
388
389
390
391
392
393
394
395
396
397
398
399
400
401
402
403
404
405
406
407
408
409
410
411
412
413
414
415
416
417
418
419
420
421
422
423
424
425
426
427
428
429
430
431
432
433
434
435
436
437
438
439
440
441
442
443
444
445
446
447
448
449
450
451
452
453
454
455
456
457
458
459
460
461
462
463
464
465
466
467
468
469
470
471
472
473
474
475
476
477
478
479
480
481
482
483
484
485
486
487
488
489
490
491
492
493
494
495
496
497
498
499
500
501
502
503
504
505
506
507
508
509
510
511
512
513
514
515
516
517
518
519
520
521
522
523
524
525
526
527
528
529
530
531
532
533
534
535
536
537
538
539
540
541
542
543
544
545
546
547
548
549
550
551
//! The [`Pipeline`] contract and the per-item [`BulkCtx`] handed to it.

use std::collections::HashMap;
use std::future::Future;

use serde::Serialize;
use serde::de::DeserializeOwned;
use taquba_workflow::{Memo, StepError};
use tokio_util::sync::CancellationToken;

use crate::cost::CostReport;

/// Defines a per-item processing pipeline. Each bulk run executes one
/// `Pipeline` for every input item independently, materialised internally
/// as a [`taquba_workflow`] run.
///
/// A `Pipeline` is a single async [`run`](Pipeline::run) method: the bulk
/// runner deserializes one input item, builds a [`BulkCtx`] around it, and
/// awaits `run`. The expensive logical steps inside `run` (LLM calls, paid
/// APIs, CPU-bound work) are wrapped in [`BulkCtx::memoized`] or
/// [`BulkCtx::memoized_by_content`] so an at-least-once retry of the item
/// replays cached step results instead of paying for them twice.
///
/// # Error classification
///
/// [`Self::Error`] must convert into a [`StepError`], which is what decides
/// retry behaviour: a [`StepError::transient`] error nacks and retries with
/// the queue's backoff up to `max_attempts` (then dead-letters and the item
/// terminates failed); a [`StepError::permanent`] error dead-letters the
/// item immediately. The simplest choice is to use `StepError` directly as
/// `type Error` (as the example below does); otherwise implement
/// `From<YourError> for StepError`.
///
/// # Example
///
/// ```no_run
/// use serde::{Deserialize, Serialize};
/// use taquba_bulk::{BulkCtx, CostReport, Pipeline, StepError};
///
/// #[derive(Serialize, Deserialize)]
/// struct Ticket { id: String, body: String }
///
/// #[derive(Serialize, Deserialize)]
/// struct Processed { id: String, classification: String }
///
/// struct TicketPipeline;
///
/// impl Pipeline for TicketPipeline {
///     type Input = Ticket;
///     type Output = Processed;
///     type Error = StepError;
///
///     async fn run(&self, ctx: &BulkCtx<Ticket>) -> Result<Processed, StepError> {
///         let classification = ctx
///             .memoized_with_cached_cost("classify", async {
///                 let cost = CostReport::new();
///                 cost.record("llm_calls", 1.0);
///                 Ok::<_, StepError>(("billing".to_string(), cost))
///             })
///             .await?;
///         Ok(Processed { id: ctx.input.id.clone(), classification })
///     }
/// }
/// ```
pub trait Pipeline: Send + Sync + 'static {
    /// One input item. Deserialized from the bulk input source and handed to
    /// [`run`](Pipeline::run) via [`BulkCtx::input`].
    type Input: Serialize + DeserializeOwned + Send + 'static;
    /// The per-item result. Serialized into the bulk output stream once the
    /// item completes.
    type Output: Serialize + DeserializeOwned + Send + 'static;
    /// Failure type. Must convert into a [`StepError`] so the runner can
    /// decide transient vs. permanent handling. Use `StepError` directly for
    /// the common case.
    type Error: Into<StepError> + Send + 'static;

    /// Process one input item. Wrap expensive logical steps in
    /// [`BulkCtx::memoized`] or [`BulkCtx::memoized_by_content`] to make
    /// retries cheap.
    fn run(
        &self,
        ctx: &BulkCtx<Self::Input>,
    ) -> impl Future<Output = Result<Self::Output, Self::Error>> + Send;
}

/// Per-item execution context handed to [`Pipeline::run`].
///
/// Wraps the typed input together with the durable per-item
/// [memo](taquba_workflow::Memo), a [cost accumulator](CostReport), and the
/// run's cooperative [cancellation token](CancellationToken).
pub struct BulkCtx<T> {
    /// The deserialized input item for this run.
    pub input: T,
    /// The run identifier for this item (the value the bulk runner derived
    /// from the input, or a positional `item-{i}` default).
    pub run_id: String,
    /// Submitter-supplied metadata threaded through from the bulk run.
    pub headers: HashMap<String, String>,
    memo: Memo,
    cost: CostReport,
    cancel_token: CancellationToken,
}

impl<T> BulkCtx<T> {
    pub(crate) fn new(
        input: T,
        run_id: String,
        headers: HashMap<String, String>,
        memo: Memo,
        cancel_token: CancellationToken,
    ) -> Self {
        Self {
            input,
            run_id,
            headers,
            memo,
            cost: CostReport::new(),
            cancel_token,
        }
    }

    /// Run `f` once and cache its result durably under `key`, or return the
    /// previously cached result on a retry.
    ///
    /// On the first execution of a step, `f` runs and its `Ok` value is
    /// rmp-serialized into the item's [`Memo`] under `key` before being
    /// returned. If the step is later re-executed (at-least-once retry after
    /// a lease expiry), the cached bytes are returned without running `f`
    /// again, so a paid call inside `f` bills once, not once per attempt.
    /// An `Err` from `f` is never cached and propagates unchanged.
    ///
    /// `key` namespaces the cache within this item; use a distinct key per
    /// logical step. A cached entry that fails to deserialize (e.g. an
    /// output type changed shape between runs) is treated as a miss and `f`
    /// re-runs, overwriting it. Memo I/O failures surface as a transient
    /// [`StepError`]; serializing the computed value fails
    /// deterministically, so that surfaces as a permanent [`StepError`].
    /// Both are converted into the caller's error type.
    ///
    /// Calls to [`record_cost`](Self::record_cost) inside `f` run only on a
    /// cache miss. Use [`memoized_with_cached_cost`](Self::memoized_with_cached_cost)
    /// when cached results should also contribute to the final cost report.
    pub async fn memoized<R, F, E>(&self, key: &str, f: F) -> Result<R, E>
    where
        R: Serialize + DeserializeOwned,
        F: Future<Output = Result<R, E>>,
        E: From<StepError>,
    {
        if let Some(value) = self
            .decode_cached::<R>(self.memo.get(key).await, Some(key))
            .map_err(E::from)?
        {
            return Ok(value);
        }

        let value = f.await?;
        let bytes = encode_memo_value(&value).map_err(E::from)?;
        self.memo
            .put(key, &bytes)
            .await
            .map_err(|e| E::from(StepError::from(e)))?;
        Ok(value)
    }

    /// Run `f` once and cache its result under a key derived from
    /// serialized `input`, or return the previously cached result on a
    /// retry.
    ///
    /// This has the same typed compute-on-miss behaviour as
    /// [`memoized`](Self::memoized), but the memo key is derived by
    /// serializing `input` as MessagePack and hashing it with SHA-256 via
    /// [`taquba_workflow::Memo::content_get`] and
    /// [`taquba_workflow::Memo::content_put`]. The entry remains scoped to
    /// this item's workflow run and step; this method does not create a
    /// cross-item cache.
    ///
    /// The derived key is stable only when `input` serializes
    /// deterministically; types with unordered iteration, such as
    /// `HashMap`, can serialize the same logical content into different
    /// bytes and therefore different keys. If several
    /// logical operations may receive the same input shape, include an
    /// operation name in the serialized input.
    pub async fn memoized_by_content<K, R, F, E>(&self, input: &K, f: F) -> Result<R, E>
    where
        K: Serialize + ?Sized,
        R: Serialize + DeserializeOwned,
        F: Future<Output = Result<R, E>>,
        E: From<StepError>,
    {
        if let Some(value) = self
            .decode_cached::<R>(self.memo.content_get(input).await, None)
            .map_err(E::from)?
        {
            return Ok(value);
        }

        let value = f.await?;
        let bytes = encode_memo_value(&value).map_err(E::from)?;
        self.memo
            .content_put(input, &bytes)
            .await
            .map_err(|e| E::from(StepError::from(e)))?;
        Ok(value)
    }

    /// Run `f` once and cache both its value and counters under `key`,
    /// or return the cached value and replay its counters on a retry.
    ///
    /// Use this when cost counters are known only inside a memoized step.
    /// The closure returns `(value, cost)`, and the helper records the
    /// `CostReport` after memoization returns, so counters are included
    /// whether the step computes freshly or hits memo state.
    pub async fn memoized_with_cached_cost<R, F, E>(&self, key: &str, f: F) -> Result<R, E>
    where
        R: Serialize + DeserializeOwned,
        F: Future<Output = Result<(R, CostReport), E>>,
        E: From<StepError>,
    {
        let (value, cost) = self.memoized(key, f).await?;
        self.cost.merge(&cost);
        Ok(value)
    }

    /// Run `f` once and cache both its value and counters under a key
    /// derived from serialized `input`, or return the cached value and
    /// replay its counters on a retry.
    ///
    /// Use this when the memo key should be content-derived and cost
    /// counters are known only inside the memoized step. The closure
    /// returns `(value, cost)`, and the helper records the `CostReport`
    /// after memoization returns, so counters are included whether the
    /// step computes freshly or hits memo state.
    pub async fn memoized_by_content_with_cached_cost<K, R, F, E>(
        &self,
        input: &K,
        f: F,
    ) -> Result<R, E>
    where
        K: Serialize + ?Sized,
        R: Serialize + DeserializeOwned,
        F: Future<Output = Result<(R, CostReport), E>>,
        E: From<StepError>,
    {
        let (value, cost) = self.memoized_by_content(input, f).await?;
        self.cost.merge(&cost);
        Ok(value)
    }

    /// Add `amount` to the cost counter named `metric` for this item. The
    /// per-item totals roll up into the batch-level
    /// [`ProgressSnapshot`](crate::ProgressSnapshot) and
    /// [`BulkReport`](crate::BulkReport).
    pub fn record_cost(&self, metric: &str, amount: f64) {
        self.cost.record(metric, amount);
    }

    /// The run's cooperative cancellation token. Watch it to short-circuit a
    /// long-running step when the bulk run is draining (e.g. on spot
    /// preemption); see [`taquba_workflow::Step::cancel_token`].
    pub fn cancel_token(&self) -> &CancellationToken {
        &self.cancel_token
    }

    /// Snapshot of the cost accumulated so far for this item.
    pub(crate) fn cost(&self) -> CostReport {
        self.cost.clone()
    }

    /// Decode a cached memo entry: `Ok(Some(value))` on a hit,
    /// `Ok(None)` on a miss, and `Err` when the memo read itself
    /// failed. An entry that fails to deserialize is treated as a miss
    /// (self-healing: the caller recomputes and overwrites it), logged
    /// with `key` when the entry is key-addressed.
    fn decode_cached<R: DeserializeOwned>(
        &self,
        cached: taquba_workflow::Result<Option<Vec<u8>>>,
        key: Option<&str>,
    ) -> Result<Option<R>, StepError> {
        let bytes = match cached {
            Ok(Some(bytes)) => bytes,
            Ok(None) => return Ok(None),
            Err(err) => return Err(err.into()),
        };
        match rmp_serde::from_slice::<R>(&bytes) {
            Ok(value) => Ok(Some(value)),
            Err(err) => {
                match key {
                    Some(key) => tracing::warn!(
                        run_id = %self.run_id,
                        key = %key,
                        error = %err,
                        "memoized cache entry failed to deserialize; recomputing",
                    ),
                    None => tracing::warn!(
                        run_id = %self.run_id,
                        error = %err,
                        "content-addressed memoized cache entry failed to deserialize; recomputing",
                    ),
                }
                Ok(None)
            }
        }
    }
}

/// Serialize a computed value for the memo. A serialization failure is
/// deterministic, so a retry produces the same error and the failure
/// is classified permanent.
fn encode_memo_value<R: Serialize>(value: &R) -> Result<Vec<u8>, StepError> {
    rmp_serde::to_vec_named(value)
        .map_err(|e| StepError::permanent(format!("memo serialize failed: {e}")))
}

#[cfg(test)]
mod tests {
    use super::*;
    use std::sync::Arc;
    use std::sync::atomic::{AtomicU32, Ordering};
    use taquba::object_store::memory::InMemory;
    use taquba_workflow::MemoStore;

    #[derive(Serialize)]
    struct ContentInput<'a> {
        operation: &'static str,
        payload: &'a [u8],
    }

    fn ctx_for_tests() -> BulkCtx<()> {
        let memo = MemoStore::new(Arc::new(InMemory::new()), "memo").new_memo("run-1", 0);
        BulkCtx::new(
            (),
            "run-1".into(),
            HashMap::new(),
            memo,
            CancellationToken::new(),
        )
    }

    #[tokio::test]
    async fn memoized_runs_once_then_serves_cache() {
        let ctx = ctx_for_tests();
        let calls = AtomicU32::new(0);
        let compute = || async {
            calls.fetch_add(1, Ordering::SeqCst);
            Ok::<_, StepError>(7u32)
        };

        let first = ctx.memoized("k", compute()).await.unwrap();
        let second = ctx.memoized("k", compute()).await.unwrap();
        assert_eq!(first, 7);
        assert_eq!(second, 7);
        assert_eq!(calls.load(Ordering::SeqCst), 1, "closure ran exactly once");
    }

    #[tokio::test]
    async fn memoized_does_not_cache_errors() {
        let ctx = ctx_for_tests();
        let calls = AtomicU32::new(0);

        let err = ctx
            .memoized::<u32, _, StepError>("k", async {
                calls.fetch_add(1, Ordering::SeqCst);
                Err(StepError::transient("boom"))
            })
            .await
            .unwrap_err();
        assert_eq!(err.message, "boom");

        // A second attempt re-runs because the error was not cached.
        let ok = ctx
            .memoized::<u32, _, StepError>("k", async {
                calls.fetch_add(1, Ordering::SeqCst);
                Ok(99)
            })
            .await
            .unwrap();
        assert_eq!(ok, 99);
        assert_eq!(calls.load(Ordering::SeqCst), 2);
    }

    #[tokio::test]
    async fn distinct_keys_cache_independently() {
        let ctx = ctx_for_tests();
        let a = ctx
            .memoized::<String, _, StepError>("a", async { Ok("first".to_string()) })
            .await
            .unwrap();
        let b = ctx
            .memoized::<String, _, StepError>("b", async { Ok("second".to_string()) })
            .await
            .unwrap();
        assert_eq!(a, "first");
        assert_eq!(b, "second");
    }

    #[tokio::test]
    async fn memoized_by_content_runs_once_then_serves_cache() {
        let ctx = ctx_for_tests();
        let calls = AtomicU32::new(0);
        let input = ContentInput {
            operation: "classify",
            payload: b"ticket",
        };
        let compute = || async {
            calls.fetch_add(1, Ordering::SeqCst);
            Ok::<_, StepError>("billing".to_string())
        };

        let first = ctx.memoized_by_content(&input, compute()).await.unwrap();
        let second = ctx.memoized_by_content(&input, compute()).await.unwrap();

        assert_eq!(first, "billing");
        assert_eq!(second, "billing");
        assert_eq!(calls.load(Ordering::SeqCst), 1, "closure ran exactly once");
    }

    #[tokio::test]
    async fn memoized_by_content_distinguishes_serialized_inputs() {
        let ctx = ctx_for_tests();
        let classify = ContentInput {
            operation: "classify",
            payload: b"ticket",
        };
        let summarize = ContentInput {
            operation: "summarize",
            payload: b"ticket",
        };

        let first = ctx
            .memoized_by_content::<_, String, _, StepError>(&classify, async {
                Ok("class-a".to_string())
            })
            .await
            .unwrap();
        let second = ctx
            .memoized_by_content::<_, String, _, StepError>(&summarize, async {
                Ok("summary".to_string())
            })
            .await
            .unwrap();

        assert_eq!(first, "class-a");
        assert_eq!(second, "summary");
    }

    #[tokio::test]
    async fn memoized_with_cached_cost_records_cost_on_compute_and_memo_hit() {
        let memo = MemoStore::new(Arc::new(InMemory::new()), "memo").new_memo("run-1", 0);
        let first_ctx = BulkCtx::new(
            (),
            "run-1".into(),
            HashMap::new(),
            memo.clone(),
            CancellationToken::new(),
        );
        let replay_ctx = BulkCtx::new(
            (),
            "run-1".into(),
            HashMap::new(),
            memo,
            CancellationToken::new(),
        );
        let calls = AtomicU32::new(0);

        let first = first_ctx
            .memoized_with_cached_cost("k", async {
                calls.fetch_add(1, Ordering::SeqCst);
                let cost = CostReport::new();
                cost.record("tokens", 42.0);
                Ok::<_, StepError>(("value".to_string(), cost))
            })
            .await
            .unwrap();
        let second = replay_ctx
            .memoized_with_cached_cost("k", async {
                calls.fetch_add(1, Ordering::SeqCst);
                Ok::<_, StepError>(("other".to_string(), CostReport::new()))
            })
            .await
            .unwrap();

        assert_eq!(first, "value");
        assert_eq!(second, "value");
        assert_eq!(
            calls.load(Ordering::SeqCst),
            1,
            "memo hit did not run closure"
        );
        assert_eq!(first_ctx.cost().get("tokens"), 42.0);
        assert_eq!(replay_ctx.cost().get("tokens"), 42.0);
    }

    #[tokio::test]
    async fn memoized_by_content_with_cached_cost_records_cost_on_compute_and_memo_hit() {
        let memo = MemoStore::new(Arc::new(InMemory::new()), "memo").new_memo("run-1", 0);
        let first_ctx = BulkCtx::new(
            (),
            "run-1".into(),
            HashMap::new(),
            memo.clone(),
            CancellationToken::new(),
        );
        let replay_ctx = BulkCtx::new(
            (),
            "run-1".into(),
            HashMap::new(),
            memo,
            CancellationToken::new(),
        );
        let calls = AtomicU32::new(0);
        let input = ContentInput {
            operation: "classify",
            payload: b"ticket",
        };

        let first = first_ctx
            .memoized_by_content_with_cached_cost(&input, async {
                calls.fetch_add(1, Ordering::SeqCst);
                let cost = CostReport::new();
                cost.record("tokens", 42.0);
                Ok::<_, StepError>(("value".to_string(), cost))
            })
            .await
            .unwrap();
        let second = replay_ctx
            .memoized_by_content_with_cached_cost(&input, async {
                calls.fetch_add(1, Ordering::SeqCst);
                Ok::<_, StepError>(("other".to_string(), CostReport::new()))
            })
            .await
            .unwrap();

        assert_eq!(first, "value");
        assert_eq!(second, "value");
        assert_eq!(
            calls.load(Ordering::SeqCst),
            1,
            "memo hit did not run closure"
        );
        assert_eq!(first_ctx.cost().get("tokens"), 42.0);
        assert_eq!(replay_ctx.cost().get("tokens"), 42.0);
    }

    #[tokio::test]
    async fn record_cost_accumulates_into_snapshot() {
        let ctx = ctx_for_tests();
        ctx.record_cost("tokens", 100.0);
        ctx.record_cost("tokens", 50.0);
        assert_eq!(ctx.cost().get("tokens"), 150.0);
    }
}