Skip to main content

dynamo_runtime/storage/kv/
etcd.rs

1// SPDX-FileCopyrightText: Copyright (c) 2024-2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved.
2// SPDX-License-Identifier: Apache-2.0
3
4use std::collections::HashMap;
5use std::pin::Pin;
6use std::time::Duration;
7
8use crate::transports::etcd;
9use async_stream::stream;
10use async_trait::async_trait;
11use etcd_client::{Compare, CompareOp, EventType, PutOptions, Txn, TxnOp, WatchOptions};
12
13use super::{Bucket, Key, KeyValue, Store, StoreError, StoreOutcome, WatchEvent};
14
15#[derive(Clone)]
16pub struct EtcdStore {
17    client: etcd::Client,
18}
19
20impl EtcdStore {
21    pub fn new(client: etcd::Client) -> Self {
22        Self { client }
23    }
24}
25
26#[async_trait]
27impl Store for EtcdStore {
28    type Bucket = EtcdBucket;
29
30    /// A "bucket" in etcd is a path prefix
31    async fn get_or_create_bucket(
32        &self,
33        bucket_name: &str,
34        _ttl: Option<Duration>, // TODO ttl not used yet
35    ) -> Result<Self::Bucket, StoreError> {
36        Ok(EtcdBucket {
37            client: self.client.clone(),
38            bucket_name: bucket_name.to_string(),
39        })
40    }
41
42    /// A "bucket" in etcd is a path prefix. This creates an EtcdBucket object without doing
43    /// any network calls.
44    async fn get_bucket(&self, bucket_name: &str) -> Result<Option<Self::Bucket>, StoreError> {
45        Ok(Some(EtcdBucket {
46            client: self.client.clone(),
47            bucket_name: bucket_name.to_string(),
48        }))
49    }
50
51    fn connection_id(&self) -> u64 {
52        self.client.lease_id()
53    }
54
55    fn shutdown(&self) {
56        // Revoke the lease? etcd will do it for us on disconnect.
57    }
58}
59
60pub struct EtcdBucket {
61    client: etcd::Client,
62    bucket_name: String,
63}
64
65#[async_trait]
66impl Bucket for EtcdBucket {
67    async fn insert(
68        &self,
69        key: &Key,
70        value: bytes::Bytes,
71        // "version" in etcd speak. revision is a global cluster-wide value
72        revision: u64,
73    ) -> Result<StoreOutcome, StoreError> {
74        let version = revision;
75        if version == 0 {
76            self.create(key, value).await
77        } else {
78            self.update(key, value, version).await
79        }
80    }
81
82    async fn get(&self, key: &Key) -> Result<Option<bytes::Bytes>, StoreError> {
83        let k = make_key(&self.bucket_name, key);
84        tracing::trace!("etcd get: {k}");
85
86        let mut kvs = self
87            .client
88            .kv_get(k, None)
89            .await
90            .map_err(|e| StoreError::EtcdError(e.to_string()))?;
91        if kvs.is_empty() {
92            return Ok(None);
93        }
94        let (_, val) = kvs.swap_remove(0).into_key_value();
95        Ok(Some(val.into()))
96    }
97
98    async fn delete(&self, key: &Key) -> Result<(), StoreError> {
99        let k = make_key(&self.bucket_name, key);
100        tracing::trace!("etcd delete: {k}");
101        let _ = self
102            .client
103            .kv_delete(k, None)
104            .await
105            .map_err(|e| StoreError::EtcdError(e.to_string()))?;
106        Ok(())
107    }
108
109    async fn watch(
110        &self,
111    ) -> Result<Pin<Box<dyn futures::Stream<Item = WatchEvent> + Send + 'life0>>, StoreError> {
112        let prefix = make_key(&self.bucket_name, &"".into());
113        tracing::trace!("etcd watch: {prefix}");
114        let watcher = self
115            .client
116            .kv_watch_prefix(&prefix)
117            .await
118            .map_err(|e| StoreError::EtcdError(e.to_string()))?;
119        let (_, mut watch_stream) = watcher.dissolve();
120        let output = stream! {
121            while let Some(event) = watch_stream.recv().await {
122                match event {
123                    etcd::WatchEvent::Put(kv) => {
124                        let (k, v) = kv.into_key_value();
125                        let key = match String::from_utf8(k) {
126                            Ok(k) => Key::new(k),
127                            Err(err) => {
128                                tracing::error!(%err, prefix, "Invalid UTF8 in etcd key");
129                                continue;
130                            }
131                        };
132                        let item = KeyValue::new(key, v.into());
133                        yield WatchEvent::Put(item);
134                    }
135                    etcd::WatchEvent::Delete(kv) => {
136                        let (k, _) = kv.into_key_value();
137                        let key = match String::from_utf8(k) {
138                            Ok(k) => Key::new(k),
139                            Err(err) => {
140                                tracing::error!(%err, prefix, "Invalid UTF8 in etcd key");
141                                continue;
142                            }
143                        };
144                        yield WatchEvent::Delete(key);
145                    }
146                    etcd::WatchEvent::Resync(kvs) => {
147                        let mut snapshot = HashMap::with_capacity(kvs.len());
148                        for kv in kvs {
149                            let (k, v) = kv.into_key_value();
150                            let key = match String::from_utf8(k) {
151                                Ok(k) => Key::new(k),
152                                Err(err) => {
153                                    tracing::error!(%err, prefix, "Invalid UTF8 in etcd resync key");
154                                    continue;
155                                }
156                            };
157                            snapshot.insert(key, v.into());
158                        }
159                        yield WatchEvent::Resync(snapshot);
160                    }
161                }
162            }
163        };
164        Ok(Box::pin(output))
165    }
166
167    async fn entries(&self) -> Result<HashMap<Key, bytes::Bytes>, StoreError> {
168        let k = make_key(&self.bucket_name, &"".into());
169        tracing::trace!("etcd entries: {k}");
170
171        let resp = self
172            .client
173            .kv_get_prefix(k)
174            .await
175            .map_err(|e| StoreError::EtcdError(e.to_string()))?;
176        let out: HashMap<Key, bytes::Bytes> = resp
177            .into_iter()
178            .map(|kv| {
179                let (k, v) = kv.into_key_value();
180                (Key::new(String::from_utf8_lossy(&k).to_string()), v.into())
181            })
182            .collect();
183
184        Ok(out)
185    }
186}
187
188impl EtcdBucket {
189    async fn create(
190        &self,
191        key: &Key,
192        value: impl Into<Vec<u8>>,
193    ) -> Result<StoreOutcome, StoreError> {
194        let k = make_key(&self.bucket_name, key);
195        tracing::trace!("etcd create: {k}");
196
197        match self
198            .client
199            .kv_create(k.as_str(), value.into(), None)
200            .await
201            .map_err(|e| StoreError::EtcdError(e.to_string()))?
202        {
203            None => {
204                // Key was created successfully
205                Ok(StoreOutcome::Created(1)) // version of new key is always 1
206            }
207            Some(revision) => Ok(StoreOutcome::Exists(revision)),
208        }
209    }
210
211    async fn update(
212        &self,
213        key: &Key,
214        value: impl AsRef<[u8]>,
215        revision: u64,
216    ) -> Result<StoreOutcome, StoreError> {
217        let version = revision;
218        let k = make_key(&self.bucket_name, key);
219        tracing::trace!("etcd update: {k}");
220
221        let kvs = self
222            .client
223            .kv_get(k.clone(), None)
224            .await
225            .map_err(|e| StoreError::EtcdError(e.to_string()))?;
226        if kvs.is_empty() {
227            return Err(StoreError::MissingKey(key.to_string()));
228        }
229        let current_version = kvs.first().unwrap().version() as u64;
230        if current_version != version + 1 {
231            tracing::warn!(
232                current_version,
233                attempted_next_version = version,
234                %key,
235                "update: Wrong revision"
236            );
237            // NATS does a resync_update, overwriting the key anyway and getting the new revision.
238            // So we do too in etcd.
239        }
240
241        let put_options = PutOptions::new()
242            .with_lease(self.client.lease_id() as i64)
243            .with_prev_key();
244        let mut put_resp = self
245            .client
246            .kv_put_with_options(k, value, Some(put_options))
247            .await
248            .map_err(|e| StoreError::EtcdError(e.to_string()))?;
249        Ok(match put_resp.take_prev_key() {
250            // Should this be an error?
251            // The key was deleted between our get and put. We re-created it.
252            // Version of new key is always 1.
253            // <https://etcd.io/docs/v3.5/learning/data_model/>
254            None => StoreOutcome::Created(1),
255            // Expected case, success
256            Some(kv) if kv.version() as u64 == version + 1 => StoreOutcome::Created(version),
257            // Should this be an error? Something updated the version between our get and put
258            Some(kv) => StoreOutcome::Created(kv.version() as u64 + 1),
259        })
260    }
261}
262
263fn make_key(bucket_name: &str, key: &Key) -> String {
264    [bucket_name.to_string(), key.to_string()].join("/")
265}
266
267#[cfg(feature = "integration")]
268#[cfg(test)]
269mod concurrent_create_tests {
270    use super::*;
271    use crate::Runtime;
272    use crate::transports::etcd as etcd_transport;
273    use std::sync::Arc;
274    use tokio::sync::Barrier;
275
276    #[test]
277    fn test_concurrent_etcd_create_race_condition() {
278        let rt = Runtime::single_threaded().unwrap();
279        let rt_clone = rt.clone();
280
281        rt_clone.primary().block_on(async move {
282            let etcd_client =
283                etcd_transport::Client::new(etcd_transport::ClientOptions::default(), rt)
284                    .await
285                    .unwrap();
286            let storage = crate::storage::kv::Manager::etcd(etcd_client);
287            test_concurrent_create(&storage).await.unwrap();
288        });
289    }
290
291    async fn test_concurrent_create(
292        storage: &crate::storage::kv::Manager,
293    ) -> Result<(), StoreError> {
294        // Create a bucket for testing
295        let bucket = Arc::new(tokio::sync::Mutex::new(
296            storage
297                .get_or_create_bucket("test_concurrent_bucket", None)
298                .await?,
299        ));
300
301        // Number of concurrent workers
302        let num_workers = 10;
303        let barrier = Arc::new(Barrier::new(num_workers));
304
305        // Shared test data
306        let test_key: Key = Key::new(format!("concurrent_test_key_{}", uuid::Uuid::new_v4()));
307        let test_value = "test_value";
308
309        // Spawn multiple tasks that will all try to create the same key simultaneously
310        let mut handles = Vec::new();
311        let success_count = Arc::new(tokio::sync::Mutex::new(0));
312        let exists_count = Arc::new(tokio::sync::Mutex::new(0));
313
314        for worker_id in 0..num_workers {
315            let bucket_clone = bucket.clone();
316            let barrier_clone = barrier.clone();
317            let key_clone = test_key.clone();
318            let value_clone = format!("{}_from_worker_{}", test_value, worker_id);
319            let success_count_clone = success_count.clone();
320            let exists_count_clone = exists_count.clone();
321
322            let handle = tokio::spawn(async move {
323                // Wait for all workers to be ready
324                barrier_clone.wait().await;
325
326                // All workers try to create the same key at the same time
327                let result = bucket_clone
328                    .lock()
329                    .await
330                    .insert(&key_clone, value_clone.into(), 0)
331                    .await;
332
333                match result {
334                    Ok(StoreOutcome::Created(version)) => {
335                        println!(
336                            "Worker {} successfully created key with version {}",
337                            worker_id, version
338                        );
339                        let mut count = success_count_clone.lock().await;
340                        *count += 1;
341                        Ok(version)
342                    }
343                    Ok(StoreOutcome::Exists(version)) => {
344                        println!(
345                            "Worker {} found key already exists with version {}",
346                            worker_id, version
347                        );
348                        let mut count = exists_count_clone.lock().await;
349                        *count += 1;
350                        Ok(version)
351                    }
352                    Err(e) => {
353                        println!("Worker {} got error: {:?}", worker_id, e);
354                        Err(e)
355                    }
356                }
357            });
358
359            handles.push(handle);
360        }
361
362        // Wait for all workers to complete
363        let mut results = Vec::new();
364        for handle in handles {
365            let result = handle.await.unwrap();
366            if let Ok(version) = result {
367                results.push(version);
368            }
369        }
370
371        // Verify results
372        let final_success_count = *success_count.lock().await;
373        let final_exists_count = *exists_count.lock().await;
374
375        println!(
376            "Final counts - Created: {}, Exists: {}",
377            final_success_count, final_exists_count
378        );
379
380        // CRITICAL ASSERTIONS:
381        // 1. Exactly ONE worker should have successfully created the key
382        assert_eq!(
383            final_success_count, 1,
384            "Exactly one worker should create the key"
385        );
386
387        // 2. All other workers should have gotten "Exists" response
388        assert_eq!(
389            final_exists_count,
390            num_workers - 1,
391            "All other workers should see key exists"
392        );
393
394        // 3. Total successful operations should equal number of workers
395        assert_eq!(
396            results.len(),
397            num_workers,
398            "All workers should complete successfully"
399        );
400
401        // 4. Verify the key actually exists in etcd
402        let stored_value = bucket.lock().await.get(&test_key).await?;
403        assert!(stored_value.is_some(), "Key should exist in etcd");
404
405        // 5. The stored value should be from one of the workers
406        let stored_str = String::from_utf8(stored_value.unwrap().to_vec()).unwrap();
407        assert!(
408            stored_str.starts_with(test_value),
409            "Stored value should match expected prefix"
410        );
411
412        // Clean up
413        bucket.lock().await.delete(&test_key).await?;
414
415        Ok(())
416    }
417}