1use std::collections::{HashMap, VecDeque};
7use std::marker::PhantomData;
8use std::sync::Arc;
9use std::sync::atomic::{AtomicU64, Ordering};
10use std::time::SystemTime;
11
12use dashmap::DashMap;
13use serde_json::Value;
14
15const DEFAULT_MUTATION_JOURNAL_CAPACITY: usize = 1024;
16
17pub struct StateKey<T> {
29 key: &'static str,
30 _phantom: PhantomData<fn() -> T>,
31}
32
33impl<T> StateKey<T> {
34 pub const fn new(key: &'static str) -> Self {
36 Self {
37 key,
38 _phantom: PhantomData,
39 }
40 }
41
42 pub const fn key(&self) -> &'static str {
44 self.key
45 }
46}
47
48#[derive(Debug, Clone, Copy, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
50#[serde(rename_all = "snake_case")]
51pub enum StateMutationOrigin {
52 Set,
54 SetCommitted,
56 Remove,
58 ClearPrefix,
60 Commit,
62}
63
64#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
69pub struct StateMutation {
70 pub sequence: u64,
72 pub key: String,
74 pub old: Option<Value>,
76 pub new: Option<Value>,
78 pub origin: StateMutationOrigin,
80 #[serde(rename = "timestamp_ms", with = "systemtime_epoch_millis")]
83 pub timestamp: SystemTime,
84 pub delta: bool,
86}
87
88mod systemtime_epoch_millis {
90 use std::time::{Duration, SystemTime, UNIX_EPOCH};
91
92 use serde::{Deserialize, Deserializer, Serializer};
93
94 pub(super) fn serialize<S: Serializer>(t: &SystemTime, ser: S) -> Result<S::Ok, S::Error> {
95 let millis = t
96 .duration_since(UNIX_EPOCH)
97 .map(|d| d.as_millis() as u64)
98 .unwrap_or(0);
99 ser.serialize_u64(millis)
100 }
101
102 pub(super) fn deserialize<'de, D: Deserializer<'de>>(de: D) -> Result<SystemTime, D::Error> {
103 let millis = u64::deserialize(de)?;
104 Ok(UNIX_EPOCH + Duration::from_millis(millis))
105 }
106}
107
108pub trait JournalSink: Send + Sync {
118 fn write(&self, m: &StateMutation);
120}
121
122#[derive(Clone, Default)]
125struct JournalSinkSlot(Arc<parking_lot::RwLock<Option<Arc<dyn JournalSink>>>>);
126
127impl std::fmt::Debug for JournalSinkSlot {
128 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
129 let installed = self.0.read().is_some();
130 f.debug_tuple("JournalSinkSlot").field(&installed).finish()
131 }
132}
133
134const JOURNAL_FLUSH_INTERVAL: std::time::Duration = std::time::Duration::from_secs(1);
135
136fn journal_log_error(context: &'static str, e: &dyn std::fmt::Display) {
139 tracing::warn!(error = %e, "{context}");
140}
141
142struct FileJournalInner {
143 writer: std::io::BufWriter<std::fs::File>,
144 last_flush: std::time::Instant,
145}
146
147pub struct FileJournalSink {
157 inner: parking_lot::Mutex<FileJournalInner>,
158}
159
160impl FileJournalSink {
161 pub fn create(path: impl AsRef<std::path::Path>) -> std::io::Result<Self> {
163 let file = std::fs::File::create(path)?;
164 Ok(Self {
165 inner: parking_lot::Mutex::new(FileJournalInner {
166 writer: std::io::BufWriter::new(file),
167 last_flush: std::time::Instant::now(),
168 }),
169 })
170 }
171
172 pub fn flush(&self) {
174 let mut inner = self.inner.lock();
175 if let Err(e) = std::io::Write::flush(&mut inner.writer) {
176 journal_log_error("FileJournalSink flush failed", &e);
177 }
178 inner.last_flush = std::time::Instant::now();
179 }
180}
181
182impl JournalSink for FileJournalSink {
183 fn write(&self, m: &StateMutation) {
184 let line = match serde_json::to_string(m) {
185 Ok(line) => line,
186 Err(e) => {
187 journal_log_error("FileJournalSink serialize failed", &e);
188 return;
189 }
190 };
191 let mut inner = self.inner.lock();
192 if let Err(e) = std::io::Write::write_all(&mut inner.writer, line.as_bytes())
193 .and_then(|()| std::io::Write::write_all(&mut inner.writer, b"\n"))
194 {
195 journal_log_error("FileJournalSink write failed", &e);
196 return;
197 }
198 if inner.last_flush.elapsed() >= JOURNAL_FLUSH_INTERVAL {
199 if let Err(e) = std::io::Write::flush(&mut inner.writer) {
200 journal_log_error("FileJournalSink flush failed", &e);
201 }
202 inner.last_flush = std::time::Instant::now();
203 }
204 }
205}
206
207impl Drop for FileJournalSink {
208 fn drop(&mut self) {
209 if let Err(e) = std::io::Write::flush(&mut self.inner.lock().writer) {
210 journal_log_error("FileJournalSink final flush failed", &e);
211 }
212 }
213}
214
215#[derive(Default)]
217pub struct MemoryJournalSink {
218 entries: parking_lot::Mutex<Vec<StateMutation>>,
219}
220
221impl MemoryJournalSink {
222 pub fn new() -> Self {
224 Self::default()
225 }
226
227 pub fn entries(&self) -> Vec<StateMutation> {
229 self.entries.lock().clone()
230 }
231
232 pub fn len(&self) -> usize {
234 self.entries.lock().len()
235 }
236
237 pub fn is_empty(&self) -> bool {
239 self.entries.lock().is_empty()
240 }
241}
242
243impl JournalSink for MemoryJournalSink {
244 fn write(&self, m: &StateMutation) {
245 self.entries.lock().push(m.clone());
246 }
247}
248
249#[derive(Debug, thiserror::Error)]
251pub enum StateError {
252 #[error("failed to serialize state value for key '{key}': {source}")]
254 Serialize {
255 key: String,
257 source: serde_json::Error,
259 },
260 #[error("state value at key '{key}' is not the requested type: {source}")]
263 WrongType {
264 key: String,
266 source: serde_json::Error,
268 },
269}
270
271#[derive(Debug, Clone)]
277enum DeltaOp {
278 Put(Value),
280 Delete,
282}
283
284#[derive(Debug, Clone, serde::Serialize)]
291pub struct SlotEvidence {
292 pub key: String,
294 pub present: bool,
296 pub value: Option<Value>,
298 pub source: Option<String>,
301 pub confidence: Option<f64>,
303 pub last_sequence: Option<u64>,
306 pub last_origin: Option<StateMutationOrigin>,
308}
309
310#[derive(Debug, Clone)]
316pub struct State {
317 inner: Arc<DashMap<String, Value>>,
318 delta: Arc<DashMap<String, DeltaOp>>,
319 mutations: Arc<std::sync::Mutex<VecDeque<StateMutation>>>,
320 next_mutation_sequence: Arc<AtomicU64>,
321 mutation_capacity: usize,
322 journal_sink: JournalSinkSlot,
323 track_delta: bool,
324}
325
326impl Default for State {
327 fn default() -> Self {
328 Self::new()
329 }
330}
331
332impl State {
333 pub fn new() -> Self {
335 Self {
336 inner: Arc::new(DashMap::new()),
337 delta: Arc::new(DashMap::new()),
338 mutations: Arc::new(std::sync::Mutex::new(VecDeque::new())),
339 next_mutation_sequence: Arc::new(AtomicU64::new(1)),
340 mutation_capacity: DEFAULT_MUTATION_JOURNAL_CAPACITY,
341 journal_sink: JournalSinkSlot::default(),
342 track_delta: false,
343 }
344 }
345
346 pub fn with_delta_tracking(&self) -> State {
349 State {
350 inner: self.inner.clone(),
351 delta: Arc::new(DashMap::new()),
352 mutations: self.mutations.clone(),
353 next_mutation_sequence: self.next_mutation_sequence.clone(),
354 mutation_capacity: self.mutation_capacity,
355 journal_sink: self.journal_sink.clone(),
356 track_delta: true,
357 }
358 }
359
360 pub fn set_journal_sink(&self, sink: Arc<dyn JournalSink>) {
368 *self.journal_sink.0.write() = Some(sink);
369 }
370
371 pub fn with_journal_sink(self, sink: Arc<dyn JournalSink>) -> Self {
373 self.set_journal_sink(sink);
374 self
375 }
376
377 pub fn get<T: serde::de::DeserializeOwned>(&self, key: &str) -> Option<T> {
384 self.get_raw(key)
385 .and_then(|v| serde_json::from_value(v).ok())
386 }
387
388 pub fn try_get<T: serde::de::DeserializeOwned>(
398 &self,
399 key: &str,
400 ) -> Result<Option<T>, StateError> {
401 match self.get_raw(key) {
402 None => Ok(None),
403 Some(v) => {
404 serde_json::from_value(v)
405 .map(Some)
406 .map_err(|source| StateError::WrongType {
407 key: key.to_string(),
408 source,
409 })
410 }
411 }
412 }
413
414 pub fn with<F, R>(&self, key: &str, f: F) -> Option<R>
422 where
423 F: FnOnce(&Value) -> R,
424 {
425 if self.track_delta {
426 match self.delta.get(key).map(|r| r.value().clone()) {
427 Some(DeltaOp::Put(v)) => return Some(f(&v)),
428 Some(DeltaOp::Delete) => return None, None => {}
430 }
431 }
432 if let Some(ref_multi) = self.inner.get(key) {
433 return Some(f(ref_multi.value()));
434 }
435 if !key.contains(':') {
436 let mut derived_key = String::with_capacity(8 + key.len());
437 use std::fmt::Write;
438 let _ = write!(derived_key, "derived:{key}");
439 if self.track_delta {
440 match self.delta.get(&derived_key).map(|r| r.value().clone()) {
441 Some(DeltaOp::Put(v)) => return Some(f(&v)),
442 Some(DeltaOp::Delete) => return None,
443 None => {}
444 }
445 }
446 if let Some(ref_multi) = self.inner.get(&derived_key) {
447 return Some(f(ref_multi.value()));
448 }
449 }
450 None
451 }
452
453 pub fn get_raw(&self, key: &str) -> Option<Value> {
458 if self.track_delta {
459 match self.delta.get(key).map(|r| r.value().clone()) {
460 Some(DeltaOp::Put(v)) => return Some(v),
461 Some(DeltaOp::Delete) => return None, None => {}
463 }
464 }
465 if let Some(v) = self.inner.get(key) {
466 return Some(v.value().clone());
467 }
468 if !key.contains(':') {
470 use std::fmt::Write;
471 let mut derived_key = String::with_capacity(8 + key.len());
472 let _ = write!(derived_key, "derived:{key}");
473 if self.track_delta {
474 match self.delta.get(&derived_key).map(|r| r.value().clone()) {
475 Some(DeltaOp::Put(v)) => return Some(v),
476 Some(DeltaOp::Delete) => return None,
477 None => {}
478 }
479 }
480 return self.inner.get(&derived_key).map(|v| v.value().clone());
481 }
482 None
483 }
484
485 pub fn get_key<T: serde::de::DeserializeOwned>(&self, key: &StateKey<T>) -> Option<T> {
488 self.get(key.key())
489 }
490
491 pub fn try_get_key<T: serde::de::DeserializeOwned>(
494 &self,
495 key: &StateKey<T>,
496 ) -> Result<Option<T>, StateError> {
497 self.try_get(key.key())
498 }
499
500 pub fn set_key<T: serde::Serialize>(
504 &self,
505 key: &StateKey<T>,
506 value: T,
507 ) -> Result<(), StateError> {
508 self.set(key.key(), value)
509 }
510
511 pub fn with_key<T, F, R>(&self, key: &StateKey<T>, f: F) -> Option<R>
513 where
514 F: FnOnce(&Value) -> R,
515 {
516 self.with(key.key(), f)
517 }
518
519 pub fn set(
525 &self,
526 key: impl Into<String>,
527 value: impl serde::Serialize,
528 ) -> Result<(), StateError> {
529 let key = key.into();
530 let v = serde_json::to_value(value).map_err(|source| StateError::Serialize {
531 key: key.clone(),
532 source,
533 })?;
534 self.put_value(key, v, StateMutationOrigin::Set);
535 Ok(())
536 }
537
538 fn put_value(&self, key: String, v: Value, origin: StateMutationOrigin) {
543 let old = self.get_raw(&key);
544 if self.track_delta {
545 self.delta.insert(key.clone(), DeltaOp::Put(v.clone()));
546 } else {
547 self.inner.insert(key.clone(), v.clone());
548 }
549 self.record_mutation(key, old, Some(v), origin);
550 }
551
552 pub fn set_committed(
556 &self,
557 key: impl Into<String>,
558 value: impl serde::Serialize,
559 ) -> Result<(), StateError> {
560 let key = key.into();
561 let v = serde_json::to_value(value).map_err(|source| StateError::Serialize {
562 key: key.clone(),
563 source,
564 })?;
565 let old = self.inner.insert(key.clone(), v.clone());
566 self.record_mutation(key, old, Some(v), StateMutationOrigin::SetCommitted);
567 Ok(())
568 }
569
570 pub fn modify<T, F>(&self, key: &str, default: T, f: F) -> Result<T, StateError>
578 where
579 T: serde::Serialize + serde::de::DeserializeOwned,
580 F: FnOnce(T) -> T,
581 {
582 use dashmap::mapref::entry::Entry;
583
584 let serialize = |key: &str, val: &T| {
585 serde_json::to_value(val).map_err(|source| StateError::Serialize {
586 key: key.to_string(),
587 source,
588 })
589 };
590
591 if self.track_delta {
592 match self.delta.entry(key.to_string()) {
595 Entry::Occupied(mut o) => {
596 let current = match o.get() {
597 DeltaOp::Put(v) => serde_json::from_value(v.clone()).unwrap_or(default),
598 DeltaOp::Delete => default,
599 };
600 let old = self.inner.get(key).map(|r| r.value().clone());
601 let new_val = f(current);
602 let v = serialize(key, &new_val)?;
603 o.insert(DeltaOp::Put(v.clone()));
604 self.record_mutation(key.to_string(), old, Some(v), StateMutationOrigin::Set);
605 Ok(new_val)
606 }
607 Entry::Vacant(slot) => {
608 let base = self
609 .inner
610 .get(key)
611 .and_then(|r| serde_json::from_value(r.value().clone()).ok());
612 let old = self.inner.get(key).map(|r| r.value().clone());
613 let new_val = f(base.unwrap_or(default));
614 let v = serialize(key, &new_val)?;
615 slot.insert(DeltaOp::Put(v.clone()));
616 self.record_mutation(key.to_string(), old, Some(v), StateMutationOrigin::Set);
617 Ok(new_val)
618 }
619 }
620 } else {
621 match self.inner.entry(key.to_string()) {
622 Entry::Occupied(mut o) => {
623 let old = o.get().clone();
624 let current = serde_json::from_value(old.clone()).unwrap_or(default);
625 let new_val = f(current);
626 let v = serialize(key, &new_val)?;
627 o.insert(v.clone());
628 self.record_mutation(
629 key.to_string(),
630 Some(old),
631 Some(v),
632 StateMutationOrigin::Set,
633 );
634 Ok(new_val)
635 }
636 Entry::Vacant(slot) => {
637 let new_val = f(default);
638 let v = serialize(key, &new_val)?;
639 slot.insert(v.clone());
640 self.record_mutation(key.to_string(), None, Some(v), StateMutationOrigin::Set);
641 Ok(new_val)
642 }
643 }
644 }
645 }
646
647 pub fn contains(&self, key: &str) -> bool {
655 if self.track_delta {
656 match self.delta.get(key).map(|r| r.value().clone()) {
657 Some(DeltaOp::Put(_)) => return true,
658 Some(DeltaOp::Delete) => return false, None => {}
660 }
661 }
662 if self.inner.contains_key(key) {
663 return true;
664 }
665 if !key.contains(':') {
666 let derived_key = format!("derived:{key}");
667 if self.track_delta {
668 match self.delta.get(&derived_key).map(|r| r.value().clone()) {
669 Some(DeltaOp::Put(_)) => return true,
670 Some(DeltaOp::Delete) => return false,
671 None => {}
672 }
673 }
674 return self.inner.contains_key(&derived_key);
675 }
676 false
677 }
678
679 pub fn remove(&self, key: &str) -> Option<Value> {
685 if self.track_delta {
686 let removed = self.get_raw(key);
687 self.delta.insert(key.to_string(), DeltaOp::Delete);
690 if let Some(ref old) = removed {
691 self.record_mutation(
692 key.to_string(),
693 Some(old.clone()),
694 None,
695 StateMutationOrigin::Remove,
696 );
697 }
698 removed
699 } else {
700 let removed = self.inner.remove(key).map(|(_, v)| v);
701 if let Some(ref old) = removed {
702 self.record_mutation(
703 key.to_string(),
704 Some(old.clone()),
705 None,
706 StateMutationOrigin::Remove,
707 );
708 }
709 removed
710 }
711 }
712
713 pub fn keys(&self) -> Vec<String> {
717 if !self.track_delta || self.delta.is_empty() {
718 return self.inner.iter().map(|r| r.key().clone()).collect();
719 }
720 let mut seen =
721 std::collections::HashSet::with_capacity(self.inner.len() + self.delta.len());
722 let mut keys = Vec::with_capacity(self.inner.len() + self.delta.len());
723 for entry in self.delta.iter() {
725 let key = entry.key().clone();
726 seen.insert(key.clone());
727 if matches!(entry.value(), DeltaOp::Put(_)) {
728 keys.push(key);
729 }
730 }
731 for entry in self.inner.iter() {
732 let key = entry.key().clone();
733 if seen.insert(key.clone()) {
734 keys.push(key);
735 }
736 }
737 keys
738 }
739
740 pub fn pick(&self, keys: &[&str]) -> State {
742 let new = State::new();
743 for key in keys {
744 if let Some(v) = self.get_raw(key) {
745 new.put_value((*key).to_string(), v, StateMutationOrigin::Set);
746 }
747 }
748 new
749 }
750
751 pub fn merge(&self, other: &State) {
753 for entry in other.inner.iter() {
754 self.put_value(
755 entry.key().clone(),
756 entry.value().clone(),
757 StateMutationOrigin::Set,
758 );
759 }
760 }
761
762 pub fn rename(&self, from: &str, to: &str) {
764 if let Some(v) = self.remove(from) {
765 self.put_value(to.to_string(), v, StateMutationOrigin::Set);
766 }
767 }
768
769 pub fn is_tracking_delta(&self) -> bool {
773 self.track_delta
774 }
775
776 pub fn has_delta(&self) -> bool {
778 self.track_delta && !self.delta.is_empty()
779 }
780
781 pub fn delta(&self) -> HashMap<String, Value> {
783 self.delta
784 .iter()
785 .filter_map(|entry| match entry.value() {
786 DeltaOp::Put(v) => Some((entry.key().clone(), v.clone())),
787 DeltaOp::Delete => None,
788 })
789 .collect()
790 }
791
792 pub fn commit(&self) {
797 let ops: Vec<(String, DeltaOp)> = self
799 .delta
800 .iter()
801 .map(|e| (e.key().clone(), e.value().clone()))
802 .collect();
803 for (key, op) in ops {
804 match op {
805 DeltaOp::Put(value) => {
806 let old = self.inner.insert(key.clone(), value.clone());
807 self.record_mutation_with_delta(
808 key,
809 old,
810 Some(value),
811 StateMutationOrigin::Commit,
812 false,
813 );
814 }
815 DeltaOp::Delete => {
816 if let Some((_, old)) = self.inner.remove(&key) {
817 self.record_mutation_with_delta(
818 key,
819 Some(old),
820 None,
821 StateMutationOrigin::Commit,
822 false,
823 );
824 }
825 }
826 }
827 }
828 self.delta.clear();
829 }
830
831 pub fn rollback(&self) {
837 self.delta.clear();
838 }
839
840 pub fn app(&self) -> PrefixedState<'_> {
844 PrefixedState {
845 state: self,
846 prefix: "app:",
847 }
848 }
849
850 pub fn user(&self) -> PrefixedState<'_> {
852 PrefixedState {
853 state: self,
854 prefix: "user:",
855 }
856 }
857
858 pub fn temp(&self) -> PrefixedState<'_> {
860 PrefixedState {
861 state: self,
862 prefix: "temp:",
863 }
864 }
865
866 pub fn session(&self) -> PrefixedState<'_> {
868 PrefixedState {
869 state: self,
870 prefix: "session:",
871 }
872 }
873
874 pub fn turn(&self) -> PrefixedState<'_> {
876 PrefixedState {
877 state: self,
878 prefix: "turn:",
879 }
880 }
881
882 pub fn bg(&self) -> PrefixedState<'_> {
884 PrefixedState {
885 state: self,
886 prefix: "bg:",
887 }
888 }
889
890 pub fn derived(&self) -> ReadOnlyPrefixedState<'_> {
892 ReadOnlyPrefixedState {
893 state: self,
894 prefix: "derived:",
895 }
896 }
897
898 pub fn snapshot_values(&self, keys: &[&str]) -> HashMap<String, Value> {
903 keys.iter()
904 .filter_map(|&k| self.get_raw(k).map(|v| (k.to_string(), v)))
905 .collect()
906 }
907
908 pub fn diff_values(
911 &self,
912 prev: &HashMap<String, Value>,
913 keys: &[&str],
914 ) -> Vec<(String, Value, Value)> {
915 keys.iter()
916 .filter_map(|&k| {
917 let old = prev.get(k);
918 let new = self.get_raw(k);
919 match (old, new) {
920 (Some(o), Some(n)) if o != &n => Some((k.to_string(), o.clone(), n)),
921 (None, Some(n)) => Some((k.to_string(), Value::Null, n)),
922 (Some(o), None) => Some((k.to_string(), o.clone(), Value::Null)),
923 _ => None,
924 }
925 })
926 .collect()
927 }
928
929 pub fn to_hashmap(&self) -> std::collections::HashMap<String, serde_json::Value> {
931 self.inner
932 .iter()
933 .map(|entry| (entry.key().clone(), entry.value().clone()))
934 .collect()
935 }
936
937 pub fn from_hashmap(&self, map: std::collections::HashMap<String, serde_json::Value>) {
939 for (key, value) in map {
940 let old = self.inner.insert(key.clone(), value.clone());
942 self.record_mutation(key, old, Some(value), StateMutationOrigin::SetCommitted);
943 }
944 }
945
946 pub fn clear_prefix(&self, prefix: &str) {
952 if self.track_delta {
953 let keys: Vec<String> = self
954 .keys()
955 .into_iter()
956 .filter(|k| k.starts_with(prefix))
957 .collect();
958 for key in keys {
959 let old = self.get_raw(&key);
960 self.delta.insert(key.clone(), DeltaOp::Delete);
961 if let Some(old) = old {
962 self.record_mutation(key, Some(old), None, StateMutationOrigin::ClearPrefix);
963 }
964 }
965 return;
966 }
967 let keys_to_remove: Vec<String> = self
968 .inner
969 .iter()
970 .filter(|entry| entry.key().starts_with(prefix))
971 .map(|entry| entry.key().clone())
972 .collect();
973 for key in keys_to_remove {
974 if let Some((_, old)) = self.inner.remove(&key) {
975 self.record_mutation(key, Some(old), None, StateMutationOrigin::ClearPrefix);
976 }
977 }
978 }
979
980 pub fn recent_mutations(&self) -> Vec<StateMutation> {
982 self.mutations
983 .lock()
984 .expect("state mutation journal poisoned")
985 .iter()
986 .cloned()
987 .collect()
988 }
989
990 pub fn mutation_cursor(&self) -> u64 {
992 self.next_mutation_sequence.load(Ordering::Relaxed) - 1
993 }
994
995 pub fn mutations_since(&self, cursor: u64) -> Vec<StateMutation> {
997 let mutations = self
998 .mutations
999 .lock()
1000 .expect("state mutation journal poisoned");
1001 mutations
1002 .iter()
1003 .filter(|mutation| mutation.sequence > cursor)
1004 .cloned()
1005 .collect()
1006 }
1007
1008 pub fn drain_mutations(&self) -> Vec<StateMutation> {
1010 self.mutations
1011 .lock()
1012 .expect("state mutation journal poisoned")
1013 .drain(..)
1014 .collect()
1015 }
1016
1017 pub fn evidence(&self, key: &str) -> SlotEvidence {
1020 let value = self.get_raw(key);
1021 let meta = self.get::<Value>(&format!("state_meta:{key}"));
1022 let source = meta
1023 .as_ref()
1024 .and_then(|m| m.get("source"))
1025 .and_then(|s| s.as_str().map(String::from));
1026 let confidence = meta
1027 .as_ref()
1028 .and_then(|m| m.get("confidence"))
1029 .and_then(serde_json::Value::as_f64);
1030
1031 let mut last_sequence: Option<u64> = None;
1032 let mut last_origin: Option<StateMutationOrigin> = None;
1033 for m in self.recent_mutations() {
1034 if m.key == key && last_sequence.is_none_or(|s| m.sequence > s) {
1035 last_sequence = Some(m.sequence);
1036 last_origin = Some(m.origin);
1037 }
1038 }
1039
1040 SlotEvidence {
1041 key: key.to_string(),
1042 present: value.is_some(),
1043 value,
1044 source,
1045 confidence,
1046 last_sequence,
1047 last_origin,
1048 }
1049 }
1050
1051 fn record_mutation(
1052 &self,
1053 key: String,
1054 old: Option<Value>,
1055 new: Option<Value>,
1056 origin: StateMutationOrigin,
1057 ) {
1058 self.record_mutation_with_delta(key, old, new, origin, self.track_delta);
1059 }
1060
1061 fn record_mutation_with_delta(
1062 &self,
1063 key: String,
1064 old: Option<Value>,
1065 new: Option<Value>,
1066 origin: StateMutationOrigin,
1067 delta: bool,
1068 ) {
1069 let mut mutations = self
1070 .mutations
1071 .lock()
1072 .expect("state mutation journal poisoned");
1073 if mutations.len() >= self.mutation_capacity {
1074 mutations.pop_front();
1075 }
1076 let sequence = self.next_mutation_sequence.fetch_add(1, Ordering::Relaxed);
1077 let mutation = StateMutation {
1078 sequence,
1079 key,
1080 old,
1081 new,
1082 origin,
1083 timestamp: SystemTime::now(),
1084 delta,
1085 };
1086 if let Some(sink) = self.journal_sink.0.read().as_ref() {
1089 sink.write(&mutation);
1090 }
1091 mutations.push_back(mutation);
1092 }
1093}
1094
1095pub struct PrefixedState<'a> {
1097 state: &'a State,
1098 prefix: &'static str,
1099}
1100
1101impl<'a> PrefixedState<'a> {
1102 fn prefixed_key(&self, key: &str) -> String {
1103 format!("{}{}", self.prefix, key)
1104 }
1105
1106 pub fn get<T: serde::de::DeserializeOwned>(&self, key: &str) -> Option<T> {
1108 self.state.get(&self.prefixed_key(key))
1109 }
1110
1111 pub fn get_raw(&self, key: &str) -> Option<Value> {
1113 self.state.get_raw(&self.prefixed_key(key))
1114 }
1115
1116 pub fn with<F, R>(&self, key: &str, f: F) -> Option<R>
1118 where
1119 F: FnOnce(&Value) -> R,
1120 {
1121 self.state.with(&self.prefixed_key(key), f)
1122 }
1123
1124 pub fn set(
1128 &self,
1129 key: impl AsRef<str>,
1130 value: impl serde::Serialize,
1131 ) -> Result<(), StateError> {
1132 self.state.set(self.prefixed_key(key.as_ref()), value)
1133 }
1134
1135 pub fn contains(&self, key: &str) -> bool {
1137 self.state.contains(&self.prefixed_key(key))
1138 }
1139
1140 pub fn remove(&self, key: &str) -> Option<Value> {
1142 self.state.remove(&self.prefixed_key(key))
1143 }
1144
1145 pub fn keys(&self) -> Vec<String> {
1147 self.state
1148 .keys()
1149 .into_iter()
1150 .filter_map(|k| {
1151 k.strip_prefix(self.prefix)
1152 .map(std::string::ToString::to_string)
1153 })
1154 .collect()
1155 }
1156}
1157
1158pub struct ReadOnlyPrefixedState<'a> {
1163 state: &'a State,
1164 prefix: &'static str,
1165}
1166
1167impl<'a> ReadOnlyPrefixedState<'a> {
1168 fn prefixed_key(&self, key: &str) -> String {
1169 format!("{}{}", self.prefix, key)
1170 }
1171
1172 pub fn get<T: serde::de::DeserializeOwned>(&self, key: &str) -> Option<T> {
1174 self.state.get(&self.prefixed_key(key))
1175 }
1176
1177 pub fn get_raw(&self, key: &str) -> Option<Value> {
1179 self.state.get_raw(&self.prefixed_key(key))
1180 }
1181
1182 pub fn with<F, R>(&self, key: &str, f: F) -> Option<R>
1184 where
1185 F: FnOnce(&Value) -> R,
1186 {
1187 self.state.with(&self.prefixed_key(key), f)
1188 }
1189
1190 pub fn contains(&self, key: &str) -> bool {
1192 self.state.contains(&self.prefixed_key(key))
1193 }
1194
1195 pub fn keys(&self) -> Vec<String> {
1197 self.state
1198 .keys()
1199 .into_iter()
1200 .filter_map(|k| {
1201 k.strip_prefix(self.prefix)
1202 .map(std::string::ToString::to_string)
1203 })
1204 .collect()
1205 }
1206}
1207
1208#[cfg(test)]
1209mod tests {
1210 use super::*;
1211
1212 #[test]
1213 fn journal_sink_receives_every_mutation_in_ring_order() {
1214 let state = State::new();
1215 let sink = Arc::new(MemoryJournalSink::new());
1216 state.set_journal_sink(sink.clone());
1217
1218 let _ = state.set("a", 1);
1219 let _ = state.set("b", "two");
1220 state.remove("a");
1221
1222 let entries = sink.entries();
1223 assert_eq!(entries.len(), 3);
1224 assert_eq!(entries, state.recent_mutations());
1225 assert_eq!(entries[0].key, "a");
1226 assert_eq!(entries[2].origin, StateMutationOrigin::Remove);
1227 }
1228
1229 #[test]
1230 fn journal_sink_is_shared_with_clones_and_delta_views() {
1231 let state = State::new();
1232 let sink = Arc::new(MemoryJournalSink::new());
1233 state.set_journal_sink(sink.clone());
1234
1235 let clone = state.clone();
1236 let _ = clone.set("from_clone", true);
1237
1238 let tracked = state.with_delta_tracking();
1239 let _ = tracked.set("from_delta", 1);
1240 tracked.commit();
1241
1242 let keys: Vec<_> = sink.entries().iter().map(|m| m.key.clone()).collect();
1243 assert!(keys.contains(&"from_clone".to_string()));
1244 assert!(keys.contains(&"from_delta".to_string()));
1245 assert!(
1247 sink.entries()
1248 .iter()
1249 .any(|m| m.origin == StateMutationOrigin::Commit)
1250 );
1251 }
1252
1253 #[test]
1254 fn journal_sink_outlives_ring_capacity() {
1255 let state = State::new();
1257 let sink = Arc::new(MemoryJournalSink::new());
1258 state.set_journal_sink(sink.clone());
1259
1260 for i in 0..(DEFAULT_MUTATION_JOURNAL_CAPACITY + 10) {
1261 let _ = state.set(format!("k{i}"), i);
1262 }
1263
1264 assert_eq!(
1265 state.recent_mutations().len(),
1266 DEFAULT_MUTATION_JOURNAL_CAPACITY
1267 );
1268 assert_eq!(sink.len(), DEFAULT_MUTATION_JOURNAL_CAPACITY + 10);
1269 assert_eq!(sink.entries()[0].key, "k0");
1270 }
1271
1272 #[test]
1273 fn state_mutation_serde_round_trip_uses_epoch_millis() {
1274 let m = StateMutation {
1275 sequence: 42,
1276 key: "app:last_city".into(),
1277 old: None,
1278 new: Some(serde_json::json!("London")),
1279 origin: StateMutationOrigin::Set,
1280 timestamp: std::time::UNIX_EPOCH + std::time::Duration::from_millis(1_718_000_000_123),
1281 delta: false,
1282 };
1283 let json = serde_json::to_string(&m).unwrap();
1284 assert!(json.contains("\"timestamp_ms\":1718000000123"));
1285 assert!(json.contains("\"origin\":\"set\""));
1286 let back: StateMutation = serde_json::from_str(&json).unwrap();
1287 assert_eq!(back, m);
1288 }
1289
1290 #[test]
1291 fn file_journal_sink_round_trip() {
1292 let dir = std::env::temp_dir().join(format!(
1293 "gemini-rs-journal-test-{}-{}",
1294 std::process::id(),
1295 std::time::SystemTime::now()
1296 .duration_since(std::time::UNIX_EPOCH)
1297 .unwrap()
1298 .as_nanos()
1299 ));
1300 std::fs::create_dir_all(&dir).unwrap();
1301 let path = dir.join("session.journal.jsonl");
1302
1303 let state = State::new();
1304 {
1305 let sink = Arc::new(FileJournalSink::create(&path).unwrap());
1306 state.set_journal_sink(sink);
1307 let _ = state.set("a", 1);
1308 let _ = state.set("a", 2);
1309 state.remove("a");
1310 state.set_journal_sink(Arc::new(MemoryJournalSink::new()));
1312 }
1313
1314 let data = std::fs::read_to_string(&path).unwrap();
1315 let parsed: Vec<StateMutation> = data
1316 .lines()
1317 .filter(|l| !l.trim().is_empty())
1318 .map(|l| serde_json::from_str(l).unwrap())
1319 .collect();
1320 assert_eq!(parsed.len(), 3);
1321 assert_eq!(parsed[0].new, Some(serde_json::json!(1)));
1322 assert_eq!(parsed[2].origin, StateMutationOrigin::Remove);
1323
1324 let _ = std::fs::remove_dir_all(&dir);
1325 }
1326
1327 #[test]
1328 fn set_and_get_string() {
1329 let state = State::new();
1330 let _ = state.set("name", "Alice");
1331 assert_eq!(state.get::<String>("name"), Some("Alice".to_string()));
1332 }
1333
1334 #[test]
1335 fn set_and_get_json() {
1336 let state = State::new();
1337 let _ = state.set("data", serde_json::json!({"temp": 22}));
1338 let v: Value = state.get("data").unwrap();
1339 assert_eq!(v["temp"], 22);
1340 }
1341
1342 #[test]
1343 fn pick_subset() {
1344 let state = State::new();
1345 let _ = state.set("a", 1);
1346 let _ = state.set("b", 2);
1347 let _ = state.set("c", 3);
1348 let picked = state.pick(&["a", "c"]);
1349 assert!(picked.contains("a"));
1350 assert!(!picked.contains("b"));
1351 assert!(picked.contains("c"));
1352 }
1353
1354 #[test]
1355 fn merge_states() {
1356 let s1 = State::new();
1357 let _ = s1.set("a", 1);
1358 let s2 = State::new();
1359 let _ = s2.set("b", 2);
1360 s1.merge(&s2);
1361 assert!(s1.contains("a"));
1362 assert!(s1.contains("b"));
1363 }
1364
1365 #[test]
1366 fn rename_key() {
1367 let state = State::new();
1368 let _ = state.set("old", "value");
1369 state.rename("old", "new");
1370 assert!(!state.contains("old"));
1371 assert_eq!(state.get::<String>("new"), Some("value".to_string()));
1372 }
1373
1374 #[test]
1375 fn remove_returns_value() {
1376 let state = State::new();
1377 let _ = state.set("key", 42);
1378 let removed = state.remove("key");
1379 assert!(removed.is_some());
1380 assert!(!state.contains("key"));
1381 }
1382
1383 #[test]
1384 fn get_missing_returns_none() {
1385 let state = State::new();
1386 assert_eq!(state.get::<String>("nope"), None);
1387 }
1388
1389 #[test]
1392 fn delta_tracking_writes_to_delta() {
1393 let state = State::new();
1394 let _ = state.set("committed", "yes");
1395
1396 let tracked = state.with_delta_tracking();
1397 let _ = tracked.set("new_key", "new_value");
1398
1399 assert_eq!(
1401 tracked.get::<String>("new_key"),
1402 Some("new_value".to_string())
1403 );
1404 assert!(!state.contains("new_key"));
1406 assert_eq!(tracked.get::<String>("committed"), Some("yes".to_string()));
1408 }
1409
1410 #[test]
1411 fn delta_has_delta_reports_correctly() {
1412 let state = State::new();
1413 let tracked = state.with_delta_tracking();
1414 assert!(!tracked.has_delta());
1415
1416 let _ = tracked.set("key", "val");
1417 assert!(tracked.has_delta());
1418 }
1419
1420 #[test]
1421 fn delta_commit_merges_to_inner() {
1422 let state = State::new();
1423 let tracked = state.with_delta_tracking();
1424 let _ = tracked.set("key", "val");
1425 assert!(!state.contains("key"));
1426
1427 tracked.commit();
1428 assert_eq!(state.get::<String>("key"), Some("val".to_string()));
1430 assert!(!tracked.has_delta());
1431 }
1432
1433 #[test]
1434 fn delta_rollback_discards_changes() {
1435 let state = State::new();
1436 let tracked = state.with_delta_tracking();
1437 let _ = tracked.set("key", "val");
1438 assert!(tracked.has_delta());
1439
1440 tracked.rollback();
1441 assert!(!tracked.has_delta());
1442 assert!(!state.contains("key"));
1443 assert!(!tracked.contains("key"));
1444 }
1445
1446 #[test]
1447 fn delta_snapshot() {
1448 let state = State::new();
1449 let tracked = state.with_delta_tracking();
1450 let _ = tracked.set("a", 1);
1451 let _ = tracked.set("b", 2);
1452
1453 let snapshot = tracked.delta();
1454 assert_eq!(snapshot.len(), 2);
1455 assert!(snapshot.contains_key("a"));
1456 assert!(snapshot.contains_key("b"));
1457 }
1458
1459 #[test]
1460 fn set_committed_bypasses_delta() {
1461 let state = State::new();
1462 let tracked = state.with_delta_tracking();
1463 let _ = tracked.set_committed("direct", "value");
1464
1465 assert_eq!(state.get::<String>("direct"), Some("value".to_string()));
1467 assert!(!tracked.has_delta());
1469 assert_eq!(tracked.get::<String>("direct"), Some("value".to_string()));
1471 }
1472
1473 #[test]
1474 fn mutation_journal_records_set_and_remove() {
1475 let state = State::new();
1476 let _ = state.set("key", "first");
1477 let _ = state.set("key", "second");
1478 state.remove("key");
1479
1480 let mutations = state.recent_mutations();
1481 assert_eq!(mutations.len(), 3);
1482 assert_eq!(mutations[0].key, "key");
1483 assert_eq!(mutations[0].old, None);
1484 assert_eq!(mutations[0].new, Some(serde_json::json!("first")));
1485 assert_eq!(mutations[0].origin, StateMutationOrigin::Set);
1486
1487 assert_eq!(mutations[1].old, Some(serde_json::json!("first")));
1488 assert_eq!(mutations[1].new, Some(serde_json::json!("second")));
1489
1490 assert_eq!(mutations[2].old, Some(serde_json::json!("second")));
1491 assert_eq!(mutations[2].new, None);
1492 assert_eq!(mutations[2].origin, StateMutationOrigin::Remove);
1493 }
1494
1495 #[test]
1496 fn mutation_journal_is_shared_with_delta_tracking() {
1497 let state = State::new();
1498 let _ = state.set("committed", "yes");
1499
1500 let tracked = state.with_delta_tracking();
1501 let _ = tracked.set("committed", "maybe");
1502 tracked.commit();
1503
1504 let mutations = state.recent_mutations();
1505 assert_eq!(mutations.len(), 3);
1506 assert_eq!(mutations[1].key, "committed");
1507 assert_eq!(mutations[1].old, Some(serde_json::json!("yes")));
1508 assert_eq!(mutations[1].new, Some(serde_json::json!("maybe")));
1509 assert_eq!(mutations[1].origin, StateMutationOrigin::Set);
1510 assert!(mutations[1].delta);
1511
1512 assert_eq!(mutations[2].origin, StateMutationOrigin::Commit);
1513 assert!(!mutations[2].delta);
1514 }
1515
1516 #[test]
1517 fn drain_mutations_clears_journal() {
1518 let state = State::new();
1519 let _ = state.set("a", 1);
1520 let _ = state.set("b", 2);
1521
1522 let drained = state.drain_mutations();
1523 assert_eq!(drained.len(), 2);
1524 assert!(state.recent_mutations().is_empty());
1525 }
1526
1527 #[test]
1528 fn mutation_cursor_reads_only_later_changes() {
1529 let state = State::new();
1530 let _ = state.set("before", 1);
1531 let cursor = state.mutation_cursor();
1532
1533 let _ = state.set("after", 2);
1534 state.remove("before");
1535
1536 let mutations = state.mutations_since(cursor);
1537 assert_eq!(mutations.len(), 2);
1538 assert_eq!(mutations[0].key, "after");
1539 assert_eq!(mutations[1].key, "before");
1540 }
1541
1542 #[test]
1543 fn no_delta_tracking_preserves_existing_behavior() {
1544 let state = State::new();
1545 assert!(!state.is_tracking_delta());
1546 let _ = state.set("key", "val");
1547 assert_eq!(state.get::<String>("key"), Some("val".to_string()));
1548 assert!(!state.has_delta());
1549 }
1550
1551 #[test]
1554 fn prefix_app_set_and_get() {
1555 let state = State::new();
1556 let _ = state.app().set("flag", true);
1557
1558 assert_eq!(state.app().get::<bool>("flag"), Some(true));
1560 assert_eq!(state.get::<bool>("app:flag"), Some(true));
1562 }
1563
1564 #[test]
1565 fn prefix_user_set_and_get() {
1566 let state = State::new();
1567 let _ = state.user().set("name", "Alice");
1568 assert_eq!(
1569 state.user().get::<String>("name"),
1570 Some("Alice".to_string())
1571 );
1572 assert_eq!(state.get::<String>("user:name"), Some("Alice".to_string()));
1573 }
1574
1575 #[test]
1576 fn prefix_temp_set_and_get() {
1577 let state = State::new();
1578 let _ = state.temp().set("scratch", 42);
1579 assert_eq!(state.temp().get::<i32>("scratch"), Some(42));
1580 }
1581
1582 #[test]
1583 fn prefix_contains_and_remove() {
1584 let state = State::new();
1585 let _ = state.app().set("x", 1);
1586 assert!(state.app().contains("x"));
1587 state.app().remove("x");
1588 assert!(!state.app().contains("x"));
1589 }
1590
1591 #[test]
1592 fn prefix_keys() {
1593 let state = State::new();
1594 let _ = state.app().set("a", 1);
1595 let _ = state.app().set("b", 2);
1596 let _ = state.user().set("c", 3);
1597
1598 let app_keys = state.app().keys();
1599 assert_eq!(app_keys.len(), 2);
1600 assert!(app_keys.contains(&"a".to_string()));
1601 assert!(app_keys.contains(&"b".to_string()));
1602
1603 let user_keys = state.user().keys();
1604 assert_eq!(user_keys.len(), 1);
1605 assert!(user_keys.contains(&"c".to_string()));
1606 }
1607
1608 #[test]
1609 fn prefix_with_delta_tracking() {
1610 let state = State::new();
1611 let tracked = state.with_delta_tracking();
1612 let _ = tracked.app().set("flag", true);
1613
1614 assert_eq!(tracked.app().get::<bool>("flag"), Some(true));
1616 assert!(tracked.has_delta());
1618 assert!(!state.contains("app:flag"));
1619
1620 tracked.commit();
1621 assert_eq!(state.get::<bool>("app:flag"), Some(true));
1622 }
1623
1624 #[test]
1627 fn prefix_session_set_and_get() {
1628 let state = State::new();
1629 let _ = state.session().set("turn_count", 5);
1630 assert_eq!(state.session().get::<i32>("turn_count"), Some(5));
1631 assert_eq!(state.get::<i32>("session:turn_count"), Some(5));
1632 }
1633
1634 #[test]
1635 fn prefix_turn_set_and_get() {
1636 let state = State::new();
1637 let _ = state.turn().set("transcript", "hello");
1638 assert_eq!(
1639 state.turn().get::<String>("transcript"),
1640 Some("hello".to_string())
1641 );
1642 assert_eq!(
1643 state.get::<String>("turn:transcript"),
1644 Some("hello".to_string())
1645 );
1646 }
1647
1648 #[test]
1649 fn prefix_bg_set_and_get() {
1650 let state = State::new();
1651 let _ = state.bg().set("task_id", "abc-123");
1652 assert_eq!(
1653 state.bg().get::<String>("task_id"),
1654 Some("abc-123".to_string())
1655 );
1656 assert_eq!(
1657 state.get::<String>("bg:task_id"),
1658 Some("abc-123".to_string())
1659 );
1660 }
1661
1662 #[test]
1663 fn prefix_session_contains_and_remove() {
1664 let state = State::new();
1665 let _ = state.session().set("x", 1);
1666 assert!(state.session().contains("x"));
1667 state.session().remove("x");
1668 assert!(!state.session().contains("x"));
1669 }
1670
1671 #[test]
1672 fn prefix_turn_keys() {
1673 let state = State::new();
1674 let _ = state.turn().set("a", 1);
1675 let _ = state.turn().set("b", 2);
1676 let _ = state.session().set("c", 3);
1677
1678 let turn_keys = state.turn().keys();
1679 assert_eq!(turn_keys.len(), 2);
1680 assert!(turn_keys.contains(&"a".to_string()));
1681 assert!(turn_keys.contains(&"b".to_string()));
1682 }
1683
1684 #[test]
1687 fn try_get_distinguishes_absent_from_wrong_type() {
1688 let state = State::new();
1689 assert!(matches!(state.try_get::<u32>("missing"), Ok(None)));
1690
1691 state.set("n", 5u32).unwrap();
1692 assert_eq!(state.try_get::<u32>("n").unwrap(), Some(5));
1693
1694 state.set("s", "not a number").unwrap();
1695 assert_eq!(state.get::<u32>("s"), None);
1697 match state.try_get::<u32>("s") {
1699 Err(StateError::WrongType { key, .. }) => assert_eq!(key, "s"),
1700 other => panic!("expected WrongType, got {other:?}"),
1701 }
1702
1703 state.set("derived:risk", 0.5f64).unwrap();
1705 assert_eq!(state.try_get::<f64>("risk").unwrap(), Some(0.5));
1706
1707 const N: StateKey<u32> = StateKey::new("n");
1708 assert_eq!(state.try_get_key(&N).unwrap(), Some(5));
1709 const S: StateKey<u32> = StateKey::new("s");
1710 assert!(state.try_get_key(&S).is_err());
1711 }
1712
1713 #[test]
1716 fn derived_read_only_get() {
1717 let state = State::new();
1718 let _ = state.set("derived:sentiment", "positive");
1720 assert_eq!(
1721 state.derived().get::<String>("sentiment"),
1722 Some("positive".to_string())
1723 );
1724 }
1725
1726 #[test]
1727 fn derived_read_only_get_raw() {
1728 let state = State::new();
1729 let _ = state.set("derived:score", serde_json::json!(0.95));
1730 let raw = state.derived().get_raw("score");
1731 assert!(raw.is_some());
1732 assert_eq!(raw.unwrap(), serde_json::json!(0.95));
1733 }
1734
1735 #[test]
1736 fn derived_read_only_contains() {
1737 let state = State::new();
1738 let _ = state.set("derived:exists", true);
1739 assert!(state.derived().contains("exists"));
1740 assert!(!state.derived().contains("missing"));
1741 }
1742
1743 #[test]
1744 fn derived_read_only_keys() {
1745 let state = State::new();
1746 let _ = state.set("derived:a", 1);
1747 let _ = state.set("derived:b", 2);
1748 let _ = state.set("app:c", 3);
1749
1750 let derived_keys = state.derived().keys();
1751 assert_eq!(derived_keys.len(), 2);
1752 assert!(derived_keys.contains(&"a".to_string()));
1753 assert!(derived_keys.contains(&"b".to_string()));
1754 }
1755
1756 #[test]
1757 fn derived_missing_key_returns_none() {
1758 let state = State::new();
1759 assert_eq!(state.derived().get::<String>("nope"), None);
1760 assert_eq!(state.derived().get_raw("nope"), None);
1761 }
1762
1763 #[test]
1766 fn snapshot_values_captures_existing_keys() {
1767 let state = State::new();
1768 let _ = state.set("a", 1);
1769 let _ = state.set("b", "hello");
1770 let _ = state.set("c", true);
1771
1772 let snap = state.snapshot_values(&["a", "b", "missing"]);
1773 assert_eq!(snap.len(), 2);
1774 assert_eq!(snap.get("a"), Some(&serde_json::json!(1)));
1775 assert_eq!(snap.get("b"), Some(&serde_json::json!("hello")));
1776 assert!(!snap.contains_key("missing"));
1777 }
1778
1779 #[test]
1780 fn snapshot_values_empty_keys() {
1781 let state = State::new();
1782 let _ = state.set("a", 1);
1783 let snap = state.snapshot_values(&[]);
1784 assert!(snap.is_empty());
1785 }
1786
1787 #[test]
1790 fn diff_values_detects_changed_value() {
1791 let state = State::new();
1792 let _ = state.set("x", 1);
1793 let snap = state.snapshot_values(&["x"]);
1794
1795 let _ = state.set("x", 2);
1796 let diffs = state.diff_values(&snap, &["x"]);
1797 assert_eq!(diffs.len(), 1);
1798 assert_eq!(diffs[0].0, "x");
1799 assert_eq!(diffs[0].1, serde_json::json!(1));
1800 assert_eq!(diffs[0].2, serde_json::json!(2));
1801 }
1802
1803 #[test]
1804 fn diff_values_detects_new_key() {
1805 let state = State::new();
1806 let snap = state.snapshot_values(&["y"]);
1807
1808 let _ = state.set("y", "new");
1809 let diffs = state.diff_values(&snap, &["y"]);
1810 assert_eq!(diffs.len(), 1);
1811 assert_eq!(diffs[0].0, "y");
1812 assert_eq!(diffs[0].1, Value::Null);
1813 assert_eq!(diffs[0].2, serde_json::json!("new"));
1814 }
1815
1816 #[test]
1817 fn diff_values_detects_removed_key() {
1818 let state = State::new();
1819 let _ = state.set("z", 42);
1820 let snap = state.snapshot_values(&["z"]);
1821
1822 state.remove("z");
1823 let diffs = state.diff_values(&snap, &["z"]);
1824 assert_eq!(diffs.len(), 1);
1825 assert_eq!(diffs[0].0, "z");
1826 assert_eq!(diffs[0].1, serde_json::json!(42));
1827 assert_eq!(diffs[0].2, Value::Null);
1828 }
1829
1830 #[test]
1831 fn diff_values_no_change() {
1832 let state = State::new();
1833 let _ = state.set("stable", 10);
1834 let snap = state.snapshot_values(&["stable"]);
1835
1836 let diffs = state.diff_values(&snap, &["stable"]);
1838 assert!(diffs.is_empty());
1839 }
1840
1841 #[test]
1842 fn diff_values_multiple_keys_mixed_changes() {
1843 let state = State::new();
1844 let _ = state.set("a", 1);
1845 let _ = state.set("b", 2);
1846 let snap = state.snapshot_values(&["a", "b", "c"]);
1847
1848 let _ = state.set("a", 10); let _ = state.set("c", 3); let diffs = state.diff_values(&snap, &["a", "b", "c"]);
1853 assert_eq!(diffs.len(), 2); let diff_keys: Vec<&str> = diffs.iter().map(|(k, _, _)| k.as_str()).collect();
1855 assert!(diff_keys.contains(&"a"));
1856 assert!(diff_keys.contains(&"c"));
1857 }
1858
1859 #[test]
1862 fn clear_prefix_removes_matching_keys() {
1863 let state = State::new();
1864 let _ = state.set("turn:a", 1);
1865 let _ = state.set("turn:b", 2);
1866 let _ = state.set("app:c", 3);
1867 let _ = state.set("session:d", 4);
1868
1869 state.clear_prefix("turn:");
1870 assert!(!state.contains("turn:a"));
1871 assert!(!state.contains("turn:b"));
1872 assert!(state.contains("app:c"));
1873 assert!(state.contains("session:d"));
1874 }
1875
1876 #[test]
1877 fn clear_prefix_no_matching_keys_is_noop() {
1878 let state = State::new();
1879 let _ = state.set("app:x", 1);
1880 state.clear_prefix("turn:");
1881 assert!(state.contains("app:x"));
1882 }
1883
1884 #[test]
1885 fn clear_prefix_also_clears_delta() {
1886 let state = State::new();
1887 let _ = state.set("turn:committed", 1);
1888 let tracked = state.with_delta_tracking();
1889 let _ = tracked.set("turn:delta_val", 2);
1890
1891 assert!(tracked.contains("turn:committed"));
1893 assert!(tracked.contains("turn:delta_val"));
1894
1895 tracked.clear_prefix("turn:");
1896 assert!(!tracked.contains("turn:committed"));
1897 assert!(!tracked.contains("turn:delta_val"));
1898 }
1899
1900 #[test]
1901 fn clear_prefix_via_turn_accessor() {
1902 let state = State::new();
1903 let _ = state.turn().set("x", 1);
1904 let _ = state.turn().set("y", 2);
1905 let _ = state.app().set("z", 3);
1906
1907 state.clear_prefix("turn:");
1908 assert!(state.turn().keys().is_empty());
1909 assert!(state.app().contains("z"));
1910 }
1911
1912 #[test]
1915 fn modify_increment_existing() {
1916 let state = State::new();
1917 let _ = state.set("count", 5u32);
1918 let result = state.modify("count", 0u32, |n| n + 1).unwrap();
1919 assert_eq!(result, 6);
1920 assert_eq!(state.get::<u32>("count"), Some(6));
1921 }
1922
1923 #[test]
1924 fn modify_uses_default_when_missing() {
1925 let state = State::new();
1926 let result = state.modify("new_count", 0u32, |n| n + 1).unwrap();
1927 assert_eq!(result, 1);
1928 assert_eq!(state.get::<u32>("new_count"), Some(1));
1929 }
1930
1931 #[test]
1932 fn modify_with_delta_tracking() {
1933 let state = State::new();
1934 let _ = state.set("x", 10u32);
1935 let tracked = state.with_delta_tracking();
1936 let result = tracked.modify("x", 0u32, |n| n * 2).unwrap();
1937 assert_eq!(result, 20);
1938 assert_eq!(tracked.get::<u32>("x"), Some(20));
1940 assert_eq!(state.get::<u32>("x"), Some(10)); }
1942
1943 #[test]
1946 fn get_falls_back_to_derived_prefix() {
1947 let state = State::new();
1948 let _ = state.set("derived:risk", 0.85);
1949 assert_eq!(state.get::<f64>("risk"), Some(0.85));
1951 }
1952
1953 #[test]
1954 fn get_prefers_direct_key_over_derived() {
1955 let state = State::new();
1956 let _ = state.set("score", 1.0);
1957 let _ = state.set("derived:score", 0.5);
1958 assert_eq!(state.get::<f64>("score"), Some(1.0));
1960 }
1961
1962 #[test]
1963 fn get_derived_fallback_skipped_for_prefixed_keys() {
1964 let state = State::new();
1965 let _ = state.set("derived:risk", 0.85);
1966 assert_eq!(state.get::<f64>("app:risk"), None);
1968 }
1969
1970 #[test]
1971 fn get_derived_fallback_with_delta_tracking() {
1972 let state = State::new();
1973 let tracked = state.with_delta_tracking();
1974 let _ = tracked.set("derived:computed_val", 42);
1975 assert_eq!(tracked.get::<i32>("computed_val"), Some(42));
1976 }
1977
1978 #[test]
1981 fn with_reads_from_inner() {
1982 let state = State::new();
1983 let _ = state.set("name", "Alice");
1984 let len = state.with("name", |v| v.as_str().unwrap().len());
1985 assert_eq!(len, Some(5));
1986 }
1987
1988 #[test]
1989 fn with_reads_from_delta_first() {
1990 let state = State::new();
1991 let _ = state.set("x", 1);
1992 let tracked = state.with_delta_tracking();
1993 let _ = tracked.set("x", 99);
1994 let val = tracked.with("x", |v| v.as_i64().unwrap());
1995 assert_eq!(val, Some(99));
1996 }
1997
1998 #[test]
1999 fn with_falls_back_to_inner_when_not_in_delta() {
2000 let state = State::new();
2001 let _ = state.set("committed", "yes");
2002 let tracked = state.with_delta_tracking();
2003 let val = tracked.with("committed", |v| v.as_str().unwrap().to_string());
2004 assert_eq!(val, Some("yes".to_string()));
2005 }
2006
2007 #[test]
2008 fn with_falls_back_to_derived() {
2009 let state = State::new();
2010 let _ = state.set("derived:risk", 0.85);
2011 let val = state.with("risk", |v| v.as_f64().unwrap());
2012 assert_eq!(val, Some(0.85));
2013 }
2014
2015 #[test]
2016 fn with_derived_fallback_skipped_for_prefixed() {
2017 let state = State::new();
2018 let _ = state.set("derived:risk", 0.85);
2019 let val = state.with("app:risk", |v| v.as_f64().unwrap());
2020 assert_eq!(val, None);
2021 }
2022
2023 #[test]
2024 fn with_returns_none_for_missing() {
2025 let state = State::new();
2026 let val = state.with("missing", std::clone::Clone::clone);
2027 assert_eq!(val, None);
2028 }
2029
2030 #[test]
2031 fn with_on_prefixed_state() {
2032 let state = State::new();
2033 let _ = state.app().set("flag", true);
2034 let val = state.app().with("flag", |v| v.as_bool().unwrap());
2035 assert_eq!(val, Some(true));
2036 }
2037
2038 #[test]
2039 fn with_on_read_only_prefixed_state() {
2040 let state = State::new();
2041 let _ = state.set("derived:score", serde_json::json!(0.95));
2042 let val = state.derived().with("score", |v| v.as_f64().unwrap());
2043 assert_eq!(val, Some(0.95));
2044 }
2045
2046 const TURN_COUNT: StateKey<u32> = StateKey::new("session:turn_count");
2049 const NAME: StateKey<String> = StateKey::new("user:name");
2050
2051 #[test]
2052 fn state_key_get_and_set() {
2053 let state = State::new();
2054 let _ = state.set_key(&TURN_COUNT, 5);
2055 assert_eq!(state.get_key(&TURN_COUNT), Some(5));
2056 }
2057
2058 #[test]
2059 fn state_key_get_missing() {
2060 let state = State::new();
2061 assert_eq!(state.get_key(&TURN_COUNT), None);
2062 }
2063
2064 #[test]
2065 fn state_key_string_type() {
2066 let state = State::new();
2067 let _ = state.set_key(&NAME, "Alice".to_string());
2068 assert_eq!(state.get_key(&NAME), Some("Alice".to_string()));
2069 }
2070
2071 #[test]
2072 fn state_key_with() {
2073 let state = State::new();
2074 let _ = state.set_key(&TURN_COUNT, 42);
2075 let val = state.with_key(&TURN_COUNT, |v| v.as_u64().unwrap());
2076 assert_eq!(val, Some(42));
2077 }
2078
2079 #[test]
2080 fn state_key_interop_with_raw() {
2081 let state = State::new();
2082 let _ = state.set_key(&TURN_COUNT, 10);
2083 assert_eq!(state.get::<u32>("session:turn_count"), Some(10));
2085 }
2086
2087 #[test]
2088 fn slot_evidence_aggregates_value_provenance_and_journal() {
2089 let state = State::new();
2090 let _ = state.set("party_size", 6u8);
2091 let _ = state.set(
2093 "state_meta:party_size",
2094 serde_json::json!({ "source": "extraction", "confidence": 0.9 }),
2095 );
2096
2097 let ev = state.evidence("party_size");
2098 assert!(ev.present);
2099 assert_eq!(ev.value, Some(serde_json::json!(6)));
2100 assert_eq!(ev.source.as_deref(), Some("extraction"));
2101 assert_eq!(ev.confidence, Some(0.9));
2102 assert!(ev.last_sequence.is_some());
2103 assert_eq!(ev.last_origin, Some(StateMutationOrigin::Set));
2104
2105 let missing = state.evidence("nope");
2107 assert!(!missing.present);
2108 assert!(missing.source.is_none());
2109 }
2110
2111 #[test]
2114 fn rollback_restores_base_after_remove() {
2115 let base = State::new();
2118 let _ = base.set("k", "original");
2119
2120 let tx = base.with_delta_tracking();
2121 assert_eq!(tx.remove("k"), Some(serde_json::json!("original")));
2122 assert_eq!(tx.get::<String>("k"), None); assert_eq!(base.get::<String>("k"), Some("original".into())); tx.rollback();
2126 assert_eq!(tx.get::<String>("k"), Some("original".into()));
2127 assert_eq!(base.get::<String>("k"), Some("original".into()));
2128 }
2129
2130 #[test]
2131 fn rollback_restores_base_after_clear_prefix() {
2132 let base = State::new();
2134 let _ = base.set("app:a", 1u32);
2135 let _ = base.set("app:b", 2u32);
2136 let _ = base.set("user:c", 3u32);
2137
2138 let tx = base.with_delta_tracking();
2139 tx.clear_prefix("app:");
2140 assert_eq!(tx.get::<u32>("app:a"), None);
2141 assert_eq!(tx.get::<u32>("app:b"), None);
2142 assert_eq!(tx.get::<u32>("user:c"), Some(3));
2143 assert_eq!(base.get::<u32>("app:a"), Some(1));
2145
2146 tx.rollback();
2147 assert_eq!(tx.get::<u32>("app:a"), Some(1));
2148 assert_eq!(tx.get::<u32>("app:b"), Some(2));
2149 }
2150
2151 #[test]
2152 fn commit_applies_removals() {
2153 let base = State::new();
2154 let _ = base.set("k", "v");
2155 let tx = base.with_delta_tracking();
2156 tx.remove("k");
2157 tx.commit();
2158 assert_eq!(base.get::<String>("k"), None);
2159 }
2160
2161 #[test]
2162 fn commit_applies_prefix_clear() {
2163 let base = State::new();
2164 let _ = base.set("app:a", 1u32);
2165 let _ = base.set("user:c", 3u32);
2166 let tx = base.with_delta_tracking();
2167 tx.clear_prefix("app:");
2168 tx.commit();
2169 assert_eq!(base.get::<u32>("app:a"), None);
2170 assert_eq!(base.get::<u32>("user:c"), Some(3));
2171 }
2172
2173 #[test]
2174 fn modify_is_atomic_under_concurrency() {
2175 use std::sync::Arc;
2176 use std::thread;
2177
2178 let state = Arc::new(State::new());
2179 let _ = state.set("count", 0u64);
2180
2181 let threads = 8;
2182 let per_thread = 1000;
2183 let handles: Vec<_> = (0..threads)
2184 .map(|_| {
2185 let state = state.clone();
2186 thread::spawn(move || {
2187 for _ in 0..per_thread {
2188 let _ = state.modify("count", 0u64, |n| n + 1);
2189 }
2190 })
2191 })
2192 .collect();
2193 for h in handles {
2194 h.join().unwrap();
2195 }
2196 assert_eq!(
2198 state.get::<u64>("count"),
2199 Some((threads * per_thread) as u64)
2200 );
2201 }
2202}
2203
2204#[cfg(test)]
2205mod proptests {
2206 use super::*;
2207 use proptest::prelude::*;
2208
2209 proptest! {
2212 #[test]
2213 fn rollback_always_restores_base(
2214 base_keys in proptest::collection::vec(("[a-c]", 0u32..5), 0..6),
2215 ops in proptest::collection::vec(
2216 prop_oneof![
2217 ("[a-c]", 0u32..5).prop_map(|(k, v)| (k, Some(v))),
2218 "[a-c]".prop_map(|k| (k, None)),
2219 ],
2220 0..12,
2221 ),
2222 ) {
2223 let base = State::new();
2224 for (k, v) in &base_keys {
2225 let _ = base.set(k.clone(), *v);
2226 }
2227 let snapshot = |s: &State| -> std::collections::BTreeMap<String, Value> {
2228 s.keys().into_iter().filter_map(|k| s.get_raw(&k).map(|v| (k, v))).collect()
2229 };
2230 let before = snapshot(&base);
2231
2232 let tx = base.with_delta_tracking();
2233 for (k, v) in &ops {
2234 match v {
2235 Some(v) => { let _ = tx.set(k.clone(), *v); }
2236 None => { tx.remove(k); }
2237 }
2238 }
2239 prop_assert_eq!(&before, &snapshot(&base));
2241
2242 tx.rollback();
2243 prop_assert_eq!(&before, &snapshot(&tx));
2244 }
2245 }
2246}
2247
2248#[cfg(test)]
2249mod derived_contains_fallback {
2250 use super::State;
2255
2256 #[test]
2257 fn contains_sees_a_derived_value_through_the_unprefixed_key() {
2258 let state = State::new();
2259 state.set("derived:risk", 0.85).unwrap();
2260 assert_eq!(
2261 state.get::<f64>("risk"),
2262 Some(0.85),
2263 "precondition: get falls back"
2264 );
2265 assert!(
2266 state.contains("risk"),
2267 "contains must fall back the same way get does"
2268 );
2269 }
2270
2271 #[test]
2272 fn contains_fallback_respects_delta_tracking_and_tombstones() {
2273 let state = State::new();
2274 state.set("derived:score", 1u32).unwrap();
2275 let tracked = state.with_delta_tracking();
2276 assert!(
2277 tracked.contains("score"),
2278 "inner derived value visible through tracked view"
2279 );
2280 tracked.remove("derived:score");
2281 assert!(
2282 !tracked.contains("score"),
2283 "a tombstone on the derived key shadows inner"
2284 );
2285 }
2286
2287 #[test]
2288 fn contains_does_not_fall_back_for_prefixed_keys() {
2289 let state = State::new();
2290 state.set("derived:flag", true).unwrap();
2291 assert!(
2292 !state.contains("session:flag"),
2293 "only unprefixed keys get the fallback"
2294 );
2295 }
2296}