Skip to main content

iris_chat_protocol/
storage.rs

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    /// Opens a dedicated directory for secret session state. On Unix, access
104    /// to this directory is restricted to the current user, including on reopen.
105    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        // Existing installations may contain session keys in readable files.
117        // Restrict their enclosing directory before loading any state.
118        #[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        // Set the mode at creation: chmod after writing would expose secrets
159        // briefly, and exclusive creation also refuses existing symlinks.
160        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
326/// SQLite-backed implementation of `iris_chat_protocol::StorageAdapter`.
327/// Keys are namespaced by (owner_pubkey_hex, device_pubkey_hex) so a
328/// single database serves multiple owner accounts and devices without
329/// keyspace collisions, matching the per-(owner, device) directory
330/// scoping the previous file-backed adapter used.
331pub 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        // Upgrading an old file must not preserve its permissive mode.
546        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}