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::PutOptions;
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 compare_and_replace(
83        &self,
84        key: &Key,
85        expected: bytes::Bytes,
86        value: bytes::Bytes,
87    ) -> Result<StoreOutcome, StoreError> {
88        let k = make_key(&self.bucket_name, key);
89        match self
90            .client
91            .kv_compare_and_put(k, expected, value, None)
92            .await
93            .map_err(|error| StoreError::EtcdError(error.to_string()))?
94        {
95            etcd::CompareAndPutOutcome::Updated => Ok(StoreOutcome::Created(0)),
96            etcd::CompareAndPutOutcome::Missing => Err(StoreError::MissingKey(key.to_string())),
97            etcd::CompareAndPutOutcome::Conflict => Err(StoreError::Retry),
98        }
99    }
100
101    async fn get(&self, key: &Key) -> Result<Option<bytes::Bytes>, StoreError> {
102        let k = make_key(&self.bucket_name, key);
103        tracing::trace!("etcd get: {k}");
104
105        let mut kvs = self
106            .client
107            .kv_get(k, None)
108            .await
109            .map_err(|e| StoreError::EtcdError(e.to_string()))?;
110        if kvs.is_empty() {
111            return Ok(None);
112        }
113        let (_, val) = kvs.swap_remove(0).into_key_value();
114        Ok(Some(val.into()))
115    }
116
117    async fn delete(&self, key: &Key) -> Result<(), StoreError> {
118        let k = make_key(&self.bucket_name, key);
119        tracing::trace!("etcd delete: {k}");
120        let _ = self
121            .client
122            .kv_delete(k, None)
123            .await
124            .map_err(|e| StoreError::EtcdError(e.to_string()))?;
125        Ok(())
126    }
127
128    async fn watch(
129        &self,
130    ) -> Result<Pin<Box<dyn futures::Stream<Item = WatchEvent> + Send + 'life0>>, StoreError> {
131        let prefix = make_key(&self.bucket_name, &"".into());
132        tracing::trace!("etcd watch: {prefix}");
133        let watcher = self
134            .client
135            .kv_get_and_watch_prefix(&prefix)
136            .await
137            .map_err(|e| StoreError::EtcdError(e.to_string()))?;
138        let (_, mut watch_stream) = watcher.dissolve();
139        let output = stream! {
140            while let Some(event) = watch_stream.recv().await {
141                match event {
142                    etcd::WatchEvent::Put(kv) => {
143                        let (k, v) = kv.into_key_value();
144                        let key = match String::from_utf8(k) {
145                            Ok(k) => Key::new(k),
146                            Err(err) => {
147                                tracing::error!(%err, prefix, "Invalid UTF8 in etcd key");
148                                continue;
149                            }
150                        };
151                        let item = KeyValue::new(key, v.into());
152                        yield WatchEvent::Put(item);
153                    }
154                    etcd::WatchEvent::Delete(kv) => {
155                        let (k, _) = kv.into_key_value();
156                        let key = match String::from_utf8(k) {
157                            Ok(k) => Key::new(k),
158                            Err(err) => {
159                                tracing::error!(%err, prefix, "Invalid UTF8 in etcd key");
160                                continue;
161                            }
162                        };
163                        yield WatchEvent::Delete(key);
164                    }
165                    etcd::WatchEvent::Resync(kvs) => {
166                        let mut snapshot = HashMap::with_capacity(kvs.len());
167                        for kv in kvs {
168                            let (k, v) = kv.into_key_value();
169                            let key = match String::from_utf8(k) {
170                                Ok(k) => Key::new(k),
171                                Err(err) => {
172                                    tracing::error!(%err, prefix, "Invalid UTF8 in etcd resync key");
173                                    continue;
174                                }
175                            };
176                            snapshot.insert(key, v.into());
177                        }
178                        yield WatchEvent::Resync(snapshot);
179                    }
180                }
181            }
182        };
183        Ok(Box::pin(output))
184    }
185
186    async fn entries(&self) -> Result<HashMap<Key, bytes::Bytes>, StoreError> {
187        let k = make_key(&self.bucket_name, &"".into());
188        tracing::trace!("etcd entries: {k}");
189
190        let resp = self
191            .client
192            .kv_get_prefix(k)
193            .await
194            .map_err(|e| StoreError::EtcdError(e.to_string()))?;
195        let out: HashMap<Key, bytes::Bytes> = resp
196            .into_iter()
197            .map(|kv| {
198                let (k, v) = kv.into_key_value();
199                (Key::new(String::from_utf8_lossy(&k).to_string()), v.into())
200            })
201            .collect();
202
203        Ok(out)
204    }
205}
206
207impl EtcdBucket {
208    async fn create(
209        &self,
210        key: &Key,
211        value: impl Into<Vec<u8>>,
212    ) -> Result<StoreOutcome, StoreError> {
213        let k = make_key(&self.bucket_name, key);
214        tracing::trace!("etcd create: {k}");
215
216        match self
217            .client
218            .kv_create(k.as_str(), value.into(), None)
219            .await
220            .map_err(|e| StoreError::EtcdError(e.to_string()))?
221        {
222            None => {
223                // Key was created successfully
224                Ok(StoreOutcome::Created(1)) // version of new key is always 1
225            }
226            Some(revision) => Ok(StoreOutcome::Exists(revision)),
227        }
228    }
229
230    async fn update(
231        &self,
232        key: &Key,
233        value: impl AsRef<[u8]>,
234        revision: u64,
235    ) -> Result<StoreOutcome, StoreError> {
236        let version = revision;
237        let k = make_key(&self.bucket_name, key);
238        tracing::trace!("etcd update: {k}");
239
240        let kvs = self
241            .client
242            .kv_get(k.clone(), None)
243            .await
244            .map_err(|e| StoreError::EtcdError(e.to_string()))?;
245        if kvs.is_empty() {
246            return Err(StoreError::MissingKey(key.to_string()));
247        }
248        let current_version = kvs.first().unwrap().version() as u64;
249        if current_version != version + 1 {
250            tracing::warn!(
251                current_version,
252                attempted_next_version = version,
253                %key,
254                "update: Wrong revision"
255            );
256            // NATS does a resync_update, overwriting the key anyway and getting the new revision.
257            // So we do too in etcd.
258        }
259
260        let put_options = PutOptions::new()
261            .with_lease(self.client.lease_id() as i64)
262            .with_prev_key();
263        let mut put_resp = self
264            .client
265            .kv_put_with_options(k, value, Some(put_options))
266            .await
267            .map_err(|e| StoreError::EtcdError(e.to_string()))?;
268        Ok(match put_resp.take_prev_key() {
269            // Should this be an error?
270            // The key was deleted between our get and put. We re-created it.
271            // Version of new key is always 1.
272            // <https://etcd.io/docs/v3.5/learning/data_model/>
273            None => StoreOutcome::Created(1),
274            // Expected case, success
275            Some(kv) if kv.version() as u64 == version + 1 => StoreOutcome::Created(version),
276            // Should this be an error? Something updated the version between our get and put
277            Some(kv) => StoreOutcome::Created(kv.version() as u64 + 1),
278        })
279    }
280}
281
282fn make_key(bucket_name: &str, key: &Key) -> String {
283    [bucket_name.to_string(), key.to_string()].join("/")
284}
285
286#[cfg(feature = "integration")]
287#[cfg(test)]
288mod concurrent_create_tests {
289    use super::*;
290    use crate::Runtime;
291    use crate::transports::etcd as etcd_transport;
292    use std::sync::Arc;
293    use tokio::sync::Barrier;
294
295    #[test]
296    fn test_concurrent_etcd_create_race_condition() {
297        let rt = Runtime::single_threaded().unwrap();
298        let rt_clone = rt.clone();
299
300        rt_clone.primary().block_on(async move {
301            let etcd_client =
302                etcd_transport::Client::new(etcd_transport::ClientOptions::default(), rt)
303                    .await
304                    .unwrap();
305            let storage = crate::storage::kv::Manager::etcd(etcd_client);
306            test_concurrent_create(&storage).await.unwrap();
307        });
308    }
309
310    #[test]
311    fn delete_wins_race_with_compare_and_replace() {
312        let rt = Runtime::single_threaded().unwrap();
313        let rt_clone = rt.clone();
314
315        rt_clone.primary().block_on(async move {
316            let etcd_client =
317                etcd_transport::Client::new(etcd_transport::ClientOptions::default(), rt)
318                    .await
319                    .unwrap();
320            let storage = crate::storage::kv::Manager::etcd(etcd_client);
321            let bucket = Arc::new(
322                storage
323                    .get_or_create_bucket("test_compare_and_replace_bucket", None)
324                    .await
325                    .unwrap(),
326            );
327            let key = Key::new(format!("model_{}", uuid::Uuid::new_v4()));
328            bucket.insert(&key, "old".into(), 0).await.unwrap();
329
330            let barrier = Arc::new(Barrier::new(3));
331            let task_bucket = bucket.clone();
332            let task_key = key.clone();
333            let task_barrier = barrier.clone();
334            let update = tokio::spawn(async move {
335                task_barrier.wait().await;
336                task_bucket
337                    .compare_and_replace(&task_key, "old".into(), "new".into())
338                    .await
339            });
340            let task_bucket = bucket.clone();
341            let task_key = key.clone();
342            let task_barrier = barrier.clone();
343            let delete = tokio::spawn(async move {
344                task_barrier.wait().await;
345                task_bucket.delete(&task_key).await
346            });
347
348            barrier.wait().await;
349            let update_result = update.await.unwrap();
350            delete.await.unwrap().unwrap();
351            assert!(
352                update_result.is_ok() || matches!(update_result, Err(StoreError::MissingKey(_)))
353            );
354            assert_eq!(bucket.get(&key).await.unwrap(), None);
355        });
356    }
357
358    async fn test_concurrent_create(
359        storage: &crate::storage::kv::Manager,
360    ) -> Result<(), StoreError> {
361        // Create a bucket for testing
362        let bucket = Arc::new(tokio::sync::Mutex::new(
363            storage
364                .get_or_create_bucket("test_concurrent_bucket", None)
365                .await?,
366        ));
367
368        // Number of concurrent workers
369        let num_workers = 10;
370        let barrier = Arc::new(Barrier::new(num_workers));
371
372        // Shared test data
373        let test_key: Key = Key::new(format!("concurrent_test_key_{}", uuid::Uuid::new_v4()));
374        let test_value = "test_value";
375
376        // Spawn multiple tasks that will all try to create the same key simultaneously
377        let mut handles = Vec::new();
378        let success_count = Arc::new(tokio::sync::Mutex::new(0));
379        let exists_count = Arc::new(tokio::sync::Mutex::new(0));
380
381        for worker_id in 0..num_workers {
382            let bucket_clone = bucket.clone();
383            let barrier_clone = barrier.clone();
384            let key_clone = test_key.clone();
385            let value_clone = format!("{}_from_worker_{}", test_value, worker_id);
386            let success_count_clone = success_count.clone();
387            let exists_count_clone = exists_count.clone();
388
389            let handle = tokio::spawn(async move {
390                // Wait for all workers to be ready
391                barrier_clone.wait().await;
392
393                // All workers try to create the same key at the same time
394                let result = bucket_clone
395                    .lock()
396                    .await
397                    .insert(&key_clone, value_clone.into(), 0)
398                    .await;
399
400                match result {
401                    Ok(StoreOutcome::Created(version)) => {
402                        println!(
403                            "Worker {} successfully created key with version {}",
404                            worker_id, version
405                        );
406                        let mut count = success_count_clone.lock().await;
407                        *count += 1;
408                        Ok(version)
409                    }
410                    Ok(StoreOutcome::Exists(version)) => {
411                        println!(
412                            "Worker {} found key already exists with version {}",
413                            worker_id, version
414                        );
415                        let mut count = exists_count_clone.lock().await;
416                        *count += 1;
417                        Ok(version)
418                    }
419                    Err(e) => {
420                        println!("Worker {} got error: {:?}", worker_id, e);
421                        Err(e)
422                    }
423                }
424            });
425
426            handles.push(handle);
427        }
428
429        // Wait for all workers to complete
430        let mut results = Vec::new();
431        for handle in handles {
432            let result = handle.await.unwrap();
433            if let Ok(version) = result {
434                results.push(version);
435            }
436        }
437
438        // Verify results
439        let final_success_count = *success_count.lock().await;
440        let final_exists_count = *exists_count.lock().await;
441
442        println!(
443            "Final counts - Created: {}, Exists: {}",
444            final_success_count, final_exists_count
445        );
446
447        // CRITICAL ASSERTIONS:
448        // 1. Exactly ONE worker should have successfully created the key
449        assert_eq!(
450            final_success_count, 1,
451            "Exactly one worker should create the key"
452        );
453
454        // 2. All other workers should have gotten "Exists" response
455        assert_eq!(
456            final_exists_count,
457            num_workers - 1,
458            "All other workers should see key exists"
459        );
460
461        // 3. Total successful operations should equal number of workers
462        assert_eq!(
463            results.len(),
464            num_workers,
465            "All workers should complete successfully"
466        );
467
468        // 4. Verify the key actually exists in etcd
469        let stored_value = bucket.lock().await.get(&test_key).await?;
470        assert!(stored_value.is_some(), "Key should exist in etcd");
471
472        // 5. The stored value should be from one of the workers
473        let stored_str = String::from_utf8(stored_value.unwrap().to_vec()).unwrap();
474        assert!(
475            stored_str.starts_with(test_value),
476            "Stored value should match expected prefix"
477        );
478
479        // Clean up
480        bucket.lock().await.delete(&test_key).await?;
481
482        Ok(())
483    }
484}