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