dynamo_runtime/storage/kv/
etcd.rs1use 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 async fn get_or_create_bucket(
32 &self,
33 bucket_name: &str,
34 _ttl: Option<Duration>, ) -> Result<Self::Bucket, StoreError> {
36 Ok(EtcdBucket {
37 client: self.client.clone(),
38 bucket_name: bucket_name.to_string(),
39 })
40 }
41
42 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 }
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 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 Ok(StoreOutcome::Created(1)) }
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 }
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 None => StoreOutcome::Created(1),
274 Some(kv) if kv.version() as u64 == version + 1 => StoreOutcome::Created(version),
276 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 let bucket = Arc::new(tokio::sync::Mutex::new(
363 storage
364 .get_or_create_bucket("test_concurrent_bucket", None)
365 .await?,
366 ));
367
368 let num_workers = 10;
370 let barrier = Arc::new(Barrier::new(num_workers));
371
372 let test_key: Key = Key::new(format!("concurrent_test_key_{}", uuid::Uuid::new_v4()));
374 let test_value = "test_value";
375
376 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 barrier_clone.wait().await;
392
393 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 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 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 assert_eq!(
450 final_success_count, 1,
451 "Exactly one worker should create the key"
452 );
453
454 assert_eq!(
456 final_exists_count,
457 num_workers - 1,
458 "All other workers should see key exists"
459 );
460
461 assert_eq!(
463 results.len(),
464 num_workers,
465 "All workers should complete successfully"
466 );
467
468 let stored_value = bucket.lock().await.get(&test_key).await?;
470 assert!(stored_value.is_some(), "Key should exist in etcd");
471
472 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 bucket.lock().await.delete(&test_key).await?;
481
482 Ok(())
483 }
484}