1use std::future::Future;
50use std::sync::Arc;
51
52use futures_util::StreamExt;
53use serde::Serialize;
54use serde::de::DeserializeOwned;
55use taquba::object_store::{ObjectStore, path::Path};
56
57use crate::blob::ObjectPrefix;
58use crate::durable::decode_or_absent;
59use crate::error::{Error, Result};
60use crate::keys::{RunId, hex_sha256};
61
62pub(crate) const RUN_RESULT_MEMO_KEY: &str = "workflow.outcome";
65
66#[derive(Clone)]
73pub struct MemoStore {
74 objects: ObjectPrefix,
75}
76
77impl std::fmt::Debug for MemoStore {
78 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
79 f.debug_struct("MemoStore")
82 .field("prefix", &self.objects.prefix())
83 .finish_non_exhaustive()
84 }
85}
86
87impl MemoStore {
88 pub fn new(store: Arc<dyn ObjectStore>, prefix: impl Into<String>) -> Self {
94 Self {
95 objects: ObjectPrefix::new(store, prefix),
96 }
97 }
98
99 pub fn new_memo(&self, run_id: &RunId, step_number: u32) -> Memo {
101 Memo::new(self.clone(), run_id, MemoScope::Step(step_number))
102 }
103
104 pub fn new_run_memo(&self, run_id: &RunId) -> Memo {
107 Memo::new(self.clone(), run_id, MemoScope::Run)
108 }
109
110 pub async fn clear_memos_for_run(&self, run_id: &RunId) -> Result<usize> {
116 let memo_deleted = self
117 .clear_prefix(run_id, self.memos_run_prefix(run_id), "memo")
118 .await?;
119 let step_output_deleted = self
120 .clear_prefix(run_id, self.step_outputs_run_prefix(run_id), "step output")
121 .await?;
122 Ok(memo_deleted + step_output_deleted)
123 }
124
125 async fn clear_prefix(
126 &self,
127 run_id: &RunId,
128 prefix: Path,
129 kind: &'static str,
130 ) -> Result<usize> {
131 let mut stream = self.objects.list(&prefix);
132 let mut deleted = 0usize;
133 while let Some(item) = stream.next().await {
134 let meta = item.map_err(Error::Store)?;
135 match self.objects.delete(&meta.location).await {
136 Ok(true) => deleted += 1,
137 Ok(false) => {}
138 Err(err) => {
139 tracing::warn!(
140 run_id = %run_id,
141 path = %meta.location,
142 error = %err,
143 "failed to delete {kind} entry",
144 );
145 }
146 }
147 }
148 Ok(deleted)
149 }
150
151 pub(crate) async fn get_step_output(
152 &self,
153 run_id: &RunId,
154 step_number: u32,
155 step_payload: &[u8],
156 ) -> Result<Option<Vec<u8>>> {
157 self.objects
158 .get(&self.step_output_path(run_id, step_number, step_payload))
159 .await
160 }
161
162 pub(crate) async fn put_step_output(
163 &self,
164 run_id: &RunId,
165 step_number: u32,
166 step_payload: &[u8],
167 value: &[u8],
168 ) -> Result<()> {
169 self.objects
170 .put(
171 &self.step_output_path(run_id, step_number, step_payload),
172 value,
173 )
174 .await
175 }
176
177 fn memo_path(&self, run_id: &RunId, scope: MemoScope, key: &str) -> Path {
178 let segment = match scope {
179 MemoScope::Step(step_number) => step_number.to_string(),
180 MemoScope::Run => "run".to_string(),
181 };
182 self.memos_run_prefix(run_id)
183 .join(segment)
184 .join(hex_sha256(&[key.as_bytes()]))
185 }
186
187 fn memos_run_prefix(&self, run_id: &RunId) -> Path {
188 self.objects.path(&format!("memos/{run_id}"))
189 }
190
191 fn step_outputs_run_prefix(&self, run_id: &RunId) -> Path {
192 self.objects.path(&format!("step-outputs/{run_id}"))
193 }
194
195 fn step_output_path(&self, run_id: &RunId, step_number: u32, step_payload: &[u8]) -> Path {
196 self.step_outputs_run_prefix(run_id)
197 .join(step_number.to_string())
198 .join(hex_sha256(&[step_payload]))
199 }
200}
201
202#[derive(Debug, Clone, Copy, PartialEq, Eq)]
205enum MemoScope {
206 Step(u32),
207 Run,
208}
209
210#[derive(Clone)]
213pub struct Memo {
214 store: MemoStore,
215 run_id: RunId,
216 scope: MemoScope,
217}
218
219impl Memo {
220 fn new(store: MemoStore, run_id: &RunId, scope: MemoScope) -> Self {
221 Self {
222 store,
223 run_id: run_id.clone(),
224 scope,
225 }
226 }
227
228 pub fn run_id(&self) -> &RunId {
230 &self.run_id
231 }
232
233 pub fn step_number(&self) -> Option<u32> {
236 match self.scope {
237 MemoScope::Step(step_number) => Some(step_number),
238 MemoScope::Run => None,
239 }
240 }
241
242 pub async fn get(&self, key: &str) -> Result<Option<Vec<u8>>> {
245 self.store
246 .objects
247 .get(&self.store.memo_path(&self.run_id, self.scope, key))
248 .await
249 }
250
251 pub async fn put(&self, key: &str, value: &[u8]) -> Result<()> {
258 self.store
259 .objects
260 .put(&self.store.memo_path(&self.run_id, self.scope, key), value)
261 .await
262 }
263
264 pub async fn memoized<R, F, E>(&self, key: &str, compute: F) -> std::result::Result<R, E>
276 where
277 R: Serialize + DeserializeOwned,
278 F: Future<Output = std::result::Result<R, E>>,
279 E: From<Error>,
280 {
281 if let Some(bytes) = self.get(key).await?
282 && let Some(value) =
283 decode_or_absent::<R>(&bytes, "memo entry", &format_args!("{}/{key}", self.run_id))
284 {
285 return Ok(value);
286 }
287 let value = compute.await?;
288 let bytes = rmp_serde::to_vec_named(&value).map_err(Error::Serialization)?;
289 self.put(key, &bytes).await?;
290 Ok(value)
291 }
292
293 pub async fn memoized_by_content<K, R, F, E>(
295 &self,
296 input: &K,
297 compute: F,
298 ) -> std::result::Result<R, E>
299 where
300 K: Serialize + ?Sized,
301 R: Serialize + DeserializeOwned,
302 F: Future<Output = std::result::Result<R, E>>,
303 E: From<Error>,
304 {
305 let key = Self::content_key(input)?;
306 self.memoized(&key, compute).await
307 }
308
309 pub fn content_key<T>(input: &T) -> Result<String>
320 where
321 T: Serialize + ?Sized,
322 {
323 let bytes = rmp_serde::to_vec_named(input)?;
324 Ok(format!("content:{}", hex_sha256(&[&bytes])))
325 }
326
327 pub async fn content_get<T>(&self, input: &T) -> Result<Option<Vec<u8>>>
330 where
331 T: Serialize + ?Sized,
332 {
333 let key = Self::content_key(input)?;
334 self.get(&key).await
335 }
336
337 pub async fn content_put<T>(&self, input: &T, value: &[u8]) -> Result<()>
339 where
340 T: Serialize + ?Sized,
341 {
342 let key = Self::content_key(input)?;
343 self.put(&key, value).await
344 }
345}
346
347impl std::fmt::Debug for Memo {
348 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
349 f.debug_struct("Memo")
350 .field("run_id", &self.run_id)
351 .field("scope", &self.scope)
352 .finish_non_exhaustive()
353 }
354}
355
356#[cfg(test)]
357mod tests {
358 use std::sync::atomic::{AtomicU32, Ordering};
359
360 use super::*;
361 use crate::test_util::rid;
362 use serde::Serialize;
363 use taquba::object_store::memory::InMemory;
364
365 #[derive(Serialize)]
366 struct ContentInput<'a> {
367 operation: &'static str,
368 payload: &'a [u8],
369 }
370
371 fn make_memo() -> Memo {
372 MemoStore::new(Arc::new(InMemory::new()), "memo").new_memo(&rid("run-1"), 0)
373 }
374
375 #[tokio::test]
376 async fn put_get_round_trips() {
377 let memo = make_memo();
378 assert_eq!(memo.get("missing").await.unwrap(), None);
379 memo.put("k", b"first").await.unwrap();
380 assert_eq!(memo.get("k").await.unwrap(), Some(b"first".to_vec()));
381 memo.put("k", b"second").await.unwrap();
382 assert_eq!(memo.get("k").await.unwrap(), Some(b"second".to_vec()));
383 memo.put("k2", b"").await.unwrap();
384 assert_eq!(memo.get("k2").await.unwrap(), Some(Vec::new()));
385 assert_eq!(memo.get("k").await.unwrap(), Some(b"second".to_vec()));
386 }
387
388 #[tokio::test]
389 async fn run_and_step_namespaces_are_isolated() {
390 let store = MemoStore::new(Arc::new(InMemory::new()), "memo");
391 let in_run_a = store.new_memo(&rid("run-a"), 0);
392 let in_run_a_step_1 = store.new_memo(&rid("run-a"), 1);
393 let in_run_b = store.new_memo(&rid("run-b"), 0);
394 in_run_a.put("k", b"a-0").await.unwrap();
395 in_run_a_step_1.put("k", b"a-1").await.unwrap();
396 in_run_b.put("k", b"b-0").await.unwrap();
397 assert_eq!(in_run_a.get("k").await.unwrap(), Some(b"a-0".to_vec()));
398 assert_eq!(
399 in_run_a_step_1.get("k").await.unwrap(),
400 Some(b"a-1".to_vec()),
401 );
402 assert_eq!(in_run_b.get("k").await.unwrap(), Some(b"b-0".to_vec()));
403 }
404
405 #[tokio::test]
406 async fn a_run_memo_is_scoped_beside_the_step_memos() {
407 let store = MemoStore::new(Arc::new(InMemory::new()), "memo");
408 let at_step_0 = store.new_memo(&rid("run-1"), 0);
409 let for_run = store.new_run_memo(&rid("run-1"));
410 at_step_0.put("k", b"step-0").await.unwrap();
411 for_run.put("k", b"run").await.unwrap();
412 assert_eq!(at_step_0.get("k").await.unwrap(), Some(b"step-0".to_vec()));
413 assert_eq!(for_run.get("k").await.unwrap(), Some(b"run".to_vec()));
414 assert_eq!(
415 store.new_run_memo(&rid("run-1")).get("k").await.unwrap(),
416 Some(b"run".to_vec()),
417 );
418 assert_eq!(
419 store.new_run_memo(&rid("run-2")).get("k").await.unwrap(),
420 None
421 );
422 }
423
424 #[tokio::test]
425 async fn memoized_runs_the_computation_once() {
426 let memo = make_memo();
427 let calls = AtomicU32::new(0);
428 let compute = || async {
429 calls.fetch_add(1, Ordering::SeqCst);
430 Ok::<_, Error>(7u32)
431 };
432
433 assert_eq!(memo.memoized("k", compute()).await.unwrap(), 7);
434 assert_eq!(memo.memoized("k", compute()).await.unwrap(), 7);
435 assert_eq!(calls.load(Ordering::SeqCst), 1);
436 assert_eq!(
437 memo.get("k").await.unwrap(),
438 Some(rmp_serde::to_vec_named(&7u32).unwrap()),
439 );
440 }
441
442 #[tokio::test]
443 async fn memoized_stores_nothing_for_a_failed_computation() {
444 let memo = make_memo();
445 let failed = memo
446 .memoized("k", async { Err::<u32, Error>(Error::EffectsSealed) })
447 .await;
448
449 assert!(matches!(failed, Err(Error::EffectsSealed)));
450 assert_eq!(memo.get("k").await.unwrap(), None);
451 assert_eq!(
452 memo.memoized("k", async { Ok::<_, Error>(7u32) })
453 .await
454 .unwrap(),
455 7,
456 );
457 }
458
459 #[tokio::test]
460 async fn a_memo_entry_that_fails_to_decode_is_recomputed() {
461 let memo = make_memo();
462 memo.put("k", b"not msgpack for a string").await.unwrap();
463
464 let value: String = memo
465 .memoized("k", async { Ok::<_, Error>("fresh".to_string()) })
466 .await
467 .unwrap();
468
469 assert_eq!(value, "fresh");
470 assert_eq!(
471 memo.get("k").await.unwrap(),
472 Some(rmp_serde::to_vec_named("fresh").unwrap()),
473 );
474 }
475
476 #[tokio::test]
477 async fn memoized_by_content_stores_under_the_content_key() {
478 let memo = make_memo();
479 let input = ContentInput {
480 operation: "draft",
481 payload: b"hello",
482 };
483 let calls = AtomicU32::new(0);
484 let compute = || async {
485 calls.fetch_add(1, Ordering::SeqCst);
486 Ok::<_, Error>(7u32)
487 };
488
489 assert_eq!(
490 memo.memoized_by_content(&input, compute()).await.unwrap(),
491 7
492 );
493 assert_eq!(
494 memo.memoized_by_content(&input, compute()).await.unwrap(),
495 7
496 );
497 assert_eq!(calls.load(Ordering::SeqCst), 1);
498 assert_eq!(
499 memo.content_get(&input).await.unwrap(),
500 Some(rmp_serde::to_vec_named(&7u32).unwrap()),
501 );
502 }
503
504 #[tokio::test]
505 async fn awkward_user_keys_round_trip() {
506 let memo = make_memo();
509 let keys = [
510 "",
511 "with/slash",
512 "with spaces",
513 "üñíçødé",
514 &"a".repeat(10_000),
515 ];
516 for (i, key) in keys.iter().enumerate() {
517 let expected = format!("v{i}").into_bytes();
518 memo.put(key, &expected).await.unwrap();
519 assert_eq!(memo.get(key).await.unwrap(), Some(expected));
520 }
521 }
522
523 #[tokio::test]
524 async fn content_key_distinguishes_serialized_inputs() {
525 let memo = make_memo();
526 let first = ContentInput {
527 operation: "draft",
528 payload: b"hello",
529 };
530 let second = ContentInput {
531 operation: "review",
532 payload: b"hello",
533 };
534
535 memo.content_put(&first, b"first").await.unwrap();
536
537 assert_eq!(
538 memo.content_get(&first).await.unwrap(),
539 Some(b"first".to_vec()),
540 );
541 assert!(memo.content_get(&second).await.unwrap().is_none());
542 }
543
544 #[tokio::test]
545 async fn step_output_entries_are_scoped_by_payload_hash() {
546 let store = MemoStore::new(Arc::new(InMemory::new()), "memo");
547
548 store
549 .put_step_output(&rid("run-1"), 0, b"payload-a", b"out-a")
550 .await
551 .unwrap();
552
553 assert_eq!(
554 store
555 .get_step_output(&rid("run-1"), 0, b"payload-a")
556 .await
557 .unwrap(),
558 Some(b"out-a".to_vec()),
559 );
560 assert!(
561 store
562 .get_step_output(&rid("run-1"), 0, b"payload-b")
563 .await
564 .unwrap()
565 .is_none(),
566 );
567 }
568
569 #[tokio::test]
570 async fn entries_are_stored_at_the_documented_paths() {
571 let backing = Arc::new(InMemory::new());
572 let store = MemoStore::new(backing.clone(), "memo");
573 store
574 .new_memo(&rid("run-1"), 0)
575 .put("k", b"step")
576 .await
577 .unwrap();
578 store
579 .new_run_memo(&rid("run-1"))
580 .put("k", b"run")
581 .await
582 .unwrap();
583 store
584 .new_memo(&rid("run-1"), 0)
585 .content_put("hello", b"content")
586 .await
587 .unwrap();
588 store
589 .put_step_output(&rid("run-1"), 0, b"payload", b"out")
590 .await
591 .unwrap();
592
593 let mut paths = Vec::new();
594 let mut listing = backing.list(None);
595 while let Some(item) = listing.next().await {
596 paths.push(item.unwrap().location.to_string());
597 }
598 paths.sort();
599
600 assert_eq!(
604 paths,
605 [
606 "memo/memos/run-1/0/1adc4f8ba16f15ba2172dab9b84bb9ba73f5cf4f156f50df4dda663b4f9c61ba",
607 "memo/memos/run-1/0/8254c329a92850f6d539dd376f4816ee2764517da5e0235514af433164480d7a",
608 "memo/memos/run-1/run/8254c329a92850f6d539dd376f4816ee2764517da5e0235514af433164480d7a",
609 "memo/step-outputs/run-1/0/239f59ed55e737c77147cf55ad0c1b030b6d7ee748a7426952f9b852d5a935e5",
610 ],
611 );
612 }
613
614 #[tokio::test]
615 async fn clear_memos_for_run_removes_step_output_and_run_memo_entries() {
616 let store = MemoStore::new(Arc::new(InMemory::new()), "memo");
617 store
618 .new_memo(&rid("run-1"), 0)
619 .put("k", b"memo")
620 .await
621 .unwrap();
622 store
623 .new_run_memo(&rid("run-1"))
624 .put("k", b"run")
625 .await
626 .unwrap();
627 store
628 .put_step_output(&rid("run-1"), 0, b"payload", b"out")
629 .await
630 .unwrap();
631
632 let deleted = store.clear_memos_for_run(&rid("run-1")).await.unwrap();
633
634 assert_eq!(deleted, 3);
635 assert!(
636 store
637 .new_memo(&rid("run-1"), 0)
638 .get("k")
639 .await
640 .unwrap()
641 .is_none()
642 );
643 assert!(
644 store
645 .new_run_memo(&rid("run-1"))
646 .get("k")
647 .await
648 .unwrap()
649 .is_none()
650 );
651 assert!(
652 store
653 .get_step_output(&rid("run-1"), 0, b"payload")
654 .await
655 .unwrap()
656 .is_none(),
657 );
658 }
659
660 #[tokio::test]
661 async fn content_key_reports_serialization_errors() {
662 struct BadSerialize;
663
664 impl Serialize for BadSerialize {
665 fn serialize<S>(&self, _serializer: S) -> std::result::Result<S::Ok, S::Error>
666 where
667 S: serde::Serializer,
668 {
669 Err(serde::ser::Error::custom("serialization failed"))
670 }
671 }
672
673 let memo = make_memo();
674 assert!(matches!(
675 memo.content_get(&BadSerialize).await,
676 Err(Error::Serialization(_)),
677 ));
678 }
679
680 #[tokio::test]
681 async fn instances_sharing_a_backing_store_see_the_same_entries() {
682 let backing: Arc<dyn ObjectStore> = Arc::new(InMemory::new());
686 let writer = MemoStore::new(backing.clone(), "memo").new_memo(&rid("run-1"), 0);
687 let reader = MemoStore::new(backing, "memo").new_memo(&rid("run-1"), 0);
688 writer.put("k", b"shared").await.unwrap();
689 assert_eq!(reader.get("k").await.unwrap(), Some(b"shared".to_vec()));
690 }
691
692 #[tokio::test]
693 async fn clear_memos_for_run_removes_only_that_runs_entries() {
694 let backing: Arc<dyn ObjectStore> = Arc::new(InMemory::new());
695 let store = MemoStore::new(backing, "memo");
696 let in_run_a = store.new_memo(&rid("run-a"), 0);
697 let in_run_a_step1 = store.new_memo(&rid("run-a"), 1);
698 let in_run_b = store.new_memo(&rid("run-b"), 0);
699 in_run_a.put("k", b"a-0").await.unwrap();
700 in_run_a_step1.put("k", b"a-1").await.unwrap();
701 in_run_b.put("k", b"b-0").await.unwrap();
702
703 let deleted = store.clear_memos_for_run(&rid("run-a")).await.unwrap();
704 assert_eq!(deleted, 2);
705
706 assert_eq!(in_run_a.get("k").await.unwrap(), None);
707 assert_eq!(in_run_a_step1.get("k").await.unwrap(), None);
708 assert_eq!(in_run_b.get("k").await.unwrap(), Some(b"b-0".to_vec()));
709 assert_eq!(store.clear_memos_for_run(&rid("run-a")).await.unwrap(), 0);
710 }
711
712 #[tokio::test]
713 async fn clear_memos_for_run_does_not_match_run_id_as_prefix() {
714 let store = MemoStore::new(Arc::new(InMemory::new()), "memo");
715 store
716 .new_memo(&rid("run"), 0)
717 .put("k", b"short")
718 .await
719 .unwrap();
720 store
721 .new_memo(&rid("run-suffix"), 0)
722 .put("k", b"long")
723 .await
724 .unwrap();
725
726 let deleted = store.clear_memos_for_run(&rid("run")).await.unwrap();
727 assert_eq!(deleted, 1);
728 assert_eq!(
729 store
730 .new_memo(&rid("run-suffix"), 0)
731 .get("k")
732 .await
733 .unwrap(),
734 Some(b"long".to_vec()),
735 );
736 }
737}