use std::collections::HashMap;
use std::future::Future;
use crate::{EffectsHandle, KvReadHandle, Memo, Step, StepError};
use serde::Serialize;
use serde::de::DeserializeOwned;
use taquba::LeaseHandle;
use tokio_util::sync::CancellationToken;
use crate::bulk::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 batch_id: String,
pub key: String,
pub run_id: String,
pub headers: HashMap<String, String>,
memo: Memo,
cost: CostReport,
cancel_token: CancellationToken,
lease: LeaseHandle,
effects: EffectsHandle,
kv: KvReadHandle,
}
impl<T> BulkCtx<T> {
pub(crate) fn new(batch_id: &str, key: &str, input: T, step: &Step) -> Self {
Self {
input,
batch_id: batch_id.to_string(),
key: key.to_string(),
run_id: step.run_id.clone(),
headers: step.headers.clone(),
memo: step.memo.clone(),
cost: CostReport::new(),
cancel_token: step.cancel_token.clone(),
lease: step.lease.clone(),
effects: step.effects.clone(),
kv: step.kv.clone(),
}
}
pub fn memo(&self) -> &Memo {
&self.memo
}
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<crate::Error>,
{
let (value, cost) = self.memo.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<crate::Error>,
{
let (value, cost) = self.memo.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 fn lease(&self) -> &LeaseHandle {
&self.lease
}
pub fn effects(&self) -> &EffectsHandle {
&self.effects
}
pub async fn kv_get(&self, key: &[u8]) -> Result<Option<Vec<u8>>, StepError> {
Ok(self
.kv
.get(key)
.await
.map_err(StepError::from)?
.map(|bytes| bytes.to_vec()))
}
pub(crate) fn cost(&self) -> CostReport {
self.cost.clone()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::MemoStore;
use std::sync::Arc;
use std::sync::atomic::{AtomicU32, Ordering};
use taquba::object_store::memory::InMemory;
#[derive(Serialize)]
struct ContentInput<'a> {
operation: &'static str,
payload: &'a [u8],
}
fn test_step(store: &MemoStore) -> Step {
Step {
run_id: "run-1".into(),
step_number: 0,
payload: Vec::new(),
headers: HashMap::new(),
job_id: "job-1".into(),
attempts: 1,
max_attempts: 3,
cancel_token: CancellationToken::new(),
lease: taquba::LeaseHandle::detached(),
memo: store.new_memo("run-1", 0),
run_memo: store.new_run_memo("run-1"),
effects: EffectsHandle::detached(),
kv: KvReadHandle::detached(),
signal: None,
}
}
fn ctx_for_tests() -> BulkCtx<()> {
let store = MemoStore::new(Arc::new(InMemory::new()), "memo");
BulkCtx::new("b", "item-1", (), &test_step(&store))
}
#[tokio::test]
async fn memoized_with_cached_cost_records_cost_on_compute_and_memo_hit() {
let store = MemoStore::new(Arc::new(InMemory::new()), "memo");
let first_ctx = BulkCtx::new("b", "item-1", (), &test_step(&store));
let replay_ctx = BulkCtx::new("b", "item-1", (), &test_step(&store));
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 store = MemoStore::new(Arc::new(InMemory::new()), "memo");
let first_ctx = BulkCtx::new("b", "item-1", (), &test_step(&store));
let replay_ctx = BulkCtx::new("b", "item-1", (), &test_step(&store));
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);
}
}