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::{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 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 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 Ok(StoreOutcome::Created(1)) }
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 }
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 None => StoreOutcome::Created(1),
255 Some(kv) if kv.version() as u64 == version + 1 => StoreOutcome::Created(version),
257 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 let bucket = Arc::new(tokio::sync::Mutex::new(
296 storage
297 .get_or_create_bucket("test_concurrent_bucket", None)
298 .await?,
299 ));
300
301 let num_workers = 10;
303 let barrier = Arc::new(Barrier::new(num_workers));
304
305 let test_key: Key = Key::new(format!("concurrent_test_key_{}", uuid::Uuid::new_v4()));
307 let test_value = "test_value";
308
309 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 barrier_clone.wait().await;
325
326 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 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 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 assert_eq!(
383 final_success_count, 1,
384 "Exactly one worker should create the key"
385 );
386
387 assert_eq!(
389 final_exists_count,
390 num_workers - 1,
391 "All other workers should see key exists"
392 );
393
394 assert_eq!(
396 results.len(),
397 num_workers,
398 "All workers should complete successfully"
399 );
400
401 let stored_value = bucket.lock().await.get(&test_key).await?;
403 assert!(stored_value.is_some(), "Key should exist in etcd");
404
405 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 bucket.lock().await.delete(&test_key).await?;
414
415 Ok(())
416 }
417}