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 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;
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 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);
}
}