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;
pub trait Pipeline: Send + Sync + 'static {
type Input: Serialize + DeserializeOwned + Send + 'static;
type Output: Serialize + DeserializeOwned + Send + 'static;
type Error: Into<StepError> + Send + 'static;
fn run(
&self,
ctx: &BulkCtx<Self::Input>,
) -> impl Future<Output = Result<Self::Output, Self::Error>> + Send;
}
pub struct BulkCtx<T> {
pub input: T,
pub run_id: String,
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,
}
}
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>,
{
match self.memo.get(key).await {
Ok(Some(bytes)) => match rmp_serde::from_slice::<R>(&bytes) {
Ok(value) => return Ok(value),
Err(err) => {
tracing::warn!(
run_id = %self.run_id,
key = %key,
error = %err,
"memoized cache entry failed to deserialize; recomputing",
);
}
},
Ok(None) => {}
Err(err) => return Err(E::from(memo_error(err))),
}
let value = f.await?;
let bytes = rmp_serde::to_vec_named(&value)
.map_err(|e| E::from(StepError::permanent(format!("memo serialize failed: {e}"))))?;
self.memo
.put(key, &bytes)
.await
.map_err(|e| E::from(memo_error(e)))?;
Ok(value)
}
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>,
{
match self.memo.content_get(input).await {
Ok(Some(bytes)) => match rmp_serde::from_slice::<R>(&bytes) {
Ok(value) => return Ok(value),
Err(err) => {
tracing::warn!(
run_id = %self.run_id,
error = %err,
"content-addressed memoized cache entry failed to deserialize; recomputing",
);
}
},
Ok(None) => {}
Err(err) => return Err(E::from(memo_error(err))),
}
let value = f.await?;
let bytes = rmp_serde::to_vec_named(&value)
.map_err(|e| E::from(StepError::permanent(format!("memo serialize failed: {e}"))))?;
self.memo
.content_put(input, &bytes)
.await
.map_err(|e| E::from(memo_error(e)))?;
Ok(value)
}
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)
}
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)
}
pub fn record_cost(&self, metric: &str, amount: f64) {
self.cost.record(metric, amount);
}
pub fn cancel_token(&self) -> &CancellationToken {
&self.cancel_token
}
pub(crate) fn cost(&self) -> CostReport {
self.cost.clone()
}
}
fn memo_error(err: taquba_workflow::Error) -> StepError {
err.into()
}
#[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");
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);
}
}