Skip to main content

bitwarden_core/client/
flags_client.rs

1//! Feature flag retrieval, persistence, and refresh from the server `/config` endpoint.
2
3use 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/// Errors returned by [`FlagsClient::fetch`].
22#[derive(Debug, thiserror::Error)]
23pub enum FetchFlagsError {
24    /// Network or deserialization error when fetching `/config`.
25    #[error("failed to fetch /config: {0}")]
26    Api(#[from] bitwarden_api_api::ApiError),
27    /// Error persisting flags or fetched_at timestamp to state registry.
28    #[error("state access error: {0}")]
29    State(#[from] bitwarden_state::SettingsError),
30}
31
32/// A client for inspecting and refreshing feature flags.
33#[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    /// Persist a flag map (e.g. from `/config`) into the state registry.
42    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    /// Retrieve the active feature flags from the state registry.
50    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    /// Fetch flags from `/config` and persist both the flag values and a `fetched_at` timestamp.
61    ///
62    /// Pass `force = true` from `from_authenticated_data` (PM-27624) immediately before
63    /// `save_to_state`, so the initial flag fetch is part of the persisted login state.
64    /// [`Client::load_from_state`] calls this with `force = false` to honour the 1-hour TTL.
65    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        // `/config` returns `serde_json::Value`; coerce to bool. Non-bool values are dropped
83        // because `Flags` only models boolean flags today.
84        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    /// Access to feature flag retrieval, persistence, and refresh.
96    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        // With no flags loaded yet, get should return defaults.
160        let initial = client.flags().get().await;
161        assert!(!initial.strict_cipher_decryption);
162
163        // Loading flags should persist them via the FLAGS setting.
164        let mut map = HashMap::new();
165        map.insert("pm-34500-strict-cipher-decryption".to_string(), true);
166        client.flags().load(map).await;
167
168        // get should now return the loaded values.
169        let loaded = client.flags().get().await;
170        assert!(loaded.strict_cipher_decryption);
171
172        // The values should be readable directly from the setting too.
173        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}