Skip to main content

risingwave_connector/parser/
mysql.rs

1// Copyright 2023 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::sync::LazyLock;
16
17use mysql_async::Row as MysqlRow;
18use mysql_common::constants::ColumnFlags;
19use risingwave_common::catalog::Schema;
20use risingwave_common::log::LogSuppressor;
21use risingwave_common::row::OwnedRow;
22use thiserror_ext::AsReport;
23
24use crate::parser::utils::log_error;
25
26static LOG_SUPPRESSOR: LazyLock<LogSuppressor> = LazyLock::new(LogSuppressor::default);
27use anyhow::anyhow;
28use chrono::NaiveDate;
29use risingwave_common::bail;
30use risingwave_common::types::{
31    DataType, Date, Datum, Decimal, JsonbVal, ScalarImpl, Time, Timestamp, Timestamptz,
32};
33use rust_decimal::Decimal as RustDecimal;
34
35macro_rules! handle_data_type {
36    ($row:expr, $i:expr, $name:expr, $typ:ty) => {{
37        match $row.take_opt::<Option<$typ>, _>($i) {
38            None => bail!("no value found at column: {}, index: {}", $name, $i),
39            Some(Ok(val)) => Ok(val.map(|v| ScalarImpl::from(v))),
40            Some(Err(e)) => Err(anyhow::Error::new(e.clone())
41                .context("failed to deserialize MySQL value into rust value")
42                .context(format!(
43                    "column: {}, index: {}, rust_type: {}",
44                    $name,
45                    $i,
46                    stringify!($typ),
47                ))),
48        }
49    }};
50    ($row:expr, $i:expr, $name:expr, $typ:ty, $rw_type:ty) => {{
51        match $row.take_opt::<Option<$typ>, _>($i) {
52            None => bail!("no value found at column: {}, index: {}", $name, $i),
53            Some(Ok(val)) => Ok(val.map(|v| ScalarImpl::from(<$rw_type>::from(v)))),
54            Some(Err(e)) => Err(anyhow::Error::new(e.clone())
55                .context("failed to deserialize MySQL value into rw value")
56                .context(format!(
57                    "column: {}, index: {}, rw_type: {}",
58                    $name,
59                    $i,
60                    stringify!($rw_type),
61                ))),
62        }
63    }};
64}
65
66macro_rules! handle_data_type_with_signed {
67    (
68        $mysql_row:expr,
69        $mysql_datum_index:expr,
70        $column_name:expr,
71        $signed_type:ty,
72        $unsigned_type:ty
73    ) => {{
74        let column_flags = $mysql_row.columns()[$mysql_datum_index].flags();
75
76        if column_flags.contains(ColumnFlags::UNSIGNED_FLAG) {
77            // UNSIGNED type: use unsigned type conversion, then convert to signed
78            match $mysql_row.take_opt::<Option<$unsigned_type>, _>($mysql_datum_index) {
79                // Note: We are intentionally converting unsigned to signed here.
80                // For example, 18446251075179777772u64 will be converted to -492998529773844i64.
81                Some(Ok(Some(val))) => Ok(Some(ScalarImpl::from(val as $signed_type))),
82                Some(Ok(None)) => Ok(None),
83                Some(Err(e)) => Err(anyhow::Error::new(e.clone())
84                    .context("failed to deserialize MySQL value into rust value")
85                    .context(format!(
86                        "column: {}, index: {}, rust_type: {}",
87                        $column_name,
88                        $mysql_datum_index,
89                        stringify!($unsigned_type),
90                    ))),
91                None => bail!(
92                    "no value found at column: {}, index: {}",
93                    $column_name,
94                    $mysql_datum_index
95                ),
96            }
97        } else {
98            // SIGNED type: use default signed type conversion
99            handle_data_type!($mysql_row, $mysql_datum_index, $column_name, $signed_type)
100        }
101    }};
102}
103
104/// The decoding result can be interpreted as follows:
105/// Ok(value) => The value was found and successfully decoded.
106/// Err(error) => The value was found but could not be decoded,
107///               either because it was not supported,
108///               or there was an error during conversion.
109pub fn mysql_datum_to_rw_datum(
110    mysql_row: &mut MysqlRow,
111    mysql_datum_index: usize,
112    column_name: &str,
113    rw_data_type: &DataType,
114) -> Result<Datum, anyhow::Error> {
115    match rw_data_type {
116        DataType::Boolean => {
117            // TinyInt(1) is used to represent boolean in MySQL
118            // This handles backwards compatibility,
119            // before https://github.com/risingwavelabs/risingwave/pull/19071
120            // we permit boolean and tinyint(1) to be equivalent to boolean in RW.
121            if let Some(Ok(val)) = mysql_row.get_opt::<Option<bool>, _>(mysql_datum_index) {
122                return Ok(val.map(ScalarImpl::from));
123            }
124            // Bit(1)
125            match mysql_row.take_opt::<Option<Vec<u8>>, _>(mysql_datum_index) {
126                None => bail!(
127                    "no value found at column: {}, index: {}",
128                    column_name,
129                    mysql_datum_index
130                ),
131                Some(Ok(val)) => match val {
132                    None => Ok(None),
133                    Some(val) => match val.as_slice() {
134                        [0] => Ok(Some(ScalarImpl::from(false))),
135                        [1] => Ok(Some(ScalarImpl::from(true))),
136                        _ => Err(anyhow!("invalid value for boolean: {:?}", val)),
137                    },
138                },
139                Some(Err(e)) => Err(anyhow::Error::new(e)
140                    .context("failed to deserialize MySQL value into rust value")
141                    .context(format!(
142                        "column: {}, index: {}, rust_type: Vec<u8>",
143                        column_name, mysql_datum_index,
144                    ))),
145            }
146        }
147        DataType::Int16 => {
148            handle_data_type!(mysql_row, mysql_datum_index, column_name, i16)
149        }
150        DataType::Int32 => {
151            handle_data_type!(mysql_row, mysql_datum_index, column_name, i32)
152        }
153        DataType::Int64 => {
154            handle_data_type_with_signed!(mysql_row, mysql_datum_index, column_name, i64, u64)
155        }
156        DataType::Float32 => {
157            handle_data_type!(mysql_row, mysql_datum_index, column_name, f32)
158        }
159        DataType::Float64 => {
160            handle_data_type!(mysql_row, mysql_datum_index, column_name, f64)
161        }
162        DataType::Decimal => {
163            handle_data_type!(
164                mysql_row,
165                mysql_datum_index,
166                column_name,
167                RustDecimal,
168                Decimal
169            )
170        }
171        DataType::Varchar => {
172            handle_data_type!(mysql_row, mysql_datum_index, column_name, String)
173        }
174        DataType::Date => {
175            handle_data_type!(mysql_row, mysql_datum_index, column_name, NaiveDate, Date)
176        }
177        DataType::Time => {
178            handle_data_type!(
179                mysql_row,
180                mysql_datum_index,
181                column_name,
182                chrono::NaiveTime,
183                Time
184            )
185        }
186        DataType::Timestamp => {
187            handle_data_type!(
188                mysql_row,
189                mysql_datum_index,
190                column_name,
191                chrono::NaiveDateTime,
192                Timestamp
193            )
194        }
195        DataType::Timestamptz => {
196            match mysql_row.take_opt::<Option<chrono::NaiveDateTime>, _>(mysql_datum_index) {
197                None => bail!(
198                    "no value found at column: {}, index: {}",
199                    column_name,
200                    mysql_datum_index
201                ),
202                Some(Ok(val)) => Ok(val.map(|v| {
203                    ScalarImpl::from(Timestamptz::from_micros_uncheck(
204                        v.and_utc().timestamp_micros(),
205                    ))
206                })),
207                Some(Err(err)) => Err(anyhow::Error::new(err)
208                    .context("failed to deserialize MySQL value into rust value")
209                    .context(format!(
210                        "column: {}, index: {}, rust_type: chrono::NaiveDateTime",
211                        column_name, mysql_datum_index,
212                    ))),
213            }
214        }
215        DataType::Bytea => match mysql_row.take_opt::<Option<Vec<u8>>, _>(mysql_datum_index) {
216            None => bail!(
217                "no value found at column: {}, index: {}",
218                column_name,
219                mysql_datum_index
220            ),
221            Some(Ok(val)) => Ok(val.map(ScalarImpl::from)),
222            Some(Err(err)) => Err(anyhow::Error::new(err)
223                .context("failed to deserialize MySQL value into rust value")
224                .context(format!(
225                    "column: {}, index: {}, rust_type: Vec<u8>",
226                    column_name, mysql_datum_index,
227                ))),
228        },
229        DataType::Jsonb => {
230            handle_data_type!(
231                mysql_row,
232                mysql_datum_index,
233                column_name,
234                serde_json::Value,
235                JsonbVal
236            )
237        }
238        DataType::Variant
239        | DataType::Vector(_)
240        | DataType::Interval
241        | DataType::Struct(_)
242        | DataType::List(_)
243        | DataType::Int256
244        | DataType::Serial
245        | DataType::Map(_) => Err(anyhow!(
246            "unsupported data type: {}, set to null",
247            rw_data_type
248        )),
249    }
250}
251
252pub fn mysql_row_to_owned_row(mysql_row: &mut MysqlRow, schema: &Schema) -> OwnedRow {
253    let mut datums = vec![];
254    for i in 0..schema.fields.len() {
255        let rw_field = &schema.fields[i];
256        let name = rw_field.name.as_str();
257        let datum = match mysql_datum_to_rw_datum(mysql_row, i, name, &rw_field.data_type) {
258            Ok(val) => val,
259            Err(e) => {
260                log_error!(name, e, "parse column failed");
261                None
262            }
263        };
264        datums.push(datum);
265    }
266    OwnedRow::new(datums)
267}
268
269/// Decode primary-key columns strictly while preserving the legacy lenient behavior for all
270/// other columns in a MySQL CDC snapshot row.
271pub fn mysql_row_to_owned_row_with_strict_pk(
272    mysql_row: &mut MysqlRow,
273    schema: &Schema,
274    pk_indices: &[usize],
275) -> anyhow::Result<OwnedRow> {
276    super::decode_row_with_strict_pk(
277        "MySQL",
278        schema,
279        pk_indices,
280        |index, field| mysql_datum_to_rw_datum(mysql_row, index, &field.name, &field.data_type),
281        |name, err| log_error!(name, err, "parse column failed"),
282    )
283}
284
285#[cfg(test)]
286mod tests {
287
288    use std::sync::Arc;
289
290    use futures::pin_mut;
291    use mysql_async::Row as MySqlRow;
292    use mysql_async::prelude::*;
293    use mysql_common::constants::ColumnType;
294    use mysql_common::packets::Column;
295    use mysql_common::row::new_row;
296    use mysql_common::value::Value;
297    use risingwave_common::catalog::{Field, Schema};
298    use risingwave_common::row::Row;
299    use risingwave_common::types::DataType;
300    use tokio_stream::StreamExt;
301
302    use crate::parser::{mysql_row_to_owned_row, mysql_row_to_owned_row_with_strict_pk};
303
304    fn mysql_row(values: Vec<Value>, names: &[&str]) -> MySqlRow {
305        let columns = names
306            .iter()
307            .map(|name| Column::new(ColumnType::MYSQL_TYPE_VAR_STRING).with_name(name.as_bytes()))
308            .collect::<Vec<_>>();
309        new_row(values, Arc::from(columns))
310    }
311
312    #[test]
313    fn strict_pk_decode_rejects_malformed_value() {
314        let schema = Schema::new(vec![
315            Field::with_name(DataType::Int32, "id"),
316            Field::with_name(DataType::Varchar, "payload"),
317        ]);
318        let mut row = mysql_row(
319            vec![
320                Value::Bytes(b"not-an-int".to_vec()),
321                Value::Bytes(b"ok".to_vec()),
322            ],
323            &["id", "payload"],
324        );
325
326        let err = mysql_row_to_owned_row_with_strict_pk(&mut row, &schema, &[0]).unwrap_err();
327        assert!(err.to_string().contains("primary key `id`"));
328    }
329
330    #[test]
331    fn strict_pk_decode_rejects_null_value() {
332        let schema = Schema::new(vec![
333            Field::with_name(DataType::Int32, "id"),
334            Field::with_name(DataType::Varchar, "payload"),
335        ]);
336        let mut row = mysql_row(
337            vec![Value::NULL, Value::Bytes(b"ok".to_vec())],
338            &["id", "payload"],
339        );
340
341        let err = mysql_row_to_owned_row_with_strict_pk(&mut row, &schema, &[0]).unwrap_err();
342        assert!(err.to_string().contains("primary key `id` cannot be NULL"));
343    }
344
345    #[test]
346    fn strict_pk_decode_keeps_non_pk_conversion_lenient() {
347        let schema = Schema::new(vec![
348            Field::with_name(DataType::Int32, "id"),
349            Field::with_name(DataType::Int32, "payload"),
350        ]);
351        let mut row = mysql_row(
352            vec![Value::Int(1), Value::Bytes(b"not-an-int".to_vec())],
353            &["id", "payload"],
354        );
355
356        let row = mysql_row_to_owned_row_with_strict_pk(&mut row, &schema, &[0]).unwrap();
357        assert!(row.datum_at(0).is_some());
358        assert!(row.datum_at(1).is_none());
359    }
360
361    // manual test case
362    #[ignore]
363    #[tokio::test]
364    async fn test_convert_mysql_row_to_owned_row() {
365        let pool = mysql_async::Pool::new("mysql://root:123456@localhost:8306/mydb");
366
367        let t1schema = Schema::new(vec![
368            Field::with_name(DataType::Int32, "v1"),
369            Field::with_name(DataType::Int32, "v2"),
370            Field::with_name(DataType::Timestamptz, "v3"),
371        ]);
372
373        let mut conn = pool.get_conn().await.unwrap();
374        conn.exec_drop("SET time_zone = \"+08:00\"", ())
375            .await
376            .unwrap();
377
378        let mut result_set = conn.query_iter("SELECT * FROM `t1m`").await.unwrap();
379        let s = result_set.stream::<MySqlRow>().await.unwrap().unwrap();
380        let row_stream = s.map(|row| {
381            // convert mysql row into OwnedRow
382            let mut mysql_row = row.unwrap();
383            Ok::<_, anyhow::Error>(Some(mysql_row_to_owned_row(&mut mysql_row, &t1schema)))
384        });
385        pin_mut!(row_stream);
386        while let Some(row) = row_stream.next().await {
387            if let Ok(ro) = row
388                && ro.is_some()
389            {
390                let owned_row = ro.unwrap();
391                let d = owned_row.datum_at(2);
392                if let Some(scalar) = d {
393                    let v = scalar.into_timestamptz();
394                    println!("timestamp: {:?}", v);
395                }
396            }
397        }
398    }
399}