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}