Skip to main content

iris_chat_protocol/
storage.rs

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