Skip to main content

risingwave_connector/sink/
dynamodb.rs

1// Copyright 2024 RisingWave Labs
2//
3// Licensed under the Apache License, Version 2.0 (the "License");
4// you may not use this file except in compliance with the License.
5// You may obtain a copy of the License at
6//
7//     http://www.apache.org/licenses/LICENSE-2.0
8//
9// Unless required by applicable law or agreed to in writing, software
10// distributed under the License is distributed on an "AS IS" BASIS,
11// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12// See the License for the specific language governing permissions and
13// limitations under the License.
14
15use std::collections::{BTreeMap, BTreeSet, HashMap};
16
17use anyhow::{Context, anyhow};
18use aws_sdk_dynamodb as dynamodb;
19use aws_sdk_dynamodb::client::Client;
20use aws_smithy_types::Blob;
21use dynamodb::types::{AttributeValue, KeySchemaElement, TableStatus, WriteRequest};
22use futures::TryFutureExt;
23use risingwave_common::array::{Op, RowRef, StreamChunk};
24use risingwave_common::catalog::Schema;
25use risingwave_common::row::Row as _;
26use risingwave_common::types::{DataType, ScalarRefImpl, ToText};
27use risingwave_common::util::iter_util::ZipEqDebug;
28use serde::Deserialize;
29use serde_with::{DisplayFromStr, serde_as};
30use with_options::WithOptions;
31use write_chunk_future::{DynamoDbPayloadWriter, WriteChunkFuture};
32
33use super::writer::{
34    AsyncTruncateLogSinkerOf, AsyncTruncateSinkWriter, AsyncTruncateSinkWriterExt,
35};
36use super::{Result, Sink, SinkError, SinkParam, SinkWriterParam};
37use crate::connector_common::AwsAuthProps;
38use crate::enforce_secret::EnforceSecret;
39use crate::error::ConnectorResult;
40use crate::sink::log_store::DeliveryFutureManagerAddFuture;
41
42pub const DYNAMO_DB_SINK: &str = "dynamodb";
43
44#[serde_as]
45#[derive(Deserialize, Debug, Clone, WithOptions)]
46pub struct DynamoDbConfig {
47    #[serde(rename = "table", alias = "dynamodb.table")]
48    pub table: String,
49
50    #[serde(rename = "dynamodb.max_batch_rows", default = "default_max_batch_rows")]
51    #[serde_as(as = "DisplayFromStr")]
52    #[deprecated]
53    pub max_batch_rows: usize,
54
55    #[serde(flatten)]
56    pub aws_auth_props: AwsAuthProps,
57
58    #[serde(
59        rename = "dynamodb.max_batch_item_nums",
60        default = "default_max_batch_item_nums"
61    )]
62    #[serde_as(as = "DisplayFromStr")]
63    pub max_batch_item_nums: usize,
64
65    #[serde(
66        rename = "dynamodb.max_future_send_nums",
67        default = "default_max_future_send_nums"
68    )]
69    #[serde_as(as = "DisplayFromStr")]
70    pub max_future_send_nums: usize,
71
72    #[serde(
73        rename = "dynamodb.batch_write_retry_times",
74        default = "default_batch_write_retry_times"
75    )]
76    #[serde_as(as = "DisplayFromStr")]
77    pub batch_write_retry_times: usize,
78
79    #[serde(
80        rename = "dynamodb.batch_write_retry_backoff_ms",
81        default = "default_batch_write_retry_backoff_ms"
82    )]
83    #[serde_as(as = "DisplayFromStr")]
84    pub batch_write_retry_backoff_ms: u64,
85
86    #[serde(flatten)]
87    pub unknown_fields: std::collections::HashMap<String, String>,
88}
89
90crate::impl_sink_unknown_fields!(DynamoDbConfig);
91
92impl EnforceSecret for DynamoDbConfig {
93    fn enforce_one(prop: &str) -> crate::error::ConnectorResult<()> {
94        AwsAuthProps::enforce_one(prop)
95    }
96}
97
98fn default_max_batch_item_nums() -> usize {
99    25
100}
101
102fn default_max_future_send_nums() -> usize {
103    256
104}
105
106fn default_batch_write_retry_times() -> usize {
107    3
108}
109
110fn default_batch_write_retry_backoff_ms() -> u64 {
111    100
112}
113
114fn default_max_batch_rows() -> usize {
115    1024
116}
117
118impl DynamoDbConfig {
119    pub async fn build_client(&self) -> ConnectorResult<Client> {
120        let config = &self.aws_auth_props;
121        let aws_config = config.build_config().await?;
122
123        Ok(Client::new(&aws_config))
124    }
125
126    fn from_btreemap(values: BTreeMap<String, String>) -> Result<Self> {
127        serde_json::from_value::<DynamoDbConfig>(serde_json::to_value(values).unwrap())
128            .map_err(|e| SinkError::Config(anyhow!(e)))
129    }
130}
131
132#[derive(Clone, Debug)]
133pub struct DynamoDbSink {
134    pub config: DynamoDbConfig,
135    schema: Schema,
136    pk_indices: Vec<usize>,
137}
138
139impl EnforceSecret for DynamoDbSink {
140    fn enforce_secret<'a>(
141        prop_iter: impl Iterator<Item = &'a str>,
142    ) -> crate::error::ConnectorResult<()> {
143        for prop in prop_iter {
144            DynamoDbConfig::enforce_one(prop)?;
145        }
146        Ok(())
147    }
148}
149
150impl Sink for DynamoDbSink {
151    type LogSinker = AsyncTruncateLogSinkerOf<DynamoDbSinkWriter>;
152
153    const SINK_NAME: &'static str = DYNAMO_DB_SINK;
154
155    crate::impl_validate_sink_unknown_fields!();
156
157    async fn validate(&self) -> Result<()> {
158        risingwave_common::license::Feature::DynamoDbSink
159            .check_available()
160            .map_err(|e| anyhow::anyhow!(e))?;
161        let client = (self.config.build_client().await)
162            .context("validate DynamoDB sink error")
163            .map_err(SinkError::DynamoDb)?;
164
165        let table_name = &self.config.table;
166        let output = client
167            .describe_table()
168            .table_name(table_name)
169            .send()
170            .await
171            .map_err(|e| anyhow!(e))?;
172        let Some(table) = output.table else {
173            return Err(SinkError::DynamoDb(anyhow!(
174                "table {} not found",
175                table_name
176            )));
177        };
178        if !matches!(table.table_status(), Some(TableStatus::Active)) {
179            return Err(SinkError::DynamoDb(anyhow!(
180                "table {} is not active",
181                table_name
182            )));
183        }
184        let rw_pk_names = rw_pk_names(&self.schema, &self.pk_indices)?;
185        let dynamodb_keys = dynamodb_key_schema_names(table_name, table.key_schema())?;
186        validate_pk_matches_dynamodb_key_schema(table_name, &rw_pk_names, &dynamodb_keys)?;
187
188        Ok(())
189    }
190
191    async fn new_log_sinker(&self, _writer_param: SinkWriterParam) -> Result<Self::LogSinker> {
192        Ok(
193            DynamoDbSinkWriter::new(self.config.clone(), self.schema.clone())
194                .await?
195                .into_log_sinker(self.config.max_future_send_nums),
196        )
197    }
198}
199
200impl TryFrom<SinkParam> for DynamoDbSink {
201    type Error = SinkError;
202
203    fn try_from(param: SinkParam) -> std::result::Result<Self, Self::Error> {
204        let schema = param.schema();
205        let pk_indices = param.downstream_pk_or_empty();
206        let config = DynamoDbConfig::from_btreemap(param.properties)?;
207
208        Ok(Self {
209            config,
210            schema,
211            pk_indices,
212        })
213    }
214}
215
216#[derive(Debug)]
217struct DynamoDbRequest {
218    inner: WriteRequest,
219    key_items: Vec<String>,
220}
221
222impl DynamoDbRequest {
223    fn extract_key(&self) -> Option<&HashMap<String, AttributeValue>> {
224        match (&self.inner.put_request(), &self.inner.delete_request()) {
225            (Some(put_req), None) => Some(&put_req.item),
226            (None, Some(del_req)) => Some(&del_req.key),
227            _ => None,
228        }
229    }
230
231    fn has_same_pk(&self, other: &Self) -> bool {
232        if self.key_items.is_empty() {
233            return false;
234        }
235
236        let Some(key) = self.extract_key() else {
237            return false;
238        };
239        let Some(other_key) = other.extract_key() else {
240            return false;
241        };
242
243        self.key_items.iter().all(|key_item| {
244            matches!(
245                (key.get(key_item), other_key.get(key_item)),
246                (Some(value), Some(other_value)) if value == other_value
247            )
248        })
249    }
250}
251
252pub struct DynamoDbSinkWriter {
253    payload_writer: DynamoDbPayloadWriter,
254    formatter: DynamoDbFormatter,
255    max_future_send_nums: usize,
256}
257
258impl DynamoDbSinkWriter {
259    pub async fn new(config: DynamoDbConfig, schema: Schema) -> Result<Self> {
260        let client = config.build_client().await?;
261        let table_name = &config.table;
262        let output = client
263            .describe_table()
264            .table_name(table_name)
265            .send()
266            .await
267            .map_err(|e| anyhow!(e))?;
268        let Some(table) = output.table else {
269            return Err(SinkError::DynamoDb(anyhow!(
270                "table {} not found",
271                table_name
272            )));
273        };
274        let dynamodb_keys = dynamodb_key_schema_names(table_name, table.key_schema())?;
275
276        let payload_writer = DynamoDbPayloadWriter {
277            client,
278            table: config.table.clone(),
279            dynamodb_keys,
280            max_batch_item_nums: config.max_batch_item_nums,
281            batch_write_retry_times: config.batch_write_retry_times,
282            batch_write_retry_backoff_ms: config.batch_write_retry_backoff_ms,
283        };
284
285        Ok(Self {
286            payload_writer,
287            formatter: DynamoDbFormatter { schema },
288            max_future_send_nums: config.max_future_send_nums,
289        })
290    }
291
292    fn write_chunk_inner(&mut self, chunk: StreamChunk) -> Result<WriteChunkFuture> {
293        let mut request_items = Vec::new();
294        for (op, row) in chunk.rows() {
295            let items = self.formatter.format_row(row)?;
296            match op {
297                Op::Insert | Op::UpdateInsert => {
298                    self.payload_writer
299                        .write_one_insert(items, &mut request_items);
300                }
301                Op::Delete => {
302                    self.payload_writer
303                        .write_one_delete(items, &mut request_items);
304                }
305                Op::UpdateDelete => {}
306            }
307        }
308        Ok(self
309            .payload_writer
310            .write_chunk(request_items, self.max_future_send_nums))
311    }
312}
313
314impl AsyncTruncateSinkWriter for DynamoDbSinkWriter {
315    type DeliveryFuture = WriteChunkFuture;
316
317    async fn write_chunk<'a>(
318        &'a mut self,
319        chunk: StreamChunk,
320        _add_future: DeliveryFutureManagerAddFuture<'a, Self::DeliveryFuture>,
321    ) -> Result<()> {
322        self.write_chunk_inner(chunk)?.map_ok(|_| ()).await?;
323        Ok(())
324    }
325}
326
327struct DynamoDbFormatter {
328    schema: Schema,
329}
330
331impl DynamoDbFormatter {
332    fn format_row(&self, row: RowRef<'_>) -> Result<HashMap<String, AttributeValue>> {
333        row.iter()
334            .zip_eq_debug((self.schema.clone()).into_fields())
335            .map(|(scalar, field)| {
336                map_data(scalar, &field.data_type()).map(|attr| (field.name, attr))
337            })
338            .collect()
339    }
340}
341
342fn map_data(scalar_ref: Option<ScalarRefImpl<'_>>, data_type: &DataType) -> Result<AttributeValue> {
343    let Some(scalar_ref) = scalar_ref else {
344        return Ok(AttributeValue::Null(true));
345    };
346    let attr = match data_type {
347        DataType::Int16
348        | DataType::Int32
349        | DataType::Int64
350        | DataType::Int256
351        | DataType::Float32
352        | DataType::Float64
353        | DataType::Decimal
354        | DataType::Serial => AttributeValue::N(scalar_ref.to_text_with_type(data_type)),
355        // TODO: jsonb as dynamic type (https://github.com/risingwavelabs/risingwave/issues/11699)
356        DataType::Varchar
357        | DataType::Interval
358        | DataType::Date
359        | DataType::Time
360        | DataType::Timestamp
361        | DataType::Timestamptz
362        | DataType::Jsonb => AttributeValue::S(scalar_ref.to_text_with_type(data_type)),
363        DataType::Variant => {
364            return Err(SinkError::DynamoDb(anyhow!("variant is not supported yet")));
365        }
366        DataType::Boolean => AttributeValue::Bool(scalar_ref.into_bool()),
367        DataType::Bytea => AttributeValue::B(Blob::new(scalar_ref.into_bytea())),
368        DataType::List(lt) => {
369            let list_attr = scalar_ref
370                .into_list()
371                .iter()
372                .map(|x| map_data(x, lt.elem()))
373                .collect::<Result<Vec<_>>>()?;
374            AttributeValue::L(list_attr)
375        }
376        DataType::Struct(st) => {
377            let mut map = HashMap::with_capacity(st.len());
378            for (sub_datum_ref, (name, data_type)) in scalar_ref
379                .into_struct()
380                .iter_fields_ref()
381                .zip_eq_debug(st.iter())
382            {
383                let attr = map_data(sub_datum_ref, data_type)?;
384                map.insert(name.to_owned(), attr);
385            }
386            AttributeValue::M(map)
387        }
388        DataType::Map(_m) => {
389            return Err(SinkError::DynamoDb(anyhow!("map is not supported yet")));
390        }
391        DataType::Vector(_) => {
392            return Err(SinkError::DynamoDb(anyhow!("vector is not supported yet")));
393        }
394    };
395    Ok(attr)
396}
397
398fn rw_pk_names(schema: &Schema, pk_indices: &[usize]) -> Result<Vec<String>> {
399    pk_indices
400        .iter()
401        .map(|pk_idx| {
402            schema
403                .fields()
404                .get(*pk_idx)
405                .map(|field| field.name.clone())
406                .ok_or_else(|| {
407                    SinkError::DynamoDb(anyhow!(
408                        "RisingWave primary key column index {} is out of range",
409                        pk_idx
410                    ))
411                })
412        })
413        .collect()
414}
415
416fn dynamodb_key_schema_names(
417    table_name: &str,
418    key_schema: &[KeySchemaElement],
419) -> Result<Vec<String>> {
420    if key_schema.is_empty() {
421        return Err(SinkError::DynamoDb(anyhow!(
422            "table {} key schema is empty",
423            table_name
424        )));
425    }
426
427    Ok(key_schema
428        .iter()
429        .map(|key_element| key_element.attribute_name().to_owned())
430        .collect())
431}
432
433fn validate_pk_matches_dynamodb_key_schema(
434    table_name: &str,
435    rw_pk_names: &[String],
436    dynamodb_keys: &[String],
437) -> Result<()> {
438    let rw_pk_set = rw_pk_names.iter().collect::<BTreeSet<_>>();
439    let dynamodb_key_set = dynamodb_keys.iter().collect::<BTreeSet<_>>();
440    if rw_pk_names.len() != dynamodb_keys.len() || rw_pk_set != dynamodb_key_set {
441        return Err(SinkError::DynamoDb(anyhow!(
442            "DynamoDB table {} primary key {:?} must match RisingWave primary key {:?}",
443            table_name,
444            dynamodb_keys,
445            rw_pk_names
446        )));
447    }
448
449    Ok(())
450}
451
452mod write_chunk_future {
453    use std::collections::HashMap;
454    use std::time::Duration;
455
456    use anyhow::anyhow;
457    use aws_sdk_dynamodb as dynamodb;
458    use aws_sdk_dynamodb::client::Client;
459    use dynamodb::types::{
460        AttributeValue, DeleteRequest, PutRequest, ReturnConsumedCapacity,
461        ReturnItemCollectionMetrics, WriteRequest,
462    };
463    use futures::{FutureExt, StreamExt, TryFuture, TryStreamExt, stream};
464    use itertools::Itertools;
465    use maplit::hashmap;
466    use risingwave_common::util::retry::exponential_backoff;
467    use tokio::time::sleep;
468    use tokio_retry::strategy::jitter;
469
470    use super::{DynamoDbRequest, SinkError};
471
472    const MAX_BATCH_WRITE_RETRY_DELAY_MS: u64 = 2000;
473    const MAX_BATCH_WRITE_CONCURRENCY: usize = 256;
474
475    pub struct DynamoDbPayloadWriter {
476        pub client: Client,
477        pub table: String,
478        pub dynamodb_keys: Vec<String>,
479        pub max_batch_item_nums: usize,
480        pub batch_write_retry_times: usize,
481        pub batch_write_retry_backoff_ms: u64,
482    }
483
484    pub type WriteChunkFuture = impl TryFuture<Ok = (), Error = SinkError> + Unpin + Send + 'static;
485
486    impl DynamoDbPayloadWriter {
487        pub fn write_one_insert(
488            &mut self,
489            item: HashMap<String, AttributeValue>,
490            request_items: &mut Vec<DynamoDbRequest>,
491        ) {
492            let put_req = PutRequest::builder().set_item(Some(item)).build().unwrap();
493            let req = WriteRequest::builder().put_request(put_req).build();
494            self.write_one_req(req, request_items);
495        }
496
497        pub fn write_one_delete(
498            &mut self,
499            key: HashMap<String, AttributeValue>,
500            request_items: &mut Vec<DynamoDbRequest>,
501        ) {
502            let key = key
503                .into_iter()
504                .filter(|(k, _)| self.dynamodb_keys.contains(k))
505                .collect();
506            let del_req = DeleteRequest::builder().set_key(Some(key)).build().unwrap();
507            let req = WriteRequest::builder().delete_request(del_req).build();
508            self.write_one_req(req, request_items);
509        }
510
511        pub fn write_one_req(
512            &mut self,
513            req: WriteRequest,
514            request_items: &mut Vec<DynamoDbRequest>,
515        ) {
516            let r_req = DynamoDbRequest {
517                inner: req,
518                key_items: self.dynamodb_keys.clone(),
519            };
520            request_items.retain(|item| !item.has_same_pk(&r_req));
521            request_items.push(r_req);
522        }
523
524        #[define_opaque(WriteChunkFuture)]
525        pub fn write_chunk(
526            &mut self,
527            request_items: Vec<DynamoDbRequest>,
528            max_future_send_nums: usize,
529        ) -> WriteChunkFuture {
530            let client = self.client.clone();
531            let table = self.table.clone();
532            let max_batch_item_nums = self.max_batch_item_nums;
533            let batch_write_retry_times = self.batch_write_retry_times;
534            let batch_write_retry_backoff_ms = self.batch_write_retry_backoff_ms;
535            async move {
536                let chunks = request_items
537                    .into_iter()
538                    .map(|r| r.inner)
539                    .chunks(max_batch_item_nums)
540                    .into_iter()
541                    .map(|chunk| chunk.collect::<Vec<_>>())
542                    .collect_vec();
543                let max_future_send_nums =
544                    max_future_send_nums.clamp(1, MAX_BATCH_WRITE_CONCURRENCY);
545                stream::iter(chunks.into_iter().map(|req_items| {
546                    let client = client.clone();
547                    let table = table.clone();
548                    async move {
549                        let mut req_items = req_items;
550                        let mut retry_count = 0;
551                        let mut retry_backoff = exponential_backoff(
552                            Duration::from_millis(batch_write_retry_backoff_ms),
553                            2,
554                            Duration::from_millis(MAX_BATCH_WRITE_RETRY_DELAY_MS),
555                        )
556                        .map(jitter)
557                        .take(batch_write_retry_times);
558
559                        loop {
560                            let return_consumed_capacity = if retry_count == 0 {
561                                ReturnConsumedCapacity::None
562                            } else {
563                                ReturnConsumedCapacity::Total
564                            };
565                            let reqs = hashmap! {
566                                table.clone() => req_items.clone(),
567                            };
568                            let result = client
569                                .batch_write_item()
570                                .set_request_items(Some(reqs))
571                                .return_consumed_capacity(return_consumed_capacity)
572                                .return_item_collection_metrics(ReturnItemCollectionMetrics::None)
573                                .send()
574                                .await;
575
576                            match result {
577                                Ok(output) => {
578                                    let unprocessed_items =
579                                        output.unprocessed_items().cloned().unwrap_or_default();
580                                    if unprocessed_items.is_empty() {
581                                        if retry_count > 0 {
582                                            tracing::warn!(
583                                                retry_count,
584                                                consumed_capacity = ?output.consumed_capacity(),
585                                                "DynamoDB batch write retry succeeded"
586                                            );
587                                        }
588                                        return Ok(());
589                                    }
590
591                                    req_items = unprocessed_items.into_values().flatten().collect();
592                                    if retry_count >= batch_write_retry_times {
593                                        return Err(SinkError::DynamoDb(anyhow!(
594                                            "failed to write {} unprocessed items to DynamoDB sink after {} retries",
595                                            req_items.len(),
596                                            batch_write_retry_times,
597                                        )));
598                                    }
599                                }
600                                Err(e) => {
601                                    return Err(SinkError::DynamoDb(
602                                        anyhow!(e).context("failed to write items to DynamoDB sink"),
603                                    ));
604                                }
605                            }
606
607                            retry_count += 1;
608                            let Some(delay) = retry_backoff.next() else {
609                                return Err(SinkError::DynamoDb(anyhow!(
610                                    "failed to write {} unprocessed items to DynamoDB sink after {} retries",
611                                    req_items.len(),
612                                    batch_write_retry_times,
613                                )));
614                            };
615                            tracing::warn!(
616                                retry_count,
617                                delay_ms = delay.as_millis(),
618                                unprocessed_items_count = req_items.len(),
619                                "retrying DynamoDB batch write"
620                            );
621                            sleep(delay).await;
622                        }
623                    }
624                }))
625                .buffer_unordered(max_future_send_nums)
626                .try_collect::<Vec<_>>()
627                .await?;
628                Ok(())
629            }
630            .boxed()
631        }
632    }
633}
634
635#[cfg(test)]
636mod tests {
637    use aws_sdk_dynamodb::types::{DeleteRequest, KeyType, PutRequest};
638
639    use super::*;
640
641    fn dynamodb_put_request(
642        items: impl IntoIterator<Item = (&'static str, &'static str)>,
643    ) -> DynamoDbRequest {
644        let item = dynamodb_items(items);
645        let put_req = PutRequest::builder().set_item(Some(item)).build().unwrap();
646        DynamoDbRequest {
647            inner: WriteRequest::builder().put_request(put_req).build(),
648            key_items: vec!["pk".to_owned(), "sk".to_owned()],
649        }
650    }
651
652    fn dynamodb_delete_request(
653        items: impl IntoIterator<Item = (&'static str, &'static str)>,
654    ) -> DynamoDbRequest {
655        let key = dynamodb_items(items);
656        let delete_req = DeleteRequest::builder().set_key(Some(key)).build().unwrap();
657        DynamoDbRequest {
658            inner: WriteRequest::builder().delete_request(delete_req).build(),
659            key_items: vec!["pk".to_owned(), "sk".to_owned()],
660        }
661    }
662
663    fn dynamodb_items(
664        items: impl IntoIterator<Item = (&'static str, &'static str)>,
665    ) -> HashMap<String, AttributeValue> {
666        items
667            .into_iter()
668            .map(|(k, v)| (k.to_owned(), AttributeValue::S(v.to_owned())))
669            .collect()
670    }
671
672    #[test]
673    fn dynamodb_request_compares_pk_by_key_attribute() {
674        let req = dynamodb_put_request([("pk", "a"), ("sk", "b")]);
675        let swapped_values = dynamodb_put_request([("pk", "b"), ("sk", "a")]);
676        let same_pk = dynamodb_put_request([("pk", "a"), ("sk", "b")]);
677
678        assert!(!req.has_same_pk(&swapped_values));
679        assert!(req.has_same_pk(&same_pk));
680    }
681
682    #[test]
683    fn dynamodb_request_empty_key_items_never_match() {
684        let mut req = dynamodb_put_request([("pk", "a"), ("sk", "b")]);
685        let same_pk = dynamodb_put_request([("pk", "a"), ("sk", "b")]);
686        req.key_items.clear();
687
688        assert!(!req.has_same_pk(&same_pk));
689    }
690
691    #[test]
692    fn dynamodb_request_compares_put_and_delete_by_composite_pk() {
693        let put = dynamodb_put_request([("pk", "a"), ("sk", "b"), ("value", "1")]);
694        let same_pk_delete = dynamodb_delete_request([("pk", "a"), ("sk", "b")]);
695        let different_hash_key_delete = dynamodb_delete_request([("pk", "x"), ("sk", "b")]);
696        let different_range_key_delete = dynamodb_delete_request([("pk", "a"), ("sk", "x")]);
697
698        assert!(put.has_same_pk(&same_pk_delete));
699        assert!(same_pk_delete.has_same_pk(&put));
700        assert!(!put.has_same_pk(&different_hash_key_delete));
701        assert!(!put.has_same_pk(&different_range_key_delete));
702    }
703
704    #[test]
705    fn dynamodb_key_schema_empty_errors() {
706        let err = dynamodb_key_schema_names("test_table", &[]).unwrap_err();
707
708        assert!(
709            err.to_string()
710                .contains("table test_table key schema is empty")
711        );
712    }
713
714    #[test]
715    fn dynamodb_key_schema_must_match_rw_pk() {
716        let dynamodb_keys = ["pk".to_owned(), "sk".to_owned()];
717        let same_rw_pk = ["sk".to_owned(), "pk".to_owned()];
718        let extra_rw_pk = ["pk".to_owned(), "sk".to_owned(), "extra".to_owned()];
719        let different_rw_pk = ["pk".to_owned(), "other".to_owned()];
720
721        validate_pk_matches_dynamodb_key_schema("test_table", &same_rw_pk, &dynamodb_keys).unwrap();
722
723        assert!(
724            validate_pk_matches_dynamodb_key_schema("test_table", &extra_rw_pk, &dynamodb_keys)
725                .unwrap_err()
726                .to_string()
727                .contains("must match RisingWave primary key")
728        );
729        assert!(
730            validate_pk_matches_dynamodb_key_schema("test_table", &different_rw_pk, &dynamodb_keys)
731                .unwrap_err()
732                .to_string()
733                .contains("must match RisingWave primary key")
734        );
735    }
736
737    #[test]
738    fn dynamodb_key_schema_names_uses_explicit_schema() {
739        let key_schema = vec![
740            KeySchemaElement::builder()
741                .attribute_name("pk")
742                .key_type(KeyType::Hash)
743                .build()
744                .unwrap(),
745            KeySchemaElement::builder()
746                .attribute_name("sk")
747                .key_type(KeyType::Range)
748                .build()
749                .unwrap(),
750        ];
751
752        assert_eq!(
753            dynamodb_key_schema_names("test_table", &key_schema).unwrap(),
754            vec!["pk".to_owned(), "sk".to_owned()]
755        );
756    }
757}