Skip to main content

taquba_workflow/bulk/
pipeline.rs

1//! The [`Pipeline`] contract and the per-item [`BulkCtx`] handed to it.
2
3use std::collections::HashMap;
4use std::future::Future;
5
6use crate::{EffectsHandle, KvReadHandle, Memo, Step, StepError};
7use serde::Serialize;
8use serde::de::DeserializeOwned;
9use taquba::LeaseHandle;
10use tokio_util::sync::CancellationToken;
11
12use crate::bulk::cost::CostReport;
13
14/// Defines a per-item processing pipeline. Each bulk run executes one
15/// `Pipeline` for every input item independently, materialised internally
16/// as a [`crate`] run.
17///
18/// A `Pipeline` is a single async [`run`](Pipeline::run) method: the bulk
19/// runner deserializes one input item, builds a [`BulkCtx`] around it, and
20/// awaits `run`. The expensive logical steps inside `run` (LLM calls, paid
21/// APIs, CPU-bound work) are wrapped in [`Memo::memoized`] or
22/// [`Memo::memoized_by_content`] on [`BulkCtx::memo`], so an at-least-once
23/// retry of the item reads the completed steps back and pays for none of
24/// them twice.
25///
26/// # Error classification
27///
28/// [`Self::Error`] must convert into a [`StepError`], which is what decides
29/// retry behaviour: a [`StepError::transient`] error nacks and retries with
30/// the queue's backoff up to `max_attempts` (then dead-letters and the item
31/// terminates failed); a [`StepError::permanent`] error dead-letters the
32/// item immediately. The simplest choice is to use `StepError` directly as
33/// `type Error` (as the example below does); otherwise implement
34/// `From<YourError> for StepError`.
35///
36/// # Example
37///
38/// ```no_run
39/// use serde::{Deserialize, Serialize};
40/// use taquba_workflow::bulk::{BulkCtx, CostReport, Pipeline};
41/// use taquba_workflow::StepError;
42///
43/// #[derive(Serialize, Deserialize)]
44/// struct Ticket { id: String, body: String }
45///
46/// #[derive(Serialize, Deserialize)]
47/// struct Processed { id: String, classification: String }
48///
49/// struct TicketPipeline;
50///
51/// impl Pipeline for TicketPipeline {
52///     type Input = Ticket;
53///     type Output = Processed;
54///     type Error = StepError;
55///
56///     async fn run(&self, ctx: &BulkCtx<Ticket>) -> Result<Processed, StepError> {
57///         let classification = ctx
58///             .memoized_with_cached_cost("classify", async {
59///                 let cost = CostReport::new();
60///                 cost.record("llm_calls", 1.0);
61///                 Ok::<_, StepError>(("billing".to_string(), cost))
62///             })
63///             .await?;
64///         Ok(Processed { id: ctx.input.id.clone(), classification })
65///     }
66/// }
67/// ```
68pub trait Pipeline: Send + Sync + 'static {
69    /// One input item. Deserialized from the bulk input source and handed to
70    /// [`run`](Pipeline::run) via [`BulkCtx::input`].
71    type Input: Serialize + DeserializeOwned + Send + 'static;
72    /// The per-item result. Serialized into the bulk output stream once the
73    /// item completes.
74    type Output: Serialize + DeserializeOwned + Send + 'static;
75    /// Failure type. Must convert into a [`StepError`] so the runner can
76    /// decide transient vs. permanent handling. Use `StepError` directly for
77    /// the common case.
78    type Error: Into<StepError> + Send + 'static;
79
80    /// Process one input item. Wrap expensive logical steps in
81    /// [`Memo::memoized`] or [`Memo::memoized_by_content`] on
82    /// [`BulkCtx::memo`] to make retries cheap.
83    fn run(
84        &self,
85        ctx: &BulkCtx<Self::Input>,
86    ) -> impl Future<Output = Result<Self::Output, Self::Error>> + Send;
87}
88
89/// Per-item execution context handed to [`Pipeline::run`].
90///
91/// Wraps the typed input together with the durable per-item
92/// [memo](crate::Memo), a [cost accumulator](CostReport), the
93/// run's cooperative [cancellation token](CancellationToken), the
94/// delivery's [lease handle](LeaseHandle), the item's staged
95/// [KV effects](EffectsHandle) and read access to the caller KV
96/// namespace.
97pub struct BulkCtx<T> {
98    /// The deserialized input item for this run.
99    pub input: T,
100    /// The batch this item belongs to.
101    pub batch_id: String,
102    /// The item's key: the value of
103    /// [`BulkBuilder::key_fn`](crate::bulk::BulkBuilder::key_fn) for the
104    /// input, or the positional `item-{i}` default.
105    pub key: String,
106    /// The workflow run identifier of this item, derived from the batch id
107    /// and the key.
108    pub run_id: String,
109    /// Submitter-supplied metadata threaded through from the bulk run.
110    pub headers: HashMap<String, String>,
111    memo: Memo,
112    cost: CostReport,
113    cancel_token: CancellationToken,
114    lease: LeaseHandle,
115    effects: EffectsHandle,
116    kv: KvReadHandle,
117}
118
119impl<T> BulkCtx<T> {
120    pub(crate) fn new(batch_id: &str, key: &str, input: T, step: &Step) -> Self {
121        Self {
122            input,
123            batch_id: batch_id.to_string(),
124            key: key.to_string(),
125            run_id: step.run_id.clone(),
126            headers: step.headers.clone(),
127            memo: step.memo.clone(),
128            cost: CostReport::new(),
129            cancel_token: step.cancel_token.clone(),
130            lease: step.lease.clone(),
131            effects: step.effects.clone(),
132            kv: step.kv.clone(),
133        }
134    }
135
136    /// The item's durable memo. Wrap each expensive phase of
137    /// [`Pipeline::run`] in [`Memo::memoized`] or
138    /// [`Memo::memoized_by_content`] so a retried item reads the completed
139    /// phases back. Calls to [`record_cost`](Self::record_cost) inside a
140    /// memoized future run only when the future runs; use
141    /// [`memoized_with_cached_cost`](Self::memoized_with_cached_cost) when
142    /// memoized results must also contribute to the cost report.
143    pub fn memo(&self) -> &Memo {
144        &self.memo
145    }
146
147    /// Run `f` once and memoize both its value and counters under `key`,
148    /// or return the memoized value and replay its counters on a retry.
149    ///
150    /// `f` returns `(value, cost)`, and the counters are recorded after
151    /// memoization returns, so they are included whether the phase runs or
152    /// reads its memo entry.
153    pub async fn memoized_with_cached_cost<R, F, E>(&self, key: &str, f: F) -> Result<R, E>
154    where
155        R: Serialize + DeserializeOwned,
156        F: Future<Output = Result<(R, CostReport), E>>,
157        E: From<crate::Error>,
158    {
159        let (value, cost) = self.memo.memoized(key, f).await?;
160        self.cost.merge(&cost);
161        Ok(value)
162    }
163
164    /// [`Self::memoized_with_cached_cost`] under [`Memo::content_key`] of
165    /// `input`.
166    pub async fn memoized_by_content_with_cached_cost<K, R, F, E>(
167        &self,
168        input: &K,
169        f: F,
170    ) -> Result<R, E>
171    where
172        K: Serialize + ?Sized,
173        R: Serialize + DeserializeOwned,
174        F: Future<Output = Result<(R, CostReport), E>>,
175        E: From<crate::Error>,
176    {
177        let (value, cost) = self.memo.memoized_by_content(input, f).await?;
178        self.cost.merge(&cost);
179        Ok(value)
180    }
181
182    /// Add `amount` to the cost counter named `metric` for this item. The
183    /// per-item totals roll up into the batch-level
184    /// [`ProgressSnapshot`](crate::bulk::ProgressSnapshot) and
185    /// [`BulkReport`](crate::bulk::BulkReport).
186    pub fn record_cost(&self, metric: &str, amount: f64) {
187        self.cost.record(metric, amount);
188    }
189
190    /// The run's cooperative cancellation token. Watch it to short-circuit a
191    /// long-running step when the bulk run is draining (e.g. on spot
192    /// preemption); see [`crate::Step::cancel_token`].
193    pub fn cancel_token(&self) -> &CancellationToken {
194        &self.cancel_token
195    }
196
197    /// The lease handle for this item's delivery. A long-running
198    /// pipeline calls [`LeaseHandle::ensure_at_least`] at progress
199    /// points (or once, with a slow call's timeout, before issuing it)
200    /// so the item is not re-queued while it still runs; see
201    /// [`crate::Step::lease`].
202    pub fn lease(&self) -> &LeaseHandle {
203        &self.lease
204    }
205
206    /// The item's staged application KV effects. Writes and deletes
207    /// staged here are applied atomically with the item's successful
208    /// completion; an item that fails applies nothing, and a staged
209    /// value must be correct when applied more than once, since
210    /// delivery is at-least-once. See [`EffectsHandle`] for the
211    /// staging rules.
212    pub fn effects(&self) -> &EffectsHandle {
213        &self.effects
214    }
215
216    /// Read the committed value under `key` from the caller KV
217    /// namespace, `None` when no value exists. Effects staged by this
218    /// item become readable only once its completion commits; see
219    /// [`crate::KvReadHandle`] for the read semantics.
220    pub async fn kv_get(&self, key: &[u8]) -> Result<Option<Vec<u8>>, StepError> {
221        Ok(self
222            .kv
223            .get(key)
224            .await
225            .map_err(StepError::from)?
226            .map(|bytes| bytes.to_vec()))
227    }
228
229    /// Snapshot of the cost accumulated so far for this item.
230    pub(crate) fn cost(&self) -> CostReport {
231        self.cost.clone()
232    }
233}
234
235#[cfg(test)]
236mod tests {
237    use super::*;
238    use crate::MemoStore;
239    use std::sync::Arc;
240    use std::sync::atomic::{AtomicU32, Ordering};
241    use taquba::object_store::memory::InMemory;
242
243    #[derive(Serialize)]
244    struct ContentInput<'a> {
245        operation: &'static str,
246        payload: &'a [u8],
247    }
248
249    fn test_step(store: &MemoStore) -> Step {
250        Step {
251            run_id: "run-1".into(),
252            step_number: 0,
253            payload: Vec::new(),
254            headers: HashMap::new(),
255            job_id: "job-1".into(),
256            attempts: 1,
257            max_attempts: 3,
258            cancel_token: CancellationToken::new(),
259            lease: taquba::LeaseHandle::detached(),
260            memo: store.new_memo("run-1", 0),
261            run_memo: store.new_run_memo("run-1"),
262            effects: EffectsHandle::detached(),
263            kv: KvReadHandle::detached(),
264            signal: None,
265        }
266    }
267
268    fn ctx_for_tests() -> BulkCtx<()> {
269        let store = MemoStore::new(Arc::new(InMemory::new()), "memo");
270        BulkCtx::new("b", "item-1", (), &test_step(&store))
271    }
272
273    #[tokio::test]
274    async fn memoized_with_cached_cost_records_cost_on_compute_and_memo_hit() {
275        let store = MemoStore::new(Arc::new(InMemory::new()), "memo");
276        let first_ctx = BulkCtx::new("b", "item-1", (), &test_step(&store));
277        let replay_ctx = BulkCtx::new("b", "item-1", (), &test_step(&store));
278        let calls = AtomicU32::new(0);
279
280        let first = first_ctx
281            .memoized_with_cached_cost("k", async {
282                calls.fetch_add(1, Ordering::SeqCst);
283                let cost = CostReport::new();
284                cost.record("tokens", 42.0);
285                Ok::<_, StepError>(("value".to_string(), cost))
286            })
287            .await
288            .unwrap();
289        let second = replay_ctx
290            .memoized_with_cached_cost("k", async {
291                calls.fetch_add(1, Ordering::SeqCst);
292                Ok::<_, StepError>(("other".to_string(), CostReport::new()))
293            })
294            .await
295            .unwrap();
296
297        assert_eq!(first, "value");
298        assert_eq!(second, "value");
299        assert_eq!(
300            calls.load(Ordering::SeqCst),
301            1,
302            "memo hit did not run closure"
303        );
304        assert_eq!(first_ctx.cost().get("tokens"), 42.0);
305        assert_eq!(replay_ctx.cost().get("tokens"), 42.0);
306    }
307
308    #[tokio::test]
309    async fn memoized_by_content_with_cached_cost_records_cost_on_compute_and_memo_hit() {
310        let store = MemoStore::new(Arc::new(InMemory::new()), "memo");
311        let first_ctx = BulkCtx::new("b", "item-1", (), &test_step(&store));
312        let replay_ctx = BulkCtx::new("b", "item-1", (), &test_step(&store));
313        let calls = AtomicU32::new(0);
314        let input = ContentInput {
315            operation: "classify",
316            payload: b"ticket",
317        };
318
319        let first = first_ctx
320            .memoized_by_content_with_cached_cost(&input, async {
321                calls.fetch_add(1, Ordering::SeqCst);
322                let cost = CostReport::new();
323                cost.record("tokens", 42.0);
324                Ok::<_, StepError>(("value".to_string(), cost))
325            })
326            .await
327            .unwrap();
328        let second = replay_ctx
329            .memoized_by_content_with_cached_cost(&input, async {
330                calls.fetch_add(1, Ordering::SeqCst);
331                Ok::<_, StepError>(("other".to_string(), CostReport::new()))
332            })
333            .await
334            .unwrap();
335
336        assert_eq!(first, "value");
337        assert_eq!(second, "value");
338        assert_eq!(
339            calls.load(Ordering::SeqCst),
340            1,
341            "memo hit did not run closure"
342        );
343        assert_eq!(first_ctx.cost().get("tokens"), 42.0);
344        assert_eq!(replay_ctx.cost().get("tokens"), 42.0);
345    }
346
347    #[tokio::test]
348    async fn record_cost_accumulates_into_snapshot() {
349        let ctx = ctx_for_tests();
350        ctx.record_cost("tokens", 100.0);
351        ctx.record_cost("tokens", 50.0);
352        assert_eq!(ctx.cost().get("tokens"), 150.0);
353    }
354}