use std::future::Future;
use std::sync::Arc;
use futures_util::StreamExt;
use serde::Serialize;
use serde::de::DeserializeOwned;
use taquba::object_store::{ObjectStore, path::Path};
use crate::blob::ObjectPrefix;
use crate::durable::decode_or_absent;
use crate::error::{Error, Result};
use crate::keys::{RunId, hex_sha256};
pub(crate) const RUN_RESULT_MEMO_KEY: &str = "workflow.outcome";
#[derive(Clone)]
pub struct MemoStore {
objects: ObjectPrefix,
}
impl std::fmt::Debug for MemoStore {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("MemoStore")
.field("prefix", &self.objects.prefix())
.finish_non_exhaustive()
}
}
impl MemoStore {
pub fn new(store: Arc<dyn ObjectStore>, prefix: impl Into<String>) -> Self {
Self {
objects: ObjectPrefix::new(store, prefix),
}
}
pub fn new_memo(&self, run_id: &RunId, step_number: u32) -> Memo {
Memo::new(self.clone(), run_id, MemoScope::Step(step_number))
}
pub fn new_run_memo(&self, run_id: &RunId) -> Memo {
Memo::new(self.clone(), run_id, MemoScope::Run)
}
pub async fn clear_memos_for_run(&self, run_id: &RunId) -> Result<usize> {
let memo_deleted = self
.clear_prefix(run_id, self.memos_run_prefix(run_id), "memo")
.await?;
let step_output_deleted = self
.clear_prefix(run_id, self.step_outputs_run_prefix(run_id), "step output")
.await?;
Ok(memo_deleted + step_output_deleted)
}
async fn clear_prefix(
&self,
run_id: &RunId,
prefix: Path,
kind: &'static str,
) -> Result<usize> {
let mut stream = self.objects.list(&prefix);
let mut deleted = 0usize;
while let Some(item) = stream.next().await {
let meta = item.map_err(Error::Store)?;
match self.objects.delete(&meta.location).await {
Ok(true) => deleted += 1,
Ok(false) => {}
Err(err) => {
tracing::warn!(
run_id = %run_id,
path = %meta.location,
error = %err,
"failed to delete {kind} entry",
);
}
}
}
Ok(deleted)
}
pub(crate) async fn get_step_output(
&self,
run_id: &RunId,
step_number: u32,
step_payload: &[u8],
) -> Result<Option<Vec<u8>>> {
self.objects
.get(&self.step_output_path(run_id, step_number, step_payload))
.await
}
pub(crate) async fn put_step_output(
&self,
run_id: &RunId,
step_number: u32,
step_payload: &[u8],
value: &[u8],
) -> Result<()> {
self.objects
.put(
&self.step_output_path(run_id, step_number, step_payload),
value,
)
.await
}
fn memo_path(&self, run_id: &RunId, scope: MemoScope, key: &str) -> Path {
let segment = match scope {
MemoScope::Step(step_number) => step_number.to_string(),
MemoScope::Run => "run".to_string(),
};
self.memos_run_prefix(run_id)
.join(segment)
.join(hex_sha256(&[key.as_bytes()]))
}
fn memos_run_prefix(&self, run_id: &RunId) -> Path {
self.objects.path(&format!("memos/{run_id}"))
}
fn step_outputs_run_prefix(&self, run_id: &RunId) -> Path {
self.objects.path(&format!("step-outputs/{run_id}"))
}
fn step_output_path(&self, run_id: &RunId, step_number: u32, step_payload: &[u8]) -> Path {
self.step_outputs_run_prefix(run_id)
.join(step_number.to_string())
.join(hex_sha256(&[step_payload]))
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum MemoScope {
Step(u32),
Run,
}
#[derive(Clone)]
pub struct Memo {
store: MemoStore,
run_id: RunId,
scope: MemoScope,
}
impl Memo {
fn new(store: MemoStore, run_id: &RunId, scope: MemoScope) -> Self {
Self {
store,
run_id: run_id.clone(),
scope,
}
}
pub fn run_id(&self) -> &RunId {
&self.run_id
}
pub fn step_number(&self) -> Option<u32> {
match self.scope {
MemoScope::Step(step_number) => Some(step_number),
MemoScope::Run => None,
}
}
pub async fn get(&self, key: &str) -> Result<Option<Vec<u8>>> {
self.store
.objects
.get(&self.store.memo_path(&self.run_id, self.scope, key))
.await
}
pub async fn put(&self, key: &str, value: &[u8]) -> Result<()> {
self.store
.objects
.put(&self.store.memo_path(&self.run_id, self.scope, key), value)
.await
}
pub async fn memoized<R, F, E>(&self, key: &str, compute: F) -> std::result::Result<R, E>
where
R: Serialize + DeserializeOwned,
F: Future<Output = std::result::Result<R, E>>,
E: From<Error>,
{
if let Some(bytes) = self.get(key).await?
&& let Some(value) =
decode_or_absent::<R>(&bytes, "memo entry", &format_args!("{}/{key}", self.run_id))
{
return Ok(value);
}
let value = compute.await?;
let bytes = rmp_serde::to_vec_named(&value).map_err(Error::Serialization)?;
self.put(key, &bytes).await?;
Ok(value)
}
pub async fn memoized_by_content<K, R, F, E>(
&self,
input: &K,
compute: F,
) -> std::result::Result<R, E>
where
K: Serialize + ?Sized,
R: Serialize + DeserializeOwned,
F: Future<Output = std::result::Result<R, E>>,
E: From<Error>,
{
let key = Self::content_key(input)?;
self.memoized(&key, compute).await
}
pub fn content_key<T>(input: &T) -> Result<String>
where
T: Serialize + ?Sized,
{
let bytes = rmp_serde::to_vec_named(input)?;
Ok(format!("content:{}", hex_sha256(&[&bytes])))
}
pub async fn content_get<T>(&self, input: &T) -> Result<Option<Vec<u8>>>
where
T: Serialize + ?Sized,
{
let key = Self::content_key(input)?;
self.get(&key).await
}
pub async fn content_put<T>(&self, input: &T, value: &[u8]) -> Result<()>
where
T: Serialize + ?Sized,
{
let key = Self::content_key(input)?;
self.put(&key, value).await
}
}
impl std::fmt::Debug for Memo {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Memo")
.field("run_id", &self.run_id)
.field("scope", &self.scope)
.finish_non_exhaustive()
}
}
#[cfg(test)]
mod tests {
use std::sync::atomic::{AtomicU32, Ordering};
use super::*;
use crate::test_util::rid;
use serde::Serialize;
use taquba::object_store::memory::InMemory;
#[derive(Serialize)]
struct ContentInput<'a> {
operation: &'static str,
payload: &'a [u8],
}
fn make_memo() -> Memo {
MemoStore::new(Arc::new(InMemory::new()), "memo").new_memo(&rid("run-1"), 0)
}
#[tokio::test]
async fn put_get_round_trips() {
let memo = make_memo();
assert_eq!(memo.get("missing").await.unwrap(), None);
memo.put("k", b"first").await.unwrap();
assert_eq!(memo.get("k").await.unwrap(), Some(b"first".to_vec()));
memo.put("k", b"second").await.unwrap();
assert_eq!(memo.get("k").await.unwrap(), Some(b"second".to_vec()));
memo.put("k2", b"").await.unwrap();
assert_eq!(memo.get("k2").await.unwrap(), Some(Vec::new()));
assert_eq!(memo.get("k").await.unwrap(), Some(b"second".to_vec()));
}
#[tokio::test]
async fn run_and_step_namespaces_are_isolated() {
let store = MemoStore::new(Arc::new(InMemory::new()), "memo");
let in_run_a = store.new_memo(&rid("run-a"), 0);
let in_run_a_step_1 = store.new_memo(&rid("run-a"), 1);
let in_run_b = store.new_memo(&rid("run-b"), 0);
in_run_a.put("k", b"a-0").await.unwrap();
in_run_a_step_1.put("k", b"a-1").await.unwrap();
in_run_b.put("k", b"b-0").await.unwrap();
assert_eq!(in_run_a.get("k").await.unwrap(), Some(b"a-0".to_vec()));
assert_eq!(
in_run_a_step_1.get("k").await.unwrap(),
Some(b"a-1".to_vec()),
);
assert_eq!(in_run_b.get("k").await.unwrap(), Some(b"b-0".to_vec()));
}
#[tokio::test]
async fn a_run_memo_is_scoped_beside_the_step_memos() {
let store = MemoStore::new(Arc::new(InMemory::new()), "memo");
let at_step_0 = store.new_memo(&rid("run-1"), 0);
let for_run = store.new_run_memo(&rid("run-1"));
at_step_0.put("k", b"step-0").await.unwrap();
for_run.put("k", b"run").await.unwrap();
assert_eq!(at_step_0.get("k").await.unwrap(), Some(b"step-0".to_vec()));
assert_eq!(for_run.get("k").await.unwrap(), Some(b"run".to_vec()));
assert_eq!(
store.new_run_memo(&rid("run-1")).get("k").await.unwrap(),
Some(b"run".to_vec()),
);
assert_eq!(
store.new_run_memo(&rid("run-2")).get("k").await.unwrap(),
None
);
}
#[tokio::test]
async fn memoized_runs_the_computation_once() {
let memo = make_memo();
let calls = AtomicU32::new(0);
let compute = || async {
calls.fetch_add(1, Ordering::SeqCst);
Ok::<_, Error>(7u32)
};
assert_eq!(memo.memoized("k", compute()).await.unwrap(), 7);
assert_eq!(memo.memoized("k", compute()).await.unwrap(), 7);
assert_eq!(calls.load(Ordering::SeqCst), 1);
assert_eq!(
memo.get("k").await.unwrap(),
Some(rmp_serde::to_vec_named(&7u32).unwrap()),
);
}
#[tokio::test]
async fn memoized_stores_nothing_for_a_failed_computation() {
let memo = make_memo();
let failed = memo
.memoized("k", async { Err::<u32, Error>(Error::EffectsSealed) })
.await;
assert!(matches!(failed, Err(Error::EffectsSealed)));
assert_eq!(memo.get("k").await.unwrap(), None);
assert_eq!(
memo.memoized("k", async { Ok::<_, Error>(7u32) })
.await
.unwrap(),
7,
);
}
#[tokio::test]
async fn a_memo_entry_that_fails_to_decode_is_recomputed() {
let memo = make_memo();
memo.put("k", b"not msgpack for a string").await.unwrap();
let value: String = memo
.memoized("k", async { Ok::<_, Error>("fresh".to_string()) })
.await
.unwrap();
assert_eq!(value, "fresh");
assert_eq!(
memo.get("k").await.unwrap(),
Some(rmp_serde::to_vec_named("fresh").unwrap()),
);
}
#[tokio::test]
async fn memoized_by_content_stores_under_the_content_key() {
let memo = make_memo();
let input = ContentInput {
operation: "draft",
payload: b"hello",
};
let calls = AtomicU32::new(0);
let compute = || async {
calls.fetch_add(1, Ordering::SeqCst);
Ok::<_, Error>(7u32)
};
assert_eq!(
memo.memoized_by_content(&input, compute()).await.unwrap(),
7
);
assert_eq!(
memo.memoized_by_content(&input, compute()).await.unwrap(),
7
);
assert_eq!(calls.load(Ordering::SeqCst), 1);
assert_eq!(
memo.content_get(&input).await.unwrap(),
Some(rmp_serde::to_vec_named(&7u32).unwrap()),
);
}
#[tokio::test]
async fn awkward_user_keys_round_trip() {
let memo = make_memo();
let keys = [
"",
"with/slash",
"with spaces",
"üñÃçødé",
&"a".repeat(10_000),
];
for (i, key) in keys.iter().enumerate() {
let expected = format!("v{i}").into_bytes();
memo.put(key, &expected).await.unwrap();
assert_eq!(memo.get(key).await.unwrap(), Some(expected));
}
}
#[tokio::test]
async fn content_key_distinguishes_serialized_inputs() {
let memo = make_memo();
let first = ContentInput {
operation: "draft",
payload: b"hello",
};
let second = ContentInput {
operation: "review",
payload: b"hello",
};
memo.content_put(&first, b"first").await.unwrap();
assert_eq!(
memo.content_get(&first).await.unwrap(),
Some(b"first".to_vec()),
);
assert!(memo.content_get(&second).await.unwrap().is_none());
}
#[tokio::test]
async fn step_output_entries_are_scoped_by_payload_hash() {
let store = MemoStore::new(Arc::new(InMemory::new()), "memo");
store
.put_step_output(&rid("run-1"), 0, b"payload-a", b"out-a")
.await
.unwrap();
assert_eq!(
store
.get_step_output(&rid("run-1"), 0, b"payload-a")
.await
.unwrap(),
Some(b"out-a".to_vec()),
);
assert!(
store
.get_step_output(&rid("run-1"), 0, b"payload-b")
.await
.unwrap()
.is_none(),
);
}
#[tokio::test]
async fn entries_are_stored_at_the_documented_paths() {
let backing = Arc::new(InMemory::new());
let store = MemoStore::new(backing.clone(), "memo");
store
.new_memo(&rid("run-1"), 0)
.put("k", b"step")
.await
.unwrap();
store
.new_run_memo(&rid("run-1"))
.put("k", b"run")
.await
.unwrap();
store
.new_memo(&rid("run-1"), 0)
.content_put("hello", b"content")
.await
.unwrap();
store
.put_step_output(&rid("run-1"), 0, b"payload", b"out")
.await
.unwrap();
let mut paths = Vec::new();
let mut listing = backing.list(None);
while let Some(item) = listing.next().await {
paths.push(item.unwrap().location.to_string());
}
paths.sort();
assert_eq!(
paths,
[
"memo/memos/run-1/0/1adc4f8ba16f15ba2172dab9b84bb9ba73f5cf4f156f50df4dda663b4f9c61ba",
"memo/memos/run-1/0/8254c329a92850f6d539dd376f4816ee2764517da5e0235514af433164480d7a",
"memo/memos/run-1/run/8254c329a92850f6d539dd376f4816ee2764517da5e0235514af433164480d7a",
"memo/step-outputs/run-1/0/239f59ed55e737c77147cf55ad0c1b030b6d7ee748a7426952f9b852d5a935e5",
],
);
}
#[tokio::test]
async fn clear_memos_for_run_removes_step_output_and_run_memo_entries() {
let store = MemoStore::new(Arc::new(InMemory::new()), "memo");
store
.new_memo(&rid("run-1"), 0)
.put("k", b"memo")
.await
.unwrap();
store
.new_run_memo(&rid("run-1"))
.put("k", b"run")
.await
.unwrap();
store
.put_step_output(&rid("run-1"), 0, b"payload", b"out")
.await
.unwrap();
let deleted = store.clear_memos_for_run(&rid("run-1")).await.unwrap();
assert_eq!(deleted, 3);
assert!(
store
.new_memo(&rid("run-1"), 0)
.get("k")
.await
.unwrap()
.is_none()
);
assert!(
store
.new_run_memo(&rid("run-1"))
.get("k")
.await
.unwrap()
.is_none()
);
assert!(
store
.get_step_output(&rid("run-1"), 0, b"payload")
.await
.unwrap()
.is_none(),
);
}
#[tokio::test]
async fn content_key_reports_serialization_errors() {
struct BadSerialize;
impl Serialize for BadSerialize {
fn serialize<S>(&self, _serializer: S) -> std::result::Result<S::Ok, S::Error>
where
S: serde::Serializer,
{
Err(serde::ser::Error::custom("serialization failed"))
}
}
let memo = make_memo();
assert!(matches!(
memo.content_get(&BadSerialize).await,
Err(Error::Serialization(_)),
));
}
#[tokio::test]
async fn instances_sharing_a_backing_store_see_the_same_entries() {
let backing: Arc<dyn ObjectStore> = Arc::new(InMemory::new());
let writer = MemoStore::new(backing.clone(), "memo").new_memo(&rid("run-1"), 0);
let reader = MemoStore::new(backing, "memo").new_memo(&rid("run-1"), 0);
writer.put("k", b"shared").await.unwrap();
assert_eq!(reader.get("k").await.unwrap(), Some(b"shared".to_vec()));
}
#[tokio::test]
async fn clear_memos_for_run_removes_only_that_runs_entries() {
let backing: Arc<dyn ObjectStore> = Arc::new(InMemory::new());
let store = MemoStore::new(backing, "memo");
let in_run_a = store.new_memo(&rid("run-a"), 0);
let in_run_a_step1 = store.new_memo(&rid("run-a"), 1);
let in_run_b = store.new_memo(&rid("run-b"), 0);
in_run_a.put("k", b"a-0").await.unwrap();
in_run_a_step1.put("k", b"a-1").await.unwrap();
in_run_b.put("k", b"b-0").await.unwrap();
let deleted = store.clear_memos_for_run(&rid("run-a")).await.unwrap();
assert_eq!(deleted, 2);
assert_eq!(in_run_a.get("k").await.unwrap(), None);
assert_eq!(in_run_a_step1.get("k").await.unwrap(), None);
assert_eq!(in_run_b.get("k").await.unwrap(), Some(b"b-0".to_vec()));
assert_eq!(store.clear_memos_for_run(&rid("run-a")).await.unwrap(), 0);
}
#[tokio::test]
async fn clear_memos_for_run_does_not_match_run_id_as_prefix() {
let store = MemoStore::new(Arc::new(InMemory::new()), "memo");
store
.new_memo(&rid("run"), 0)
.put("k", b"short")
.await
.unwrap();
store
.new_memo(&rid("run-suffix"), 0)
.put("k", b"long")
.await
.unwrap();
let deleted = store.clear_memos_for_run(&rid("run")).await.unwrap();
assert_eq!(deleted, 1);
assert_eq!(
store
.new_memo(&rid("run-suffix"), 0)
.get("k")
.await
.unwrap(),
Some(b"long".to_vec()),
);
}
}