Skip to main content

open_feature_flagd/resolver/in_process/storage/
mod.rs

1pub mod connector;
2use crate::error::FlagdError;
3pub use connector::{
4    Connector, FileConnector, GrpcStreamConnector, QueuePayload, QueuePayloadType,
5};
6use tracing::{debug, error, warn};
7
8use crate::resolver::in_process::model::feature_flag::FeatureFlag;
9use crate::resolver::in_process::model::flag_parser::FlagParser;
10use std::collections::{HashMap, HashSet};
11use std::sync::Arc;
12use std::sync::atomic::{AtomicBool, Ordering};
13use tokio::sync::RwLock;
14use tokio::sync::mpsc::{Receiver, Sender, channel};
15
16/// State of the flag storage
17#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default)]
18pub enum StorageState {
19    /// Storage is healthy and up-to-date
20    #[default]
21    Ok,
22    /// Storage data may be stale (connection issues)
23    Stale,
24    /// Storage encountered an error
25    Error,
26}
27
28/// Represents a change in storage state with affected flags
29#[derive(Debug, Clone, PartialEq)]
30pub struct StorageStateChange {
31    /// Current state of the storage
32    pub storage_state: StorageState,
33    /// Keys of flags that changed in this update
34    pub changed_flags_keys: Vec<String>,
35    /// Metadata from the sync operation
36    pub sync_metadata: HashMap<String, serde_json::Value>,
37}
38
39impl Default for StorageStateChange {
40    fn default() -> Self {
41        Self {
42            storage_state: StorageState::Ok,
43            changed_flags_keys: Vec::new(),
44            sync_metadata: HashMap::new(),
45        }
46    }
47}
48
49/// Result of querying a flag from storage
50#[derive(Debug, Clone)]
51pub struct StorageQueryResult {
52    /// The feature flag if found
53    pub feature_flag: Option<FeatureFlag>,
54    /// Metadata associated with the flag set
55    pub flag_set_metadata: HashMap<String, serde_json::Value>,
56    /// Static context associated with the sync payload
57    pub sync_metadata: HashMap<String, serde_json::Value>,
58}
59
60pub struct FlagStore {
61    flags: Arc<RwLock<HashMap<String, FeatureFlag>>>,
62    flag_set_metadata: Arc<RwLock<HashMap<String, serde_json::Value>>>,
63    sync_metadata: Arc<RwLock<HashMap<String, serde_json::Value>>>,
64    state_sender: Sender<StorageStateChange>,
65    connector: Arc<dyn Connector>,
66    shutdown: Arc<AtomicBool>,
67}
68
69impl FlagStore {
70    pub fn new(connector: Arc<dyn Connector>) -> (Self, Receiver<StorageStateChange>) {
71        let (state_sender, state_receiver) = channel(1000);
72
73        (
74            Self {
75                flags: Arc::new(RwLock::new(HashMap::new())),
76                flag_set_metadata: Arc::new(RwLock::new(HashMap::new())),
77                sync_metadata: Arc::new(RwLock::new(HashMap::new())),
78                state_sender,
79                connector,
80                shutdown: Arc::new(AtomicBool::new(false)),
81            },
82            state_receiver,
83        )
84    }
85
86    pub async fn init(&self) -> Result<(), FlagdError> {
87        debug!("Initializing flag store");
88        self.connector.init().await?;
89
90        // Handle initial sync
91        let stream = self.connector.get_stream();
92        let mut receiver = stream.lock().await;
93        debug!("Waiting for initial sync message");
94
95        if let Some(receiver_ref) = receiver.as_mut() {
96            match tokio::time::timeout(std::time::Duration::from_secs(5), receiver_ref.recv())
97                .await?
98            {
99                Some(payload) => {
100                    debug!("Received initial sync message");
101                    match payload.payload_type {
102                        QueuePayloadType::Data => {
103                            debug!("Parsing flag data: {}", &payload.flag_data);
104                            let parsing_result = FlagParser::parse_string(&payload.flag_data)?;
105                            let mut flags_write = self.flags.write().await;
106                            let mut metadata_write = self.flag_set_metadata.write().await;
107                            let mut sync_metadata_write = self.sync_metadata.write().await;
108                            let flag_keys: Vec<String> =
109                                parsing_result.flags.keys().cloned().collect();
110                            let sync_metadata = payload.metadata.unwrap_or_default();
111                            *flags_write = parsing_result.flags;
112                            *metadata_write = parsing_result.flag_set_metadata;
113                            *sync_metadata_write = sync_metadata.clone();
114                            debug!("Successfully parsed {} flags", flags_write.len());
115
116                            // Send initial state change so FileResolver knows init completed
117                            let _ = self
118                                .state_sender
119                                .send(StorageStateChange {
120                                    storage_state: StorageState::Ok,
121                                    changed_flags_keys: flag_keys,
122                                    sync_metadata,
123                                })
124                                .await;
125                        }
126                        QueuePayloadType::Error => {
127                            error!("Error in initial sync: {}", payload.flag_data);
128                            return Err(FlagdError::Sync(format!(
129                                "Error in initial sync: {}",
130                                payload.flag_data
131                            )));
132                        }
133                    }
134                }
135                None => {
136                    error!("No initial sync message received");
137                    return Err(FlagdError::Sync(
138                        "No initial sync message received".to_string(),
139                    ));
140                }
141            }
142        }
143
144        // Start continuous stream processing
145        self.start_stream_listener().await;
146        Ok(())
147    }
148
149    pub async fn shutdown(&self) -> Result<(), FlagdError> {
150        debug!("Shutting down flag store");
151        self.shutdown.store(true, Ordering::Relaxed);
152        self.connector.shutdown().await
153    }
154
155    pub async fn get_flag(&self, key: &str) -> StorageQueryResult {
156        let flags = self.flags.read().await;
157        let metadata = self.flag_set_metadata.read().await;
158        let sync_metadata = self.sync_metadata.read().await;
159
160        StorageQueryResult {
161            feature_flag: flags.get(key).cloned(),
162            flag_set_metadata: metadata.clone(),
163            sync_metadata: sync_metadata.clone(),
164        }
165    }
166
167    /// Compute which flags have changed between old and new flag sets
168    fn compute_changed_flags(
169        old_flags: &HashMap<String, FeatureFlag>,
170        new_flags: &HashMap<String, FeatureFlag>,
171    ) -> Vec<String> {
172        let mut changed = Vec::new();
173
174        // Check for modified or added flags
175        for (key, new_flag) in new_flags {
176            match old_flags.get(key) {
177                Some(old_flag) if old_flag != new_flag => {
178                    changed.push(key.clone());
179                }
180                None => {
181                    changed.push(key.clone());
182                }
183                _ => {}
184            }
185        }
186
187        // Check for deleted flags
188        let old_keys: HashSet<_> = old_flags.keys().collect();
189        let new_keys: HashSet<_> = new_flags.keys().collect();
190        for key in old_keys.difference(&new_keys) {
191            changed.push((*key).clone());
192        }
193
194        changed
195    }
196
197    async fn start_stream_listener(&self) {
198        let flags = self.flags.clone();
199        let metadata = self.flag_set_metadata.clone();
200        let sync_metadata = self.sync_metadata.clone();
201        let sender = self.state_sender.clone();
202        let stream = self.connector.get_stream();
203        let shutdown = self.shutdown.clone();
204
205        tokio::spawn(async move {
206            let mut receiver = stream.lock().await;
207            if let Some(receiver) = receiver.as_mut() {
208                while let Some(payload) = receiver.recv().await {
209                    if shutdown.load(Ordering::Relaxed) {
210                        debug!("Stream listener shutting down");
211                        break;
212                    }
213
214                    match payload.payload_type {
215                        QueuePayloadType::Data => {
216                            match FlagParser::parse_string(&payload.flag_data) {
217                                Ok(parsing_result) => {
218                                    let mut flags_write = flags.write().await;
219                                    let mut metadata_write = metadata.write().await;
220                                    let mut sync_metadata_write = sync_metadata.write().await;
221
222                                    // Compute changed flags before updating
223                                    let changed_keys = Self::compute_changed_flags(
224                                        &flags_write,
225                                        &parsing_result.flags,
226                                    );
227
228                                    let num_changes = changed_keys.len();
229                                    let payload_sync_metadata =
230                                        payload.metadata.unwrap_or_default();
231                                    *flags_write = parsing_result.flags;
232                                    *metadata_write = parsing_result.flag_set_metadata;
233                                    *sync_metadata_write = payload_sync_metadata.clone();
234
235                                    debug!(
236                                        "Flag store updated: {} flags changed ({} total flags)",
237                                        num_changes,
238                                        flags_write.len()
239                                    );
240
241                                    let _ = sender
242                                        .send(StorageStateChange {
243                                            storage_state: StorageState::Ok,
244                                            changed_flags_keys: changed_keys,
245                                            sync_metadata: payload_sync_metadata,
246                                        })
247                                        .await;
248                                }
249                                Err(e) => {
250                                    warn!("Failed to parse flag data: {}", e);
251                                    let _ = sender
252                                        .send(StorageStateChange {
253                                            storage_state: StorageState::Error,
254                                            changed_flags_keys: vec![],
255                                            sync_metadata: HashMap::new(),
256                                        })
257                                        .await;
258                                }
259                            }
260                        }
261                        QueuePayloadType::Error => {
262                            error!("Received error from connector: {}", payload.flag_data);
263                            let _ = sender
264                                .send(StorageStateChange {
265                                    storage_state: StorageState::Error,
266                                    changed_flags_keys: vec![],
267                                    sync_metadata: HashMap::new(),
268                                })
269                                .await;
270                        }
271                    }
272                }
273            }
274            debug!("Stream listener stopped");
275        });
276    }
277}
278
279#[cfg(test)]
280mod tests {
281    use super::*;
282    use serde_json::json;
283
284    fn create_test_flag(state: &str, default_variant: &str) -> FeatureFlag {
285        FeatureFlag {
286            state: state.to_string(),
287            default_variant: default_variant.to_string(),
288            variants: {
289                let mut map = HashMap::new();
290                map.insert("on".to_string(), json!(true));
291                map.insert("off".to_string(), json!(false));
292                map
293            },
294            targeting: None,
295            metadata: HashMap::new(),
296        }
297    }
298
299    #[test]
300    fn test_compute_changed_flags_no_changes() {
301        let mut flags = HashMap::new();
302        flags.insert("flag1".to_string(), create_test_flag("ENABLED", "on"));
303        flags.insert("flag2".to_string(), create_test_flag("ENABLED", "off"));
304
305        let changed = FlagStore::compute_changed_flags(&flags, &flags);
306        assert!(
307            changed.is_empty(),
308            "Expected no changes for identical flags"
309        );
310    }
311
312    #[test]
313    fn test_compute_changed_flags_added_flag() {
314        let old_flags = HashMap::new();
315        let mut new_flags = HashMap::new();
316        new_flags.insert("flag1".to_string(), create_test_flag("ENABLED", "on"));
317
318        let changed = FlagStore::compute_changed_flags(&old_flags, &new_flags);
319        assert_eq!(changed.len(), 1);
320        assert!(changed.contains(&"flag1".to_string()));
321    }
322
323    #[test]
324    fn test_compute_changed_flags_removed_flag() {
325        let mut old_flags = HashMap::new();
326        old_flags.insert("flag1".to_string(), create_test_flag("ENABLED", "on"));
327        let new_flags = HashMap::new();
328
329        let changed = FlagStore::compute_changed_flags(&old_flags, &new_flags);
330        assert_eq!(changed.len(), 1);
331        assert!(changed.contains(&"flag1".to_string()));
332    }
333
334    #[test]
335    fn test_compute_changed_flags_modified_flag() {
336        let mut old_flags = HashMap::new();
337        old_flags.insert("flag1".to_string(), create_test_flag("ENABLED", "on"));
338
339        let mut new_flags = HashMap::new();
340        new_flags.insert("flag1".to_string(), create_test_flag("ENABLED", "off")); // Changed default
341
342        let changed = FlagStore::compute_changed_flags(&old_flags, &new_flags);
343        assert_eq!(changed.len(), 1);
344        assert!(changed.contains(&"flag1".to_string()));
345    }
346
347    #[test]
348    fn test_compute_changed_flags_mixed_changes() {
349        let mut old_flags = HashMap::new();
350        old_flags.insert("flag1".to_string(), create_test_flag("ENABLED", "on"));
351        old_flags.insert("flag2".to_string(), create_test_flag("ENABLED", "on"));
352        old_flags.insert("flag3".to_string(), create_test_flag("ENABLED", "on"));
353
354        let mut new_flags = HashMap::new();
355        new_flags.insert("flag1".to_string(), create_test_flag("ENABLED", "on")); // Unchanged
356        new_flags.insert("flag2".to_string(), create_test_flag("DISABLED", "on")); // Modified
357        new_flags.insert("flag4".to_string(), create_test_flag("ENABLED", "on")); // Added
358        // flag3 is removed
359
360        let changed = FlagStore::compute_changed_flags(&old_flags, &new_flags);
361        assert_eq!(changed.len(), 3);
362        assert!(changed.contains(&"flag2".to_string())); // Modified
363        assert!(changed.contains(&"flag3".to_string())); // Removed
364        assert!(changed.contains(&"flag4".to_string())); // Added
365        assert!(!changed.contains(&"flag1".to_string())); // Unchanged
366    }
367
368    #[test]
369    fn test_storage_state_equality() {
370        assert_eq!(StorageState::Ok, StorageState::Ok);
371        assert_ne!(StorageState::Ok, StorageState::Error);
372        assert_ne!(StorageState::Error, StorageState::Stale);
373    }
374}