1use zeph_db::ActiveDialect;
18#[allow(unused_imports)]
19use zeph_db::sql;
20
21use super::SqliteStore;
22use crate::error::MemoryError;
23
24#[derive(Debug, Clone, PartialEq, Eq)]
30pub struct StoreItem {
31 pub owner_key: String,
32 pub namespace: String,
33 pub key: String,
34 pub value: String,
35 pub version: i64,
36 pub created_at: String,
37 pub updated_at: String,
38}
39
40type StoreItemTuple = (String, String, String, String, i64, String, String);
41
42fn item_from_tuple(t: StoreItemTuple) -> StoreItem {
43 StoreItem {
44 owner_key: t.0,
45 namespace: t.1,
46 key: t.2,
47 value: t.3,
48 version: t.4,
49 created_at: t.5,
50 updated_at: t.6,
51 }
52}
53
54fn select_columns() -> String {
58 let created_at_sel = <ActiveDialect as zeph_db::dialect::Dialect>::select_as_text("created_at");
59 let updated_at_sel = <ActiveDialect as zeph_db::dialect::Dialect>::select_as_text("updated_at");
60 format!("owner_key, namespace, key, value, version, {created_at_sel}, {updated_at_sel}")
61}
62
63impl SqliteStore {
64 pub async fn store_put(
101 &self,
102 owner_key: &str,
103 namespace: &str,
104 key: &str,
105 value: &str,
106 max_value_bytes: usize,
107 expected_version: Option<i64>,
108 ) -> Result<StoreItem, MemoryError> {
109 if value.len() > max_value_bytes {
110 return Err(MemoryError::InvalidInput(format!(
111 "store_put value for namespace={namespace:?} key={key:?} is {} bytes, \
112 exceeds max_value_bytes={max_value_bytes}",
113 value.len()
114 )));
115 }
116
117 let now = <ActiveDialect as zeph_db::dialect::Dialect>::NOW;
118 let cols = select_columns();
119
120 let row: Option<StoreItemTuple> = if let Some(expected) = expected_version {
121 let raw = format!(
122 "UPDATE cross_thread_store \
123 SET value = ?, version = version + 1, updated_at = {now} \
124 WHERE owner_key = ? AND namespace = ? AND key = ? AND version = ? \
125 RETURNING {cols}"
126 );
127 let query_sql = zeph_db::rewrite_placeholders(&raw);
128 zeph_db::query_as(sqlx::AssertSqlSafe(query_sql))
129 .bind(value)
130 .bind(owner_key)
131 .bind(namespace)
132 .bind(key)
133 .bind(expected)
134 .fetch_optional(&self.pool)
135 .await?
136 } else {
137 let raw = format!(
138 "INSERT INTO cross_thread_store (owner_key, namespace, key, value) \
139 VALUES (?, ?, ?, ?) \
140 ON CONFLICT(owner_key, namespace, key) DO UPDATE SET \
141 value = excluded.value, \
142 version = cross_thread_store.version + 1, \
143 updated_at = {now} \
144 RETURNING {cols}"
145 );
146 let query_sql = zeph_db::rewrite_placeholders(&raw);
147 zeph_db::query_as(sqlx::AssertSqlSafe(query_sql))
148 .bind(owner_key)
149 .bind(namespace)
150 .bind(key)
151 .bind(value)
152 .fetch_optional(&self.pool)
153 .await?
154 };
155
156 match row {
157 Some(t) => Ok(item_from_tuple(t)),
158 None => Err(MemoryError::VersionConflict {
159 owner_key: owner_key.to_owned(),
160 namespace: namespace.to_owned(),
161 key: key.to_owned(),
162 expected: expected_version.unwrap_or(0),
163 }),
164 }
165 }
166
167 pub async fn store_get(
187 &self,
188 owner_key: &str,
189 namespace: &str,
190 key: &str,
191 ) -> Result<Option<StoreItem>, MemoryError> {
192 let cols = select_columns();
193 let raw = format!(
194 "SELECT {cols} FROM cross_thread_store \
195 WHERE owner_key = ? AND namespace = ? AND key = ?"
196 );
197 let query_sql = zeph_db::rewrite_placeholders(&raw);
198 let row: Option<StoreItemTuple> = zeph_db::query_as(sqlx::AssertSqlSafe(query_sql))
199 .bind(owner_key)
200 .bind(namespace)
201 .bind(key)
202 .fetch_optional(&self.pool)
203 .await?;
204 Ok(row.map(item_from_tuple))
205 }
206
207 pub async fn store_delete(
233 &self,
234 owner_key: &str,
235 namespace: &str,
236 key: &str,
237 ) -> Result<bool, MemoryError> {
238 let result = zeph_db::query(sql!(
239 "DELETE FROM cross_thread_store WHERE owner_key = ? AND namespace = ? AND key = ?"
240 ))
241 .bind(owner_key)
242 .bind(namespace)
243 .bind(key)
244 .execute(&self.pool)
245 .await?;
246 Ok(result.rows_affected() > 0)
247 }
248
249 pub async fn store_list(
288 &self,
289 owner_key: &str,
290 namespace_prefix: &str,
291 limit: usize,
292 ) -> Result<Vec<StoreItem>, MemoryError> {
293 let cols = select_columns();
294 let (limit_clause, limit_bind) = zeph_db::limit_clause(limit as u64);
295 let upper = prefix_range_upper_bound(namespace_prefix);
296 let raw = if upper.is_some() {
297 format!(
298 "SELECT {cols} FROM cross_thread_store \
299 WHERE owner_key = ? AND namespace >= ? AND namespace < ? \
300 ORDER BY namespace, key{limit_clause}"
301 )
302 } else {
303 format!(
304 "SELECT {cols} FROM cross_thread_store \
305 WHERE owner_key = ? AND namespace >= ? \
306 ORDER BY namespace, key{limit_clause}"
307 )
308 };
309 let query_sql = zeph_db::rewrite_placeholders(&raw);
310 let mut query = zeph_db::query_as(sqlx::AssertSqlSafe(query_sql))
311 .bind(owner_key)
312 .bind(namespace_prefix.to_owned());
313 if let Some(ref upper) = upper {
314 query = query.bind(upper.clone());
315 }
316 if let Some(lim) = limit_bind {
317 query = query.bind(lim);
318 }
319 let rows: Vec<StoreItemTuple> = query.fetch_all(&self.pool).await?;
320 Ok(rows.into_iter().map(item_from_tuple).collect())
321 }
322
323 pub async fn store_search(
362 &self,
363 owner_key: &str,
364 namespace_prefix: &str,
365 query: &str,
366 limit: usize,
367 ) -> Result<Vec<StoreItem>, MemoryError> {
368 let cols = select_columns();
369 let (limit_clause, limit_bind) = zeph_db::limit_clause(limit as u64);
370 let upper = prefix_range_upper_bound(namespace_prefix);
371 let raw = if upper.is_some() {
372 format!(
373 "SELECT {cols} FROM cross_thread_store \
374 WHERE owner_key = ? AND namespace >= ? AND namespace < ? \
375 AND value LIKE ? ESCAPE '\\' \
376 ORDER BY namespace, key{limit_clause}"
377 )
378 } else {
379 format!(
380 "SELECT {cols} FROM cross_thread_store \
381 WHERE owner_key = ? AND namespace >= ? AND value LIKE ? ESCAPE '\\' \
382 ORDER BY namespace, key{limit_clause}"
383 )
384 };
385 let query_sql = zeph_db::rewrite_placeholders(&raw);
386 let mut q = zeph_db::query_as(sqlx::AssertSqlSafe(query_sql))
387 .bind(owner_key)
388 .bind(namespace_prefix.to_owned());
389 if let Some(ref upper) = upper {
390 q = q.bind(upper.clone());
391 }
392 q = q.bind(like_contains(query));
393 if let Some(lim) = limit_bind {
394 q = q.bind(lim);
395 }
396 let rows: Vec<StoreItemTuple> = q.fetch_all(&self.pool).await?;
397 Ok(rows.into_iter().map(item_from_tuple).collect())
398 }
399}
400
401fn prefix_range_upper_bound(prefix: &str) -> Option<String> {
416 let mut chars: Vec<char> = prefix.chars().collect();
417 while let Some(last) = chars.pop() {
418 let mut next = last as u32 + 1;
419 if (0xD800..=0xDFFF).contains(&next) {
420 next = 0xE000; }
422 if let Some(incremented) = char::from_u32(next) {
423 chars.push(incremented);
424 return Some(chars.into_iter().collect());
425 }
426 }
428 None
429}
430
431fn like_contains(substring: &str) -> String {
433 format!("%{}%", escape_like(substring))
434}
435
436fn escape_like(s: &str) -> String {
437 s.replace('\\', "\\\\")
438 .replace('%', "\\%")
439 .replace('_', "\\_")
440}
441
442#[cfg(test)]
443mod tests {
444 use super::*;
445
446 async fn store() -> SqliteStore {
447 SqliteStore::new(":memory:").await.unwrap()
448 }
449
450 const MAX_BYTES: usize = 65536;
451
452 #[test]
455 fn prefix_range_upper_bound_bumps_last_char() {
456 assert_eq!(
457 prefix_range_upper_bound("orch/g1").as_deref(),
458 Some("orch/g2")
459 );
460 }
461
462 #[test]
463 fn prefix_range_upper_bound_bumps_slash_to_digit_zero() {
464 assert_eq!(prefix_range_upper_bound("orch/").as_deref(), Some("orch0"));
467 }
468
469 #[test]
470 fn prefix_range_upper_bound_empty_prefix_returns_none() {
471 assert_eq!(prefix_range_upper_bound(""), None);
472 }
473
474 #[test]
475 fn prefix_range_upper_bound_brackets_every_continuation_and_nothing_else() {
476 let prefix = "ns";
477 let upper = prefix_range_upper_bound(prefix).unwrap();
478 assert!(
479 prefix < upper.as_str(),
480 "prefix itself must fall in [prefix, upper)"
481 );
482 assert!(
483 format!("{prefix}-anything") < upper,
484 "any continuation of prefix must sort before upper bound"
485 );
486 assert!(
487 "nt" >= upper.as_str(),
488 "an unrelated namespace one step past the prefix family must not be < upper"
489 );
490 }
491
492 #[tokio::test]
493 async fn put_get_roundtrip() {
494 let s = store().await;
495 let item = s
496 .store_put("local", "orch/g1", "finding", "{\"x\":1}", MAX_BYTES, None)
497 .await
498 .unwrap();
499 assert_eq!(item.version, 1);
500 assert_eq!(item.value, "{\"x\":1}");
501
502 let fetched = s
503 .store_get("local", "orch/g1", "finding")
504 .await
505 .unwrap()
506 .expect("row must exist");
507 assert_eq!(fetched.value, "{\"x\":1}");
508 assert_eq!(fetched.version, 1);
509 assert_eq!(fetched.owner_key, "local");
510 assert_eq!(fetched.namespace, "orch/g1");
511 assert_eq!(fetched.key, "finding");
512 }
513
514 #[tokio::test]
515 async fn put_upserts_and_bumps_version() {
516 let s = store().await;
517 s.store_put("local", "ns", "k", "v1", MAX_BYTES, None)
518 .await
519 .unwrap();
520 let updated = s
521 .store_put("local", "ns", "k", "v2", MAX_BYTES, None)
522 .await
523 .unwrap();
524 assert_eq!(updated.version, 2);
525 assert_eq!(updated.value, "v2");
526
527 let fetched = s.store_get("local", "ns", "k").await.unwrap().unwrap();
528 assert_eq!(fetched.value, "v2");
529 assert_eq!(fetched.version, 2);
530 }
531
532 #[tokio::test]
534 async fn namespace_isolation() {
535 let s = store().await;
536 s.store_put("local", "ns-a", "k", "value-a", MAX_BYTES, None)
537 .await
538 .unwrap();
539 s.store_put("local", "ns-b", "k", "value-b", MAX_BYTES, None)
540 .await
541 .unwrap();
542
543 let a = s.store_get("local", "ns-a", "k").await.unwrap().unwrap();
544 let b = s.store_get("local", "ns-b", "k").await.unwrap().unwrap();
545 assert_eq!(a.value, "value-a");
546 assert_eq!(b.value, "value-b");
547 }
548
549 #[tokio::test]
552 async fn owner_key_isolation() {
553 let s = store().await;
554 s.store_put("owner-a", "ns", "k", "value-a", MAX_BYTES, None)
555 .await
556 .unwrap();
557 s.store_put("owner-b", "ns", "k", "value-b", MAX_BYTES, None)
558 .await
559 .unwrap();
560
561 let a = s.store_get("owner-a", "ns", "k").await.unwrap().unwrap();
562 let b = s.store_get("owner-b", "ns", "k").await.unwrap().unwrap();
563 assert_eq!(a.value, "value-a");
564 assert_eq!(b.value, "value-b");
565
566 assert!(s.store_delete("owner-a", "ns", "k").await.unwrap());
567 assert!(s.store_get("owner-b", "ns", "k").await.unwrap().is_some());
569 }
570
571 #[tokio::test]
572 async fn version_conflict_on_stale_expected_version() {
573 let s = store().await;
574 let first = s
575 .store_put("local", "ns", "k", "v1", MAX_BYTES, None)
576 .await
577 .unwrap();
578 assert_eq!(first.version, 1);
579
580 let second = s
582 .store_put("local", "ns", "k", "v2", MAX_BYTES, Some(1))
583 .await
584 .unwrap();
585 assert_eq!(second.version, 2);
586
587 let err = s
589 .store_put("local", "ns", "k", "v3", MAX_BYTES, Some(1))
590 .await
591 .unwrap_err();
592 assert!(matches!(err, MemoryError::VersionConflict { .. }));
593
594 let fetched = s.store_get("local", "ns", "k").await.unwrap().unwrap();
596 assert_eq!(fetched.value, "v2");
597 assert_eq!(fetched.version, 2);
598 }
599
600 #[tokio::test]
601 async fn version_conflict_when_row_does_not_exist() {
602 let s = store().await;
603 let err = s
604 .store_put("local", "ns", "no-such-key", "v", MAX_BYTES, Some(1))
605 .await
606 .unwrap_err();
607 assert!(matches!(err, MemoryError::VersionConflict { .. }));
608 }
609
610 #[tokio::test]
611 async fn put_rejects_value_exceeding_max_bytes() {
612 let s = store().await;
613 let err = s
614 .store_put("local", "ns", "k", "0123456789", 5, None)
615 .await
616 .unwrap_err();
617 assert!(matches!(err, MemoryError::InvalidInput(_)));
618
619 assert!(s.store_get("local", "ns", "k").await.unwrap().is_none());
621 }
622
623 #[tokio::test]
624 async fn delete_returns_false_for_missing_row() {
625 let s = store().await;
626 assert!(!s.store_delete("local", "ns", "no-such").await.unwrap());
627 }
628
629 #[tokio::test]
630 async fn get_returns_none_for_missing_row() {
631 let s = store().await;
632 assert!(
633 s.store_get("local", "ns", "no-such")
634 .await
635 .unwrap()
636 .is_none()
637 );
638 }
639
640 #[tokio::test]
641 async fn list_by_namespace_prefix() {
642 let s = store().await;
643 s.store_put("local", "orch/g1", "a", "1", MAX_BYTES, None)
644 .await
645 .unwrap();
646 s.store_put("local", "orch/g1", "b", "2", MAX_BYTES, None)
647 .await
648 .unwrap();
649 s.store_put("local", "orch/g2", "c", "3", MAX_BYTES, None)
650 .await
651 .unwrap();
652
653 let g1 = s.store_list("local", "orch/g1", 0).await.unwrap();
654 assert_eq!(g1.len(), 2);
655 assert!(g1.iter().all(|i| i.namespace == "orch/g1"));
656
657 let all_orch = s.store_list("local", "orch/", 0).await.unwrap();
658 assert_eq!(all_orch.len(), 3);
659 }
660
661 #[tokio::test]
662 async fn list_respects_limit() {
663 let s = store().await;
664 for i in 0..5u8 {
665 s.store_put("local", "ns", &format!("k{i}"), "v", MAX_BYTES, None)
666 .await
667 .unwrap();
668 }
669 let limited = s.store_list("local", "ns", 2).await.unwrap();
670 assert_eq!(limited.len(), 2);
671 }
672
673 #[tokio::test]
674 async fn list_empty_namespace_returns_empty_vec() {
675 let s = store().await;
676 let rows = s.store_list("local", "no/such/ns", 0).await.unwrap();
677 assert!(rows.is_empty());
678 }
679
680 #[tokio::test]
681 async fn search_matches_value_keyword() {
682 let s = store().await;
683 s.store_put(
684 "local",
685 "orch/g1",
686 "a",
687 "{\"finding\":\"needle in haystack\"}",
688 MAX_BYTES,
689 None,
690 )
691 .await
692 .unwrap();
693 s.store_put(
694 "local",
695 "orch/g1",
696 "b",
697 "{\"finding\":\"nothing here\"}",
698 MAX_BYTES,
699 None,
700 )
701 .await
702 .unwrap();
703
704 let hits = s
705 .store_search("local", "orch/g1", "needle", 0)
706 .await
707 .unwrap();
708 assert_eq!(hits.len(), 1);
709 assert_eq!(hits[0].key, "a");
710 }
711
712 #[tokio::test]
713 async fn search_scoped_by_namespace_prefix() {
714 let s = store().await;
715 s.store_put("local", "orch/g1", "a", "needle", MAX_BYTES, None)
716 .await
717 .unwrap();
718 s.store_put("local", "orch/g2", "b", "needle", MAX_BYTES, None)
719 .await
720 .unwrap();
721
722 let hits = s
723 .store_search("local", "orch/g1", "needle", 0)
724 .await
725 .unwrap();
726 assert_eq!(hits.len(), 1);
727 assert_eq!(hits[0].namespace, "orch/g1");
728 }
729
730 #[tokio::test]
731 async fn like_wildcards_in_query_are_escaped() {
732 let s = store().await;
733 s.store_put("local", "ns", "a", "50% off", MAX_BYTES, None)
734 .await
735 .unwrap();
736 s.store_put("local", "ns", "b", "50x off", MAX_BYTES, None)
737 .await
738 .unwrap();
739
740 let hits = s.store_search("local", "ns", "50%", 0).await.unwrap();
742 assert_eq!(hits.len(), 1);
743 assert_eq!(hits[0].key, "a");
744 }
745}