bitwarden_core/client/
flags_client.rs1use std::{collections::HashMap, sync::Arc};
4
5use bitwarden_state::Setting;
6use chrono::{DateTime, Duration, Utc};
7#[cfg(feature = "wasm")]
8use wasm_bindgen::prelude::*;
9
10use crate::{
11 Client,
12 client::{
13 flags::Flags,
14 internal::ApiConfigurations,
15 persisted_state::{FLAGS, FLAGS_FETCHED_AT},
16 },
17};
18
19const FLAGS_TTL: Duration = Duration::hours(1);
20
21#[derive(Debug, thiserror::Error)]
23pub enum FetchFlagsError {
24 #[error("failed to fetch /config: {0}")]
26 Api(#[from] bitwarden_api_api::ApiError),
27 #[error("state access error: {0}")]
29 State(#[from] bitwarden_state::SettingsError),
30}
31
32#[cfg_attr(feature = "wasm", wasm_bindgen)]
34pub struct FlagsClient {
35 flags: Setting<Flags>,
36 flags_fetched_at: Setting<DateTime<Utc>>,
37 api_configurations: Arc<ApiConfigurations>,
38}
39
40impl FlagsClient {
41 pub async fn load(&self, flags: HashMap<String, bool>) {
43 let flags = Flags::load_from_map(flags);
44 if let Err(e) = self.flags.update(flags).await {
45 tracing::warn!("Failed to persist flags: {e}");
46 }
47 }
48
49 pub async fn get(&self) -> Flags {
51 match self.flags.get().await {
52 Ok(flags) => flags.unwrap_or_default(),
53 Err(e) => {
54 tracing::warn!("Failed to read flags, using defaults: {e}");
55 Flags::default()
56 }
57 }
58 }
59
60 pub async fn fetch(&self, force: bool) -> Result<(), FetchFlagsError> {
66 if !force {
67 let last: Option<DateTime<Utc>> = self.flags_fetched_at.get().await?;
68 if let Some(fetched_at) = last
69 && Utc::now().signed_duration_since(fetched_at) < FLAGS_TTL
70 {
71 return Ok(());
72 }
73 }
74
75 let config = self
76 .api_configurations
77 .api_client
78 .config_api()
79 .get_configs()
80 .await?;
81 let feature_states = config.feature_states.unwrap_or_default();
82 let bool_map = feature_states
85 .into_iter()
86 .filter_map(|(k, v)| v.as_bool().map(|b| (k, b)))
87 .collect();
88 self.load(bool_map).await;
89 self.flags_fetched_at.update(Utc::now()).await?;
90 Ok(())
91 }
92}
93
94impl Client {
95 pub fn flags(&self) -> FlagsClient {
97 let registry = &self.internal.state_registry;
98 FlagsClient {
99 flags: registry
100 .setting(FLAGS)
101 .expect("Settings repository must be registered on the state registry"),
102 flags_fetched_at: registry
103 .setting(FLAGS_FETCHED_AT)
104 .expect("Settings repository must be registered on the state registry"),
105 api_configurations: self.internal.api_configurations.clone(),
106 }
107 }
108}
109
110#[cfg(test)]
111mod tests {
112 use serde_json::json;
113 use wiremock::{
114 Mock, MockServer, ResponseTemplate,
115 matchers::{method, path},
116 };
117
118 use super::*;
119 use crate::{ClientSettings, DeviceType};
120
121 fn settings_for(server: &MockServer) -> ClientSettings {
122 ClientSettings {
123 identity_url: format!("http://{}", server.address()),
124 api_url: format!("http://{}", server.address()),
125 user_agent: "flags-tests".to_string(),
126 device_type: DeviceType::SDK,
127 device_identifier: None,
128 bitwarden_client_version: None,
129 bitwarden_package_type: None,
130 }
131 }
132
133 async fn write_fetched_at(client: &Client, at: DateTime<Utc>) {
134 client
135 .internal
136 .state_registry
137 .setting(FLAGS_FETCHED_AT)
138 .unwrap()
139 .update(at)
140 .await
141 .unwrap();
142 }
143
144 async fn read_fetched_at(client: &Client) -> Option<DateTime<Utc>> {
145 client
146 .internal
147 .state_registry
148 .setting(FLAGS_FETCHED_AT)
149 .unwrap()
150 .get()
151 .await
152 .unwrap()
153 }
154
155 #[tokio::test]
156 async fn load_round_trips_through_setting() {
157 let client = Client::new(None);
158
159 let initial = client.flags().get().await;
161 assert!(!initial.strict_cipher_decryption);
162
163 let mut map = HashMap::new();
165 map.insert("pm-34500-strict-cipher-decryption".to_string(), true);
166 client.flags().load(map).await;
167
168 let loaded = client.flags().get().await;
170 assert!(loaded.strict_cipher_decryption);
171
172 let persisted = client
174 .internal
175 .state_registry
176 .setting(FLAGS)
177 .unwrap()
178 .get()
179 .await
180 .unwrap()
181 .expect("flags should be persisted after load");
182 assert!(persisted.strict_cipher_decryption);
183 }
184
185 #[tokio::test]
186 async fn fetch_force_persists_flags_and_timestamp() {
187 let server = MockServer::start().await;
188 Mock::given(method("GET"))
189 .and(path("/config"))
190 .respond_with(ResponseTemplate::new(200).set_body_json(json!({
191 "featureStates": { "pm-34500-strict-cipher-decryption": true }
192 })))
193 .expect(1)
194 .mount(&server)
195 .await;
196
197 let client = Client::new(Some(settings_for(&server)));
198 let before = Utc::now();
199 client.flags().fetch(true).await.unwrap();
200
201 assert!(client.flags().get().await.strict_cipher_decryption);
202 let fetched_at = read_fetched_at(&client)
203 .await
204 .expect("fetched_at must be set after a successful fetch");
205 assert!(fetched_at >= before);
206 }
207
208 #[tokio::test]
209 async fn fetch_skips_when_fresh() {
210 let server = MockServer::start().await;
211 Mock::given(method("GET"))
212 .and(path("/config"))
213 .respond_with(ResponseTemplate::new(200).set_body_json(json!({})))
214 .expect(0)
215 .mount(&server)
216 .await;
217
218 let client = Client::new(Some(settings_for(&server)));
219 write_fetched_at(&client, Utc::now() - Duration::minutes(5)).await;
220
221 client.flags().fetch(false).await.unwrap();
222 }
223
224 #[tokio::test]
225 async fn fetch_force_ignores_ttl() {
226 let server = MockServer::start().await;
227 Mock::given(method("GET"))
228 .and(path("/config"))
229 .respond_with(ResponseTemplate::new(200).set_body_json(json!({})))
230 .expect(1)
231 .mount(&server)
232 .await;
233
234 let client = Client::new(Some(settings_for(&server)));
235 write_fetched_at(&client, Utc::now() - Duration::minutes(5)).await;
236
237 client.flags().fetch(true).await.unwrap();
238 }
239
240 #[tokio::test]
241 async fn fetch_refreshes_when_stale() {
242 let server = MockServer::start().await;
243 Mock::given(method("GET"))
244 .and(path("/config"))
245 .respond_with(ResponseTemplate::new(200).set_body_json(json!({
246 "featureStates": { "pm-34500-strict-cipher-decryption": true }
247 })))
248 .expect(1)
249 .mount(&server)
250 .await;
251
252 let client = Client::new(Some(settings_for(&server)));
253 let stale = Utc::now() - Duration::hours(2);
254 write_fetched_at(&client, stale).await;
255
256 client.flags().fetch(false).await.unwrap();
257
258 assert!(client.flags().get().await.strict_cipher_decryption);
259 let fetched_at = read_fetched_at(&client).await.unwrap();
260 assert!(fetched_at > stale);
261 }
262
263 #[tokio::test]
264 async fn fetch_network_error_is_non_fatal_and_preserves_flags() {
265 let server = MockServer::start().await;
266 Mock::given(method("GET"))
267 .and(path("/config"))
268 .respond_with(ResponseTemplate::new(500))
269 .mount(&server)
270 .await;
271
272 let client = Client::new(Some(settings_for(&server)));
273 client
274 .flags()
275 .load(HashMap::from([(
276 "pm-34500-strict-cipher-decryption".to_string(),
277 true,
278 )]))
279 .await;
280
281 assert!(client.flags().fetch(true).await.is_err());
282 assert!(
283 client.flags().get().await.strict_cipher_decryption,
284 "previously persisted flags must survive a failed fetch"
285 );
286 }
287}