Skip to main content

bitwarden_state/
registry.rs

1use std::sync::Arc;
2
3use bitwarden_error::bitwarden_error;
4use thiserror::Error;
5
6use crate::{
7    any_map::AnyMap,
8    repository::{Repository, RepositoryItem, RepositoryMigrations},
9    sdk_managed::{Database, DatabaseConfiguration, DatabaseError, MemoryDatabase, SystemDatabase},
10    settings::{Key, Setting, SettingItem},
11};
12
13/// A registry that contains repositories for different types of items.
14/// These repositories can be either managed by the client or by the SDK itself.
15pub struct StateRegistry {
16    database: SystemDatabase,
17    client_managed: AnyMap,
18}
19
20impl std::fmt::Debug for StateRegistry {
21    fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
22        f.debug_struct("StateRegistry").finish()
23    }
24}
25
26#[allow(missing_docs)]
27#[bitwarden_error(flat)]
28#[derive(Debug, Error)]
29pub enum StateRegistryError {
30    #[error("Database is not initialized")]
31    DatabaseNotInitialized,
32
33    #[error(transparent)]
34    Database(#[from] DatabaseError),
35}
36
37impl StateRegistry {
38    /// Creates a new `StateRegistry` backed by an in-memory database.
39    pub fn new_with_memory_db() -> Self {
40        StateRegistry {
41            database: SystemDatabase::Memory(MemoryDatabase::new()),
42            client_managed: AnyMap::new(),
43        }
44    }
45
46    /// Creates a new `StateRegistry` backed by a database.
47    pub async fn new_with_db(
48        configuration: DatabaseConfiguration,
49        migrations: RepositoryMigrations,
50    ) -> Result<Self, DatabaseError> {
51        let database = SystemDatabase::initialize(configuration, migrations.clone()).await?;
52        Ok(StateRegistry {
53            database,
54            client_managed: AnyMap::new(),
55        })
56    }
57
58    /// Get a handle to a setting by its type-safe key.
59    pub fn setting<T>(&self, key: Key<T>) -> Result<Setting<T>, StateRegistryError> {
60        let repo = self.get::<SettingItem>()?;
61        Ok(Setting::new(repo, key))
62    }
63
64    /// Registers a client-managed repository into the map, associating it with its type.
65    pub fn register_client_managed<T: RepositoryItem>(&self, value: Arc<dyn Repository<T>>) {
66        self.client_managed.insert(value);
67    }
68
69    /// Retrieves a client-managed repository from the map given its type.
70    fn get_client_managed<T: RepositoryItem>(&self) -> Option<Arc<dyn Repository<T>>> {
71        self.client_managed.get()
72    }
73
74    /// Retrieves a SDK-managed repository from the database.
75    fn get_sdk_managed<T: RepositoryItem>(
76        &self,
77    ) -> Result<Arc<dyn Repository<T>>, StateRegistryError> {
78        Ok(self.database.get_repository::<T>())
79    }
80
81    /// Get a repository with fallback: prefer client-managed, fall back to SDK-managed.
82    ///
83    /// This method first attempts to retrieve a client-managed repository. If not found,
84    /// it falls back to an SDK-managed repository. Both are returned as `Arc<dyn Repository<T>>`.
85    ///
86    /// # Errors
87    /// This method never fails, but returns a Result for backwards compatibility.
88    pub fn get<T>(&self) -> Result<Arc<dyn Repository<T>>, StateRegistryError>
89    where
90        T: RepositoryItem,
91    {
92        if let Some(repo) = self.get_client_managed::<T>() {
93            return Ok(repo);
94        }
95
96        self.get_sdk_managed::<T>()
97    }
98
99    /// Wipes all state from this registry, and deletes any files or databases associated with it.
100    /// Intended to be used during logout, where the Client will be dropped right after.
101    ///
102    /// # Warning
103    ///
104    /// This closes the SDK-managed database and deletes persistent storage (SQLite file + WAL/SHM,
105    /// IndexedDB database). Outstanding [`Repository`] handles will return
106    /// [`DatabaseError::Closed`] on subsequent operations. Client-managed repositories are also
107    /// cleared.
108    pub async fn wipe(&self) -> Result<(), DatabaseError> {
109        // Clear client-managed first so a failure in the persistent-store wipe
110        // still releases the in-memory Arc references.
111        self.client_managed.clear();
112        self.database.wipe().await
113    }
114}
115
116#[cfg(test)]
117mod tests {
118    use super::*;
119    use crate::{
120        register_repository_item,
121        repository::{RepositoryError, RepositoryItem},
122        sdk_managed::DatabaseError,
123    };
124
125    macro_rules! impl_repository {
126        ($name:ident, $ty:ty) => {
127            #[async_trait::async_trait]
128            impl Repository<$ty> for $name {
129                async fn get(&self, _key: String) -> Result<Option<$ty>, RepositoryError> {
130                    Ok(Some(TestItem(self.0.clone())))
131                }
132                async fn list(&self) -> Result<Vec<$ty>, RepositoryError> {
133                    unimplemented!()
134                }
135                async fn set(&self, _key: String, _value: $ty) -> Result<(), RepositoryError> {
136                    unimplemented!()
137                }
138                async fn set_bulk(
139                    &self,
140                    _values: Vec<(String, $ty)>,
141                ) -> Result<(), RepositoryError> {
142                    unimplemented!()
143                }
144                async fn remove(&self, _key: String) -> Result<(), RepositoryError> {
145                    unimplemented!()
146                }
147                async fn remove_bulk(&self, _keys: Vec<String>) -> Result<(), RepositoryError> {
148                    unimplemented!()
149                }
150                async fn remove_all(&self) -> Result<(), RepositoryError> {
151                    unimplemented!()
152                }
153            }
154        };
155    }
156
157    use serde::{Deserialize, Serialize};
158
159    #[derive(PartialEq, Eq, Debug)]
160    struct TestA(usize);
161    #[derive(PartialEq, Eq, Debug)]
162    struct TestB(String);
163    #[derive(PartialEq, Eq, Debug)]
164    struct TestC(Vec<u8>);
165    /// A second implementation for the same item type as [`TestA`].
166    #[derive(PartialEq, Eq, Debug)]
167    struct TestD(usize);
168    #[derive(PartialEq, Eq, Debug, Serialize, Deserialize)]
169    struct TestItem<T>(T);
170
171    register_repository_item!(String => TestItem<usize>, "TestItem_usize");
172    register_repository_item!(String => TestItem<String>, "TestItem_String");
173    register_repository_item!(String => TestItem<Vec<u8>>, "TestItem_Vec");
174
175    impl_repository!(TestA, TestItem<usize>);
176    impl_repository!(TestB, TestItem<String>);
177    impl_repository!(TestC, TestItem<Vec<u8>>);
178    impl_repository!(TestD, TestItem<usize>);
179
180    #[tokio::test]
181    async fn test_state_registry() {
182        let a = Arc::new(TestA(145832));
183        let b = Arc::new(TestB("test".to_string()));
184        let c = Arc::new(TestC(vec![1, 2, 3, 4, 5, 6, 7, 8, 9]));
185
186        let map = StateRegistry::new_with_memory_db();
187
188        async fn get<T: RepositoryItem>(map: &StateRegistry) -> Option<T>
189        where
190            T::Key: Default,
191        {
192            map.get_client_managed::<T>()
193                .unwrap()
194                .get(Default::default())
195                .await
196                .unwrap()
197        }
198
199        assert!(map.get_client_managed::<TestItem<usize>>().is_none());
200        assert!(map.get_client_managed::<TestItem<String>>().is_none());
201        assert!(map.get_client_managed::<TestItem<Vec<u8>>>().is_none());
202
203        map.register_client_managed(a.clone());
204        assert_eq!(get(&map).await, Some(TestItem(a.0)));
205        assert!(map.get_client_managed::<TestItem<String>>().is_none());
206        assert!(map.get_client_managed::<TestItem<Vec<u8>>>().is_none());
207
208        map.register_client_managed(b.clone());
209        assert_eq!(get(&map).await, Some(TestItem(a.0)));
210        assert_eq!(get(&map).await, Some(TestItem(b.0.clone())));
211        assert!(map.get_client_managed::<TestItem<Vec<u8>>>().is_none());
212
213        map.register_client_managed(c.clone());
214        assert_eq!(get(&map).await, Some(TestItem(a.0)));
215        assert_eq!(get(&map).await, Some(TestItem(b.0.clone())));
216        assert_eq!(get(&map).await, Some(TestItem(c.0.clone())));
217    }
218
219    #[tokio::test]
220    async fn test_fallback_client_managed_found() {
221        let registry = StateRegistry::new_with_memory_db();
222        let test_repo = Arc::new(TestA(12345));
223
224        registry.register_client_managed(test_repo.clone());
225
226        let repo = registry.get::<TestItem<usize>>().unwrap();
227        let result = repo.get(String::new()).await.unwrap();
228
229        assert_eq!(result, Some(TestItem(12345)));
230    }
231
232    #[tokio::test]
233    async fn test_new_with_memory_db_sync() {
234        // Construct in sync context (no .await on the constructor itself)
235        let registry = StateRegistry::new_with_memory_db();
236        // Database must be accessible via async get after sync construction
237        let repo = registry.get::<TestItem<usize>>().unwrap();
238        let result = repo.get(String::new()).await;
239        // Should return Ok(None) — key not found, not an error
240        // (Note: TestItem<usize> is registered in this test module already)
241        assert!(result.is_ok());
242    }
243
244    #[tokio::test]
245    async fn test_wipe_disconnects_outstanding_repository_handles() {
246        let registry = StateRegistry::new_with_memory_db();
247        let repo = registry.get::<TestItem<usize>>().unwrap();
248        repo.set(String::new(), TestItem(42usize)).await.unwrap();
249
250        registry.wipe().await.unwrap();
251
252        assert!(matches!(
253            repo.get(String::new()).await,
254            Err(RepositoryError::Database(DatabaseError::Closed))
255        ));
256        assert!(matches!(
257            repo.list().await,
258            Err(RepositoryError::Database(DatabaseError::Closed))
259        ));
260    }
261
262    #[tokio::test]
263    async fn test_wipe_clears_client_managed() {
264        let registry = StateRegistry::new_with_memory_db();
265        registry.register_client_managed(Arc::new(TestA(99)));
266
267        registry.wipe().await.unwrap();
268
269        // Client-managed is gone; falls through to SDK-managed (now closed).
270        let repo = registry.get::<TestItem<usize>>().unwrap();
271        assert!(matches!(
272            repo.get(String::new()).await,
273            Err(RepositoryError::Database(DatabaseError::Closed))
274        ));
275    }
276
277    #[tokio::test]
278    async fn test_wipe_is_idempotent() {
279        let registry = StateRegistry::new_with_memory_db();
280        registry.wipe().await.unwrap();
281        registry.wipe().await.unwrap();
282    }
283
284    #[tokio::test]
285    async fn test_setting_on_memory_db() {
286        use crate::register_setting_key;
287        register_setting_key!(const TEST_SETTING: String = "test_registry_setting_key");
288
289        let registry = StateRegistry::new_with_memory_db();
290        let setting = registry.setting(TEST_SETTING).unwrap();
291
292        // Value must not exist initially
293        assert_eq!(setting.get().await.unwrap(), None::<String>);
294
295        // Update and read back
296        setting.update("hello".to_string()).await.unwrap();
297        assert_eq!(setting.get().await.unwrap(), Some("hello".to_string()));
298
299        // Delete and confirm gone
300        setting.delete().await.unwrap();
301        assert_eq!(setting.get().await.unwrap(), None::<String>);
302    }
303
304    /// The concrete implementation is erased by the coercion to `Arc<dyn Repository<T>>`, so two
305    /// implementations of the same item type share a slot and the later registration wins.
306    #[tokio::test]
307    async fn test_register_client_managed_replaces_previous_implementation() {
308        let registry = StateRegistry::new_with_memory_db();
309
310        registry.register_client_managed(Arc::new(TestA(1)));
311        registry.register_client_managed(Arc::new(TestD(2)));
312
313        let repo = registry.get_client_managed::<TestItem<usize>>().unwrap();
314        assert_eq!(repo.get(String::new()).await.unwrap(), Some(TestItem(2)));
315    }
316}