1mod incremental;
2
3use super::SharedConnection;
4use std::collections::HashMap;
5use std::fs;
6use std::io::Write;
7use std::path::PathBuf;
8use std::sync::{Arc, Mutex};
9use std::time::{Duration, Instant};
10
11#[derive(Debug, Clone, PartialEq, Eq)]
12pub struct StorageError {
13 message: String,
14}
15
16impl StorageError {
17 pub fn new(message: impl Into<String>) -> Self {
18 Self {
19 message: message.into(),
20 }
21 }
22}
23
24impl std::fmt::Display for StorageError {
25 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
26 self.message.fmt(f)
27 }
28}
29
30impl std::error::Error for StorageError {}
31
32pub type StorageResult<T> = Result<T, StorageError>;
33
34pub trait StorageAdapter: Send + Sync {
35 fn get(&self, key: &str) -> StorageResult<Option<String>>;
36 fn put(&self, key: &str, value: String) -> StorageResult<()>;
37 fn del(&self, key: &str) -> StorageResult<()>;
38 fn list(&self, prefix: &str) -> StorageResult<Vec<String>>;
39}
40
41#[derive(Clone)]
42pub struct InMemoryStorage {
43 store: Arc<Mutex<HashMap<String, String>>>,
44}
45
46impl InMemoryStorage {
47 pub fn new() -> Self {
48 Self {
49 store: Arc::new(Mutex::new(HashMap::new())),
50 }
51 }
52}
53
54impl Default for InMemoryStorage {
55 fn default() -> Self {
56 Self::new()
57 }
58}
59
60impl StorageAdapter for InMemoryStorage {
61 fn get(&self, key: &str) -> StorageResult<Option<String>> {
62 let store = self
63 .store
64 .lock()
65 .map_err(|_| StorageError::new("storage mutex poisoned"))?;
66 Ok(store.get(key).cloned())
67 }
68
69 fn put(&self, key: &str, value: String) -> StorageResult<()> {
70 let mut store = self
71 .store
72 .lock()
73 .map_err(|_| StorageError::new("storage mutex poisoned"))?;
74 store.insert(key.to_string(), value);
75 Ok(())
76 }
77
78 fn del(&self, key: &str) -> StorageResult<()> {
79 let mut store = self
80 .store
81 .lock()
82 .map_err(|_| StorageError::new("storage mutex poisoned"))?;
83 store.remove(key);
84 Ok(())
85 }
86
87 fn list(&self, prefix: &str) -> StorageResult<Vec<String>> {
88 let store = self
89 .store
90 .lock()
91 .map_err(|_| StorageError::new("storage mutex poisoned"))?;
92 Ok(store
93 .keys()
94 .filter(|key| key.starts_with(prefix))
95 .cloned()
96 .collect())
97 }
98}
99
100pub struct FileStorageAdapter {
101 base_path: PathBuf,
102}
103
104impl FileStorageAdapter {
105 pub fn new(base_path: PathBuf) -> StorageResult<Self> {
108 let mut builder = fs::DirBuilder::new();
109 builder.recursive(true);
110 #[cfg(unix)]
111 {
112 use std::os::unix::fs::DirBuilderExt;
113 builder.mode(0o700);
114 }
115 builder
116 .create(&base_path)
117 .map_err(|err| storage_io_error("failed to create storage directory", err))?;
118 #[cfg(unix)]
121 {
122 use std::os::unix::fs::PermissionsExt;
123 fs::set_permissions(&base_path, fs::Permissions::from_mode(0o700))
124 .map_err(|err| storage_io_error("failed to protect storage directory", err))?;
125 }
126 Ok(Self { base_path })
127 }
128
129 fn sanitize_key(key: &str) -> String {
130 key.replace(['/', '\\', ':'], "_")
131 }
132
133 fn key_to_path(&self, key: &str) -> PathBuf {
134 let sanitized = Self::sanitize_key(key);
135 self.base_path.join(format!("{}.json", sanitized))
136 }
137}
138
139impl StorageAdapter for FileStorageAdapter {
140 fn get(&self, key: &str) -> StorageResult<Option<String>> {
141 let path = self.key_to_path(key);
142 match fs::read_to_string(&path) {
143 Ok(contents) => Ok(Some(contents)),
144 Err(err) if err.kind() == std::io::ErrorKind::NotFound => Ok(None),
145 Err(err) => Err(storage_io_error("failed to read storage file", err)),
146 }
147 }
148
149 fn put(&self, key: &str, value: String) -> StorageResult<()> {
150 let path = self.key_to_path(key);
151
152 let tmp_path = path.with_extension(format!("json.{}.tmp", rand::random::<u128>()));
153 let mut options = fs::OpenOptions::new();
154 options.write(true).create_new(true);
155 #[cfg(unix)]
156 {
157 use std::os::unix::fs::OpenOptionsExt;
158 options.mode(0o600);
159 }
160 let mut file = options
163 .open(&tmp_path)
164 .map_err(|err| storage_io_error("failed to create storage temp file", err))?;
165 if let Err(err) = file.write_all(value.as_bytes()) {
166 drop(file);
167 let _ = fs::remove_file(&tmp_path);
168 return Err(storage_io_error("failed to write storage temp file", err));
169 }
170 drop(file);
171
172 #[cfg(windows)]
173 {
174 if path.exists() {
175 fs::remove_file(&path).map_err(|err| {
176 storage_io_error("failed to replace existing storage file", err)
177 })?;
178 }
179 }
180
181 fs::rename(&tmp_path, &path)
182 .map_err(|err| storage_io_error("failed to commit storage file", err))?;
183
184 Ok(())
185 }
186
187 fn del(&self, key: &str) -> StorageResult<()> {
188 let path = self.key_to_path(key);
189 match fs::remove_file(&path) {
190 Ok(()) => Ok(()),
191 Err(err) if err.kind() == std::io::ErrorKind::NotFound => Ok(()),
192 Err(err) => Err(storage_io_error("failed to delete storage file", err)),
193 }
194 }
195
196 fn list(&self, prefix: &str) -> StorageResult<Vec<String>> {
197 let mut keys = Vec::new();
198 let sanitized_prefix = Self::sanitize_key(prefix);
199 let entries = fs::read_dir(&self.base_path)
200 .map_err(|err| storage_io_error("failed to read storage directory", err))?;
201
202 for entry in entries {
203 let entry =
204 entry.map_err(|err| storage_io_error("failed to read storage entry", err))?;
205 let file_name = entry.file_name();
206 let file_name_str = file_name.to_string_lossy();
207
208 if !file_name_str.ends_with(".json") {
209 continue;
210 }
211
212 let key = file_name_str
213 .strip_suffix(".json")
214 .unwrap_or(&file_name_str)
215 .to_string();
216
217 if prefix.is_empty() {
218 keys.push(key);
219 continue;
220 }
221
222 if key.starts_with(&sanitized_prefix) {
223 let remainder = key.strip_prefix(&sanitized_prefix).unwrap_or("");
224 keys.push(format!("{}{}", prefix, remainder));
225 }
226 }
227
228 Ok(keys)
229 }
230}
231
232pub struct DebouncedFileStorage {
233 adapter: FileStorageAdapter,
234 pending_writes: Mutex<HashMap<String, String>>,
235 last_flush: Mutex<Instant>,
236 flush_interval: Duration,
237}
238
239impl DebouncedFileStorage {
240 pub fn new(base_path: PathBuf, flush_interval_ms: u64) -> StorageResult<Self> {
241 Ok(Self {
242 adapter: FileStorageAdapter::new(base_path)?,
243 pending_writes: Mutex::new(HashMap::new()),
244 last_flush: Mutex::new(Instant::now()),
245 flush_interval: Duration::from_millis(flush_interval_ms),
246 })
247 }
248
249 pub fn flush(&self) -> StorageResult<()> {
250 let mut pending = self
251 .pending_writes
252 .lock()
253 .map_err(|_| StorageError::new("pending file storage mutex poisoned"))?;
254 for (key, value) in pending.drain() {
255 self.adapter.put(&key, value)?;
256 }
257 *self
258 .last_flush
259 .lock()
260 .map_err(|_| StorageError::new("file storage flush mutex poisoned"))? = Instant::now();
261 Ok(())
262 }
263
264 fn maybe_flush(&self) -> StorageResult<()> {
265 let last_flush = *self
266 .last_flush
267 .lock()
268 .map_err(|_| StorageError::new("file storage flush mutex poisoned"))?;
269 let pending_count = self
270 .pending_writes
271 .lock()
272 .map_err(|_| StorageError::new("pending file storage mutex poisoned"))?
273 .len();
274
275 if last_flush.elapsed() >= self.flush_interval && pending_count > 0 {
276 self.flush()?;
277 }
278 Ok(())
279 }
280}
281
282impl StorageAdapter for DebouncedFileStorage {
283 fn get(&self, key: &str) -> StorageResult<Option<String>> {
284 let pending = self
285 .pending_writes
286 .lock()
287 .map_err(|_| StorageError::new("pending file storage mutex poisoned"))?;
288 if let Some(value) = pending.get(key) {
289 return Ok(Some(value.clone()));
290 }
291 drop(pending);
292 self.adapter.get(key)
293 }
294
295 fn put(&self, key: &str, value: String) -> StorageResult<()> {
296 self.pending_writes
297 .lock()
298 .map_err(|_| StorageError::new("pending file storage mutex poisoned"))?
299 .insert(key.to_string(), value);
300 self.maybe_flush()
301 }
302
303 fn del(&self, key: &str) -> StorageResult<()> {
304 self.pending_writes
305 .lock()
306 .map_err(|_| StorageError::new("pending file storage mutex poisoned"))?
307 .remove(key);
308 self.adapter.del(key)
309 }
310
311 fn list(&self, prefix: &str) -> StorageResult<Vec<String>> {
312 let mut keys = self.adapter.list(prefix)?;
313 let pending = self
314 .pending_writes
315 .lock()
316 .map_err(|_| StorageError::new("pending file storage mutex poisoned"))?;
317
318 for key in pending.keys() {
319 if key.starts_with(prefix) && !keys.contains(key) {
320 keys.push(key.clone());
321 }
322 }
323
324 Ok(keys)
325 }
326}
327
328pub struct SqliteStorageAdapter {
334 conn: SharedConnection,
335 owner_pubkey_hex: String,
336 device_pubkey_hex: String,
337}
338
339impl SqliteStorageAdapter {
340 pub fn new(
341 conn: SharedConnection,
342 owner_pubkey_hex: String,
343 device_pubkey_hex: String,
344 ) -> Self {
345 Self {
346 conn,
347 owner_pubkey_hex,
348 device_pubkey_hex,
349 }
350 }
351
352 fn map_err<E: std::fmt::Display>(error: E) -> StorageError {
353 StorageError::new(error.to_string())
354 }
355}
356
357impl StorageAdapter for SqliteStorageAdapter {
358 fn get(&self, key: &str) -> StorageResult<Option<String>> {
359 let conn = self
360 .conn
361 .lock()
362 .map_err(|_| StorageError::new("ndr_kv connection mutex poisoned"))?;
363 conn.query_row(
364 "SELECT value FROM ndr_kv WHERE owner_pubkey_hex = ?1 AND device_pubkey_hex = ?2 AND key = ?3",
365 (&self.owner_pubkey_hex, &self.device_pubkey_hex, key),
366 |row| row.get::<_, String>(0),
367 )
368 .map(Some)
369 .or_else(|err| match err {
370 rusqlite::Error::QueryReturnedNoRows => Ok(None),
371 other => Err(Self::map_err(other)),
372 })
373 }
374
375 fn put(&self, key: &str, value: String) -> StorageResult<()> {
376 let mut conn = self
377 .conn
378 .lock()
379 .map_err(|_| StorageError::new("ndr_kv connection mutex poisoned"))?;
380 if incremental::try_put(
381 &mut conn,
382 &self.owner_pubkey_hex,
383 &self.device_pubkey_hex,
384 key,
385 &value,
386 )? {
387 return Ok(());
388 }
389 conn.execute(
390 "INSERT INTO ndr_kv (owner_pubkey_hex, device_pubkey_hex, key, value)
391 VALUES (?1, ?2, ?3, ?4)
392 ON CONFLICT(owner_pubkey_hex, device_pubkey_hex, key) DO UPDATE SET value = excluded.value",
393 (&self.owner_pubkey_hex, &self.device_pubkey_hex, key, &value),
394 )
395 .map_err(Self::map_err)?;
396 Ok(())
397 }
398
399 fn del(&self, key: &str) -> StorageResult<()> {
400 let conn = self
401 .conn
402 .lock()
403 .map_err(|_| StorageError::new("ndr_kv connection mutex poisoned"))?;
404 conn.execute(
405 "DELETE FROM ndr_kv WHERE owner_pubkey_hex = ?1 AND device_pubkey_hex = ?2 AND key = ?3",
406 (&self.owner_pubkey_hex, &self.device_pubkey_hex, key),
407 )
408 .map_err(Self::map_err)?;
409 Ok(())
410 }
411
412 fn list(&self, prefix: &str) -> StorageResult<Vec<String>> {
413 let conn = self
414 .conn
415 .lock()
416 .map_err(|_| StorageError::new("ndr_kv connection mutex poisoned"))?;
417 let mut stmt = conn
418 .prepare(
419 "SELECT key FROM ndr_kv
420 WHERE owner_pubkey_hex = ?1 AND device_pubkey_hex = ?2 AND key LIKE ?3 ESCAPE '\\'",
421 )
422 .map_err(Self::map_err)?;
423 let pattern = format!("{}%", escape_like(prefix));
424 let rows = stmt
425 .query_map(
426 (&self.owner_pubkey_hex, &self.device_pubkey_hex, &pattern),
427 |row| row.get::<_, String>(0),
428 )
429 .map_err(Self::map_err)?;
430 let mut keys = Vec::new();
431 for row in rows {
432 keys.push(row.map_err(Self::map_err)?);
433 }
434 Ok(keys)
435 }
436}
437
438fn escape_like(input: &str) -> String {
439 let mut out = String::with_capacity(input.len());
440 for ch in input.chars() {
441 match ch {
442 '\\' | '%' | '_' => {
443 out.push('\\');
444 out.push(ch);
445 }
446 other => out.push(other),
447 }
448 }
449 out
450}
451
452fn storage_io_error(context: &str, error: std::io::Error) -> StorageError {
453 StorageError::new(format!("{}: {}", context, error))
454}
455
456#[cfg(test)]
457mod tests {
458 use super::*;
459 use std::sync::{Arc, Mutex};
460 use tempfile::TempDir;
461
462 fn fresh_connection() -> SharedConnection {
463 let conn = rusqlite::Connection::open_in_memory().unwrap();
464 conn.execute_batch(
465 "CREATE TABLE ndr_kv (
466 owner_pubkey_hex TEXT NOT NULL,
467 device_pubkey_hex TEXT NOT NULL,
468 key TEXT NOT NULL,
469 value TEXT NOT NULL,
470 PRIMARY KEY (owner_pubkey_hex, device_pubkey_hex, key)
471 );",
472 )
473 .unwrap();
474 Arc::new(Mutex::new(conn))
475 }
476
477 fn fresh_adapter() -> SqliteStorageAdapter {
478 SqliteStorageAdapter::new(
479 fresh_connection(),
480 "owner".to_string(),
481 "device".to_string(),
482 )
483 }
484
485 #[test]
486 fn put_get_del_round_trip() {
487 let adapter = fresh_adapter();
488 assert!(adapter.get("k").unwrap().is_none());
489 adapter.put("k", "v".to_string()).unwrap();
490 assert_eq!(adapter.get("k").unwrap(), Some("v".to_string()));
491 adapter.put("k", "v2".to_string()).unwrap();
492 assert_eq!(adapter.get("k").unwrap(), Some("v2".to_string()));
493 adapter.del("k").unwrap();
494 assert!(adapter.get("k").unwrap().is_none());
495 }
496
497 #[test]
498 fn list_returns_only_matching_prefix() {
499 let adapter = fresh_adapter();
500 adapter.put("user/alice", "1".to_string()).unwrap();
501 adapter.put("user/bob", "2".to_string()).unwrap();
502 adapter.put("invite/charlie", "3".to_string()).unwrap();
503 let mut keys = adapter.list("user/").unwrap();
504 keys.sort();
505 assert_eq!(keys, vec!["user/alice".to_string(), "user/bob".to_string()]);
506 }
507
508 #[test]
509 fn keys_are_isolated_per_owner_device() {
510 let conn = fresh_connection();
511 let alice = SqliteStorageAdapter::new(conn.clone(), "owner_a".into(), "device_a".into());
512 let bob = SqliteStorageAdapter::new(conn, "owner_b".into(), "device_b".into());
513 alice.put("shared-key", "alice".to_string()).unwrap();
514 bob.put("shared-key", "bob".to_string()).unwrap();
515 assert_eq!(alice.get("shared-key").unwrap(), Some("alice".to_string()));
516 assert_eq!(bob.get("shared-key").unwrap(), Some("bob".to_string()));
517 }
518
519 #[test]
520 fn file_storage_round_trips_values() {
521 let temp_dir = TempDir::new().unwrap();
522 let adapter = FileStorageAdapter::new(temp_dir.path().to_path_buf()).unwrap();
523
524 assert!(adapter.get("test-key").unwrap().is_none());
525
526 adapter.put("test-key", "test-value".to_string()).unwrap();
527 assert_eq!(
528 adapter.get("test-key").unwrap(),
529 Some("test-value".to_string())
530 );
531
532 adapter.del("test-key").unwrap();
533 assert!(adapter.get("test-key").unwrap().is_none());
534 }
535
536 #[cfg(unix)]
537 #[test]
538 fn file_storage_keeps_session_secrets_private_when_created_and_replaced() {
539 use std::os::unix::fs::PermissionsExt;
540
541 let temp_dir = TempDir::new().unwrap();
542 let storage_dir = temp_dir.path().join("protocol");
543 let adapter = FileStorageAdapter::new(storage_dir.clone()).unwrap();
544 assert_eq!(
545 fs::metadata(&storage_dir).unwrap().permissions().mode() & 0o777,
546 0o700
547 );
548
549 adapter.put("session", "secret".to_string()).unwrap();
550 let session_path = storage_dir.join("session.json");
551 assert_eq!(
552 fs::metadata(&session_path).unwrap().permissions().mode() & 0o777,
553 0o600
554 );
555
556 fs::set_permissions(&session_path, fs::Permissions::from_mode(0o644)).unwrap();
558 adapter.put("session", "new secret".to_string()).unwrap();
559 assert_eq!(
560 fs::metadata(&session_path).unwrap().permissions().mode() & 0o777,
561 0o600
562 );
563 assert_eq!(
564 adapter.get("session").unwrap().as_deref(),
565 Some("new secret")
566 );
567 }
568
569 #[cfg(unix)]
570 #[test]
571 fn file_storage_protects_existing_session_directory() {
572 use std::os::unix::fs::PermissionsExt;
573
574 let temp_dir = TempDir::new().unwrap();
575 let storage_dir = temp_dir.path().join("protocol");
576 fs::create_dir(&storage_dir).unwrap();
577 fs::set_permissions(&storage_dir, fs::Permissions::from_mode(0o755)).unwrap();
578 fs::write(storage_dir.join("session.json"), "legacy secret").unwrap();
579
580 let adapter = FileStorageAdapter::new(storage_dir.clone()).unwrap();
581
582 assert_eq!(
583 fs::metadata(&storage_dir).unwrap().permissions().mode() & 0o777,
584 0o700
585 );
586 assert_eq!(
587 adapter.get("session").unwrap().as_deref(),
588 Some("legacy secret")
589 );
590 }
591
592 #[test]
593 fn file_storage_lists_sanitized_runtime_keys() {
594 let temp_dir = TempDir::new().unwrap();
595 let adapter = FileStorageAdapter::new(temp_dir.path().to_path_buf()).unwrap();
596
597 adapter.put("user/alice", "1".to_string()).unwrap();
598 adapter.put("user/bob", "2".to_string()).unwrap();
599 adapter.put("invite/charlie", "3".to_string()).unwrap();
600
601 let mut user_keys = adapter.list("user/").unwrap();
602 user_keys.sort();
603 assert_eq!(
604 user_keys,
605 vec!["user/alice".to_string(), "user/bob".to_string()]
606 );
607
608 let mut all_keys = adapter.list("").unwrap();
609 all_keys.sort();
610 assert_eq!(
611 all_keys,
612 vec![
613 "invite_charlie".to_string(),
614 "user_alice".to_string(),
615 "user_bob".to_string()
616 ]
617 );
618 }
619
620 #[test]
621 fn debounced_file_storage_reads_pending_writes_and_flushes() {
622 let temp_dir = TempDir::new().unwrap();
623 let storage = DebouncedFileStorage::new(temp_dir.path().to_path_buf(), 1000).unwrap();
624
625 storage.put("key1", "value1".to_string()).unwrap();
626
627 assert_eq!(storage.get("key1").unwrap(), Some("value1".to_string()));
628 assert!(storage.pending_writes.lock().unwrap().contains_key("key1"));
629
630 storage.flush().unwrap();
631
632 assert!(storage.pending_writes.lock().unwrap().is_empty());
633 assert_eq!(
634 storage.adapter.get("key1").unwrap(),
635 Some("value1".to_string())
636 );
637 }
638}