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
13pub 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 pub fn new_with_memory_db() -> Self {
40 StateRegistry {
41 database: SystemDatabase::Memory(MemoryDatabase::new()),
42 client_managed: AnyMap::new(),
43 }
44 }
45
46 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 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 pub fn register_client_managed<T: RepositoryItem>(&self, value: Arc<dyn Repository<T>>) {
66 self.client_managed.insert(value);
67 }
68
69 fn get_client_managed<T: RepositoryItem>(&self) -> Option<Arc<dyn Repository<T>>> {
71 self.client_managed.get()
72 }
73
74 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 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 pub async fn wipe(&self) -> Result<(), DatabaseError> {
109 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 #[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 let registry = StateRegistry::new_with_memory_db();
236 let repo = registry.get::<TestItem<usize>>().unwrap();
238 let result = repo.get(String::new()).await;
239 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 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 assert_eq!(setting.get().await.unwrap(), None::<String>);
294
295 setting.update("hello".to_string()).await.unwrap();
297 assert_eq!(setting.get().await.unwrap(), Some("hello".to_string()));
298
299 setting.delete().await.unwrap();
301 assert_eq!(setting.get().await.unwrap(), None::<String>);
302 }
303
304 #[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}