Skip to main content

risingwave_connector/source/cdc/external/
mysql.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::HashMap;
16
17use anyhow::{Context, anyhow};
18use chrono::{DateTime, NaiveDateTime};
19use futures::stream::BoxStream;
20use futures::{StreamExt, pin_mut, stream};
21use futures_async_stream::try_stream;
22use itertools::Itertools;
23use mysql_async::prelude::*;
24use mysql_common::params::Params;
25use mysql_common::value::Value;
26use risingwave_common::bail;
27use risingwave_common::catalog::{CDC_OFFSET_COLUMN_NAME, ColumnDesc, ColumnId, Field, Schema};
28use risingwave_common::row::OwnedRow;
29use risingwave_common::types::{DataType, Datum, Decimal, F32, ScalarImpl};
30use risingwave_common::util::iter_util::ZipEqFast;
31use sea_schema::mysql::def::{ColumnDefault, ColumnKey, ColumnType, NumericAttr};
32use sea_schema::mysql::discovery::SchemaDiscovery;
33use sea_schema::mysql::query::SchemaQueryBuilder;
34use sea_schema::sea_query::{Alias, IntoIden};
35use serde::{Deserialize, Serialize};
36use sqlx::MySqlPool;
37use sqlx::mysql::MySqlConnectOptions;
38use thiserror_ext::AsReport;
39
40use crate::connector_common::SslMode;
41// Re-export SslMode for convenience
42pub use crate::connector_common::SslMode as MySqlSslMode;
43use crate::error::{ConnectorError, ConnectorResult};
44use crate::parser::mysql_row_to_owned_row_with_strict_pk;
45use crate::source::CdcTableSnapshotSplit;
46use crate::source::cdc::external::{
47    CdcOffset, CdcOffsetParseFunc, CdcTableSnapshotSplitOption, DebeziumOffset,
48    ExternalTableConfig, ExternalTableReader, SchemaTableName,
49};
50
51/// Build MySQL connection pool with proper SSL configuration.
52///
53/// This helper function creates a `mysql_async::Pool` with all necessary configurations
54/// including SSL settings. Use this function to ensure consistent MySQL connection setup
55/// across the codebase.
56///
57/// # Arguments
58/// * `host` - MySQL server hostname or IP address
59/// * `port` - MySQL server port
60/// * `username` - MySQL username
61/// * `password` - MySQL password
62/// * `database` - Database name
63/// * `ssl_mode` - SSL mode configuration (disabled, preferred, required, verify-ca, verify-full)
64///
65/// # Returns
66/// Returns a configured `mysql_async::Pool` ready for use
67pub fn build_mysql_connection_pool(
68    host: &str,
69    port: u16,
70    username: &str,
71    password: &str,
72    database: &str,
73    ssl_mode: SslMode,
74) -> mysql_async::Pool {
75    let mut opts_builder = mysql_async::OptsBuilder::default()
76        .user(Some(username))
77        .pass(Some(password))
78        .ip_or_hostname(host)
79        .tcp_port(port)
80        .db_name(Some(database));
81
82    opts_builder = match ssl_mode {
83        SslMode::Disabled | SslMode::Preferred => opts_builder.ssl_opts(None),
84        // verify-ca and verify-full are same as required for mysql now
85        SslMode::Required | SslMode::VerifyCa | SslMode::VerifyFull => {
86            let ssl_without_verify = mysql_async::SslOpts::default()
87                .with_danger_accept_invalid_certs(true)
88                .with_danger_skip_domain_validation(true);
89            opts_builder.ssl_opts(Some(ssl_without_verify))
90        }
91    };
92
93    mysql_async::Pool::new(opts_builder)
94}
95
96#[derive(Debug, Clone, Default, PartialEq, PartialOrd, Serialize, Deserialize)]
97pub struct MySqlOffset {
98    pub filename: String,
99    pub position: u64,
100}
101
102impl MySqlOffset {
103    pub fn new(filename: String, position: u64) -> Self {
104        Self { filename, position }
105    }
106}
107
108impl MySqlOffset {
109    pub fn parse_debezium_offset(offset: &str) -> ConnectorResult<Self> {
110        let dbz_offset: DebeziumOffset = serde_json::from_str(offset)
111            .with_context(|| format!("invalid upstream offset: {}", offset))?;
112
113        Ok(Self {
114            filename: dbz_offset
115                .source_offset
116                .file
117                .context("binlog file not found in offset")?,
118            position: dbz_offset
119                .source_offset
120                .pos
121                .context("binlog position not found in offset")?,
122        })
123    }
124}
125
126pub struct MySqlExternalTable {
127    column_descs: Vec<ColumnDesc>,
128    pk_names: Vec<String>,
129}
130
131impl MySqlExternalTable {
132    pub async fn connect(config: ExternalTableConfig) -> ConnectorResult<Self> {
133        tracing::debug!("connect to mysql");
134        let options = MySqlConnectOptions::new()
135            .username(&config.username)
136            .password(&config.password)
137            .host(&config.host)
138            .port(config.port.parse::<u16>().unwrap())
139            .database(&config.database)
140            .ssl_mode(match config.ssl_mode {
141                SslMode::Disabled => sqlx::mysql::MySqlSslMode::Disabled,
142                SslMode::Preferred => sqlx::mysql::MySqlSslMode::Preferred,
143                SslMode::Required => sqlx::mysql::MySqlSslMode::Required,
144                _ => {
145                    return Err(anyhow!("unsupported SSL mode").into());
146                }
147            });
148
149        let connection = MySqlPool::connect_with(options).await?;
150        let mut schema_discovery = SchemaDiscovery::new(connection, config.database.as_str());
151
152        // discover system version first
153        let system_info = schema_discovery.discover_system().await?;
154        schema_discovery.query = SchemaQueryBuilder::new(system_info.clone());
155        let schema = Alias::new(config.database.as_str()).into_iden();
156        let table = Alias::new(config.table.as_str()).into_iden();
157        let columns = schema_discovery
158            .discover_columns(schema, table, &system_info)
159            .await?;
160        let mut column_descs = vec![];
161        let mut pk_names = vec![];
162        for col in columns {
163            let data_type = mysql_type_to_rw_type(&col.col_type)?;
164            // column name in mysql is case-insensitive, convert to lowercase
165            let col_name = col.name.to_lowercase();
166            let column_desc = if let Some(default) = col.default {
167                let snapshot_value = derive_default_value(default.clone(), &data_type)
168                    .unwrap_or_else(|e| {
169                        tracing::warn!(
170                            column = col_name,
171                            ?default,
172                            %data_type,
173                            error = %e.as_report(),
174                            "failed to derive column default value, fallback to `NULL`",
175                        );
176                        None
177                    });
178
179                ColumnDesc::named_with_default_value(
180                    col_name.clone(),
181                    ColumnId::placeholder(),
182                    data_type.clone(),
183                    snapshot_value,
184                )
185            } else {
186                ColumnDesc::named(col_name.clone(), ColumnId::placeholder(), data_type)
187            };
188
189            column_descs.push(column_desc);
190            if matches!(col.key, ColumnKey::Primary) {
191                pk_names.push(col_name);
192            }
193        }
194
195        if pk_names.is_empty() {
196            return Err(anyhow!("MySQL table doesn't define the primary key").into());
197        }
198        Ok(Self {
199            column_descs,
200            pk_names,
201        })
202    }
203
204    pub fn column_descs(&self) -> &Vec<ColumnDesc> {
205        &self.column_descs
206    }
207
208    pub fn pk_names(&self) -> &Vec<String> {
209        &self.pk_names
210    }
211}
212
213fn derive_default_value(default: ColumnDefault, data_type: &DataType) -> ConnectorResult<Datum> {
214    let datum = match default {
215        ColumnDefault::Null => None,
216        ColumnDefault::Int(val) => match data_type {
217            DataType::Int16 => Some(ScalarImpl::Int16(val as _)),
218            DataType::Int32 => Some(ScalarImpl::Int32(val as _)),
219            DataType::Int64 => Some(ScalarImpl::Int64(val)),
220            DataType::Varchar => {
221                // should be the Enum type which is mapped to Varchar
222                Some(ScalarImpl::from(val.to_string()))
223            }
224            _ => bail!("unexpected default value type for integer"),
225        },
226        ColumnDefault::Real(val) => match data_type {
227            DataType::Float32 => Some(ScalarImpl::Float32(F32::from(val as f32))),
228            DataType::Float64 => Some(ScalarImpl::Float64(val.into())),
229            DataType::Decimal => Some(ScalarImpl::Decimal(
230                Decimal::try_from(val).context("failed to convert default value to decimal")?,
231            )),
232            _ => bail!("unexpected default value type for real"),
233        },
234        ColumnDefault::String(mut val) => {
235            // mysql timestamp is mapped to timestamptz, we use UTC timezone to
236            // interpret its value
237            if data_type == &DataType::Timestamptz {
238                val = timestamp_val_to_timestamptz(val.as_str())?;
239            }
240            Some(ScalarImpl::from_text(val.as_str(), data_type).map_err(|e| anyhow!(e)).context(
241                "failed to parse mysql default value expression, only constant is supported",
242            )?)
243        }
244        ColumnDefault::CurrentTimestamp | ColumnDefault::CustomExpr(_) => {
245            bail!("MySQL CURRENT_TIMESTAMP and custom expression default value not supported")
246        }
247    };
248    Ok(datum)
249}
250
251pub fn timestamp_val_to_timestamptz(value_text: &str) -> ConnectorResult<String> {
252    let format = "%Y-%m-%d %H:%M:%S";
253    let naive_datetime = NaiveDateTime::parse_from_str(value_text, format)
254        .map_err(|err| anyhow!("failed to parse mysql timestamp value").context(err))?;
255    let postgres_timestamptz: DateTime<chrono::Utc> =
256        DateTime::<chrono::Utc>::from_naive_utc_and_offset(naive_datetime, chrono::Utc);
257    Ok(postgres_timestamptz
258        .format("%Y-%m-%d %H:%M:%S%:z")
259        .to_string())
260}
261
262pub fn type_name_to_mysql_type(ty_name: &str) -> Option<ColumnType> {
263    // Debezium schema change message may include extra qualifiers, e.g. `BIGINT UNSIGNED`,
264    // `BIGINT(20) UNSIGNED`, `INT UNSIGNED ZEROFILL`, etc.
265    let ty = ty_name.trim().to_lowercase();
266    let tokens = ty
267        .split(|c: char| c.is_whitespace() || matches!(c, '(' | ')' | ','))
268        .filter(|token| !token.is_empty())
269        .collect_vec();
270    let base = tokens.first().copied().unwrap_or_default();
271    let second = tokens.get(1).copied();
272    let is_unsigned = tokens.contains(&"unsigned");
273    let is_zero_fill = tokens.contains(&"zerofill");
274
275    let make_numeric_attr = || {
276        let mut attr = NumericAttr::default();
277        if is_unsigned {
278            attr.unsigned = Some(true);
279        }
280        if is_zero_fill {
281            attr.zero_fill = Some(true);
282        }
283        attr
284    };
285
286    match (base, second) {
287        ("character", Some("varying")) => return Some(ColumnType::Varchar(Default::default())),
288        ("double", Some("precision")) => return Some(ColumnType::Double(make_numeric_attr())),
289        ("long", Some("varchar")) => return Some(ColumnType::MediumText(Default::default())),
290        ("long", Some("varbinary")) => return Some(ColumnType::MediumBlob),
291        _ => {}
292    }
293
294    match base {
295        "serial" => Some(ColumnType::Serial),
296        "bit" => Some(ColumnType::Bit(make_numeric_attr())),
297        "tinyint" | "int1" => Some(ColumnType::TinyInt(make_numeric_attr())),
298        "bool" | "boolean" => Some(ColumnType::Bool),
299        "smallint" | "int2" => Some(ColumnType::SmallInt(make_numeric_attr())),
300        "mediumint" | "middleint" | "int3" => Some(ColumnType::MediumInt(make_numeric_attr())),
301        "int" | "integer" | "int4" => Some(ColumnType::Int(make_numeric_attr())),
302        "bigint" | "int8" => Some(ColumnType::BigInt(make_numeric_attr())),
303        "decimal" | "dec" | "fixed" | "numeric" => Some(ColumnType::Decimal(make_numeric_attr())),
304        "float" | "float4" => Some(ColumnType::Float(make_numeric_attr())),
305        "double" | "float8" | "real" => Some(ColumnType::Double(make_numeric_attr())),
306        "time" => Some(ColumnType::Time(Default::default())),
307        "datetime" => Some(ColumnType::DateTime(Default::default())),
308        "timestamp" => Some(ColumnType::Timestamp(Default::default())),
309        "year" => Some(ColumnType::Year),
310        "char" | "character" => Some(ColumnType::Char(Default::default())),
311        "nchar" => Some(ColumnType::NChar(Default::default())),
312        "varchar" => Some(ColumnType::Varchar(Default::default())),
313        "nvarchar" => Some(ColumnType::NVarchar(Default::default())),
314        "binary" => Some(ColumnType::Binary(Default::default())),
315        "varbinary" => Some(ColumnType::Varbinary(Default::default())),
316        "text" => Some(ColumnType::Text(Default::default())),
317        "tinytext" => Some(ColumnType::TinyText(Default::default())),
318        "mediumtext" => Some(ColumnType::MediumText(Default::default())),
319        "longtext" => Some(ColumnType::LongText(Default::default())),
320        "blob" => Some(ColumnType::Blob(Default::default())),
321        "tinyblob" => Some(ColumnType::TinyBlob),
322        "mediumblob" => Some(ColumnType::MediumBlob),
323        "longblob" => Some(ColumnType::LongBlob),
324        "enum" => Some(ColumnType::Enum(Default::default())),
325        "set" => Some(ColumnType::Set(Default::default())),
326        "json" => Some(ColumnType::Json),
327        "date" => Some(ColumnType::Date),
328        "geometry" => Some(ColumnType::Geometry(Default::default())),
329        "point" => Some(ColumnType::Point(Default::default())),
330        "linestring" => Some(ColumnType::LineString(Default::default())),
331        "polygon" => Some(ColumnType::Polygon(Default::default())),
332        "multipoint" => Some(ColumnType::MultiPoint(Default::default())),
333        "multilinestring" => Some(ColumnType::MultiLineString(Default::default())),
334        "multipolygon" => Some(ColumnType::MultiPolygon(Default::default())),
335        "geometrycollection" => Some(ColumnType::GeometryCollection(Default::default())),
336        _ => None,
337    }
338}
339
340fn mysql_type_is_unsigned_bigint(col_type: &ColumnType) -> bool {
341    match col_type {
342        // MySQL SERIAL is an alias for BIGINT UNSIGNED NOT NULL AUTO_INCREMENT UNIQUE.
343        ColumnType::Serial => true,
344        ColumnType::BigInt(attr) => attr.unsigned == Some(true),
345        _ => false,
346    }
347}
348
349pub fn mysql_type_to_rw_type(col_type: &ColumnType) -> ConnectorResult<DataType> {
350    let dtype = match col_type {
351        ColumnType::Serial => DataType::Decimal,
352        ColumnType::Bit(attr) => {
353            if let Some(1) = attr.maximum {
354                DataType::Boolean
355            } else {
356                return Err(
357                    anyhow!("BIT({}) type not supported", attr.maximum.unwrap_or(0)).into(),
358                );
359            }
360        }
361        // Unsigned integer family needs promotion to avoid overflow.
362        ColumnType::TinyInt(_) => DataType::Int16,
363        ColumnType::SmallInt(attr) => {
364            if attr.unsigned == Some(true) {
365                DataType::Int32
366            } else {
367                DataType::Int16
368            }
369        }
370        ColumnType::Bool => DataType::Boolean,
371        ColumnType::MediumInt(_) => DataType::Int32,
372        ColumnType::Int(attr) => {
373            if attr.unsigned == Some(true) {
374                DataType::Int64
375            } else {
376                DataType::Int32
377            }
378        }
379        ColumnType::BigInt(attr) => {
380            if attr.unsigned == Some(true) {
381                DataType::Decimal
382            } else {
383                DataType::Int64
384            }
385        }
386        ColumnType::Decimal(_) => DataType::Decimal,
387        ColumnType::Float(_) => DataType::Float32,
388        ColumnType::Double(_) => DataType::Float64,
389        ColumnType::Date => DataType::Date,
390        ColumnType::Time(_) => DataType::Time,
391        ColumnType::DateTime(_) => DataType::Timestamp,
392        ColumnType::Timestamp(_) => DataType::Timestamptz,
393        ColumnType::Year => DataType::Int32,
394        ColumnType::Char(_)
395        | ColumnType::NChar(_)
396        | ColumnType::Varchar(_)
397        | ColumnType::NVarchar(_) => DataType::Varchar,
398        ColumnType::Binary(_) | ColumnType::Varbinary(_) => DataType::Bytea,
399        ColumnType::Text(_)
400        | ColumnType::TinyText(_)
401        | ColumnType::MediumText(_)
402        | ColumnType::LongText(_) => DataType::Varchar,
403        ColumnType::Blob(_)
404        | ColumnType::TinyBlob
405        | ColumnType::MediumBlob
406        | ColumnType::LongBlob => DataType::Bytea,
407        ColumnType::Enum(_) => DataType::Varchar,
408        ColumnType::Json => DataType::Jsonb,
409        ColumnType::Set(_) => {
410            return Err(anyhow!("SET type not supported").into());
411        }
412        ColumnType::Geometry(_) => {
413            return Err(anyhow!("GEOMETRY type not supported").into());
414        }
415        ColumnType::Point(_) => {
416            return Err(anyhow!("POINT type not supported").into());
417        }
418        ColumnType::LineString(_) => {
419            return Err(anyhow!("LINE string type not supported").into());
420        }
421        ColumnType::Polygon(_) => {
422            return Err(anyhow!("POLYGON type not supported").into());
423        }
424        ColumnType::MultiPoint(_) => {
425            return Err(anyhow!("MULTI POINT type not supported").into());
426        }
427        ColumnType::MultiLineString(_) => {
428            return Err(anyhow!("MULTI LINE STRING type not supported").into());
429        }
430        ColumnType::MultiPolygon(_) => {
431            return Err(anyhow!("MULTI POLYGON type not supported").into());
432        }
433        ColumnType::GeometryCollection(_) => {
434            return Err(anyhow!("GEOMETRY COLLECTION type not supported").into());
435        }
436        ColumnType::Unknown(_) => {
437            return Err(anyhow!("Unknown MySQL data type").into());
438        }
439    };
440
441    Ok(dtype)
442}
443
444pub struct MySqlExternalTableReader {
445    rw_schema: Schema,
446    pk_indices: Vec<usize>,
447    field_names: String,
448    pool: mysql_async::Pool,
449    upstream_mysql_pk_infos: Vec<(String, ColumnType)>, // (column_name, column_type)
450    mysql_version: (u8, u8),
451    is_mariadb: bool,
452}
453
454impl ExternalTableReader for MySqlExternalTableReader {
455    async fn current_cdc_offset(&self) -> ConnectorResult<CdcOffset> {
456        let mut conn = self.pool.get_conn().await?;
457
458        // Choose SQL command based on MySQL version
459        let sql = if !self.is_mariadb && self.is_mysql_8_4_or_later() {
460            "SHOW BINARY LOG STATUS"
461        } else {
462            "SHOW MASTER STATUS"
463        };
464
465        tracing::debug!(
466            "Using SQL command: {} for MySQL version {}.{} (is_mariadb={})",
467            sql,
468            self.mysql_version.0,
469            self.mysql_version.1,
470            self.is_mariadb
471        );
472        let mut rs = conn.query::<mysql_async::Row, _>(sql).await?;
473        let row = Itertools::exactly_one(rs.iter_mut())
474            .ok()
475            .context("expect exactly one row when reading binlog offset")?;
476        drop(conn);
477        Ok(CdcOffset::MySql(MySqlOffset {
478            filename: row.take("File").unwrap(),
479            position: row.take("Position").unwrap(),
480        }))
481    }
482
483    fn snapshot_read(
484        &self,
485        table_name: SchemaTableName,
486        start_pk: Option<OwnedRow>,
487        primary_keys: Vec<String>,
488        limit: u32,
489    ) -> BoxStream<'_, ConnectorResult<OwnedRow>> {
490        self.snapshot_read_inner(table_name, start_pk, primary_keys, limit)
491    }
492
493    async fn disconnect(self) -> ConnectorResult<()> {
494        self.pool.disconnect().await.map_err(|e| e.into())
495    }
496
497    fn get_parallel_cdc_splits(
498        &self,
499        _options: CdcTableSnapshotSplitOption,
500    ) -> BoxStream<'_, ConnectorResult<CdcTableSnapshotSplit>> {
501        // TODO(zw): feat: impl
502        stream::empty::<ConnectorResult<CdcTableSnapshotSplit>>().boxed()
503    }
504
505    fn split_snapshot_read(
506        &self,
507        _table_name: SchemaTableName,
508        _left: OwnedRow,
509        _right: OwnedRow,
510        _split_columns: Vec<Field>,
511    ) -> BoxStream<'_, ConnectorResult<OwnedRow>> {
512        todo!("implement MySQL CDC parallelized backfill")
513    }
514}
515
516impl MySqlExternalTableReader {
517    /// Get MySQL version from the connection
518    async fn get_mysql_version(pool: &mysql_async::Pool) -> ConnectorResult<(u8, u8, bool)> {
519        let mut conn = pool.get_conn().await?;
520        let result: Option<String> = conn.query_first("SELECT VERSION()").await?;
521
522        if let Some(version_str) = result {
523            let parts: Vec<&str> = version_str.split('.').collect();
524            if parts.len() >= 2 {
525                let major_version = parts[0]
526                    .parse::<u8>()
527                    .context("Failed to parse major version")?;
528                let minor_version = parts[1]
529                    .parse::<u8>()
530                    .context("Failed to parse minor version")?;
531                let is_mariadb = version_str.to_lowercase().contains("mariadb");
532                return Ok((major_version, minor_version, is_mariadb));
533            }
534        }
535        Err(anyhow!("Failed to get MySQL version").into())
536    }
537
538    /// Check if MySQL version is 8.4 or later
539    fn is_mysql_8_4_or_later(&self) -> bool {
540        let (major, minor) = self.mysql_version;
541        major > 8 || (major == 8 && minor >= 4)
542    }
543
544    pub async fn new(
545        config: ExternalTableConfig,
546        rw_schema: Schema,
547        pk_indices: Vec<usize>,
548    ) -> ConnectorResult<Self> {
549        let database = config.database.clone();
550        let table = config.table.clone();
551        let pool = build_mysql_connection_pool(
552            &config.host,
553            config.port.parse::<u16>().unwrap(),
554            &config.username,
555            &config.password,
556            &config.database,
557            config.ssl_mode,
558        );
559
560        let field_names = rw_schema
561            .fields
562            .iter()
563            .filter(|f| f.name != CDC_OFFSET_COLUMN_NAME)
564            .map(|f| Self::quote_column(f.name.as_str()))
565            .join(",");
566
567        // Query MySQL primary key infos for type casting.
568        let upstream_mysql_pk_infos =
569            Self::query_upstream_pk_infos(&pool, &database, &table).await?;
570        // Get MySQL version
571        let (major_version, minor_version, is_mariadb) = Self::get_mysql_version(&pool).await?;
572        let mysql_version = (major_version, minor_version);
573        tracing::info!(
574            "MySQL version detected: {}.{} (is_mariadb={})",
575            mysql_version.0,
576            mysql_version.1,
577            is_mariadb
578        );
579
580        Ok(Self {
581            rw_schema,
582            pk_indices,
583            field_names,
584            pool,
585            upstream_mysql_pk_infos,
586            mysql_version,
587            is_mariadb,
588        })
589    }
590
591    pub fn get_normalized_table_name(table_name: &SchemaTableName) -> String {
592        // schema name is the database name in mysql
593        format!("`{}`.`{}`", table_name.schema_name, table_name.table_name)
594    }
595
596    pub fn get_cdc_offset_parser() -> CdcOffsetParseFunc {
597        Box::new(move |offset| {
598            Ok(CdcOffset::MySql(MySqlOffset::parse_debezium_offset(
599                offset,
600            )?))
601        })
602    }
603
604    /// Query upstream primary key data types, used for generating filter conditions with proper type casting.
605    async fn query_upstream_pk_infos(
606        pool: &mysql_async::Pool,
607        database: &str,
608        table: &str,
609    ) -> ConnectorResult<Vec<(String, ColumnType)>> {
610        let mut conn = pool.get_conn().await?;
611
612        // Query primary key columns and their data types
613        let sql = format!(
614            "SELECT COLUMN_NAME, COLUMN_TYPE
615            FROM INFORMATION_SCHEMA.COLUMNS
616            WHERE TABLE_SCHEMA = '{}'
617            AND TABLE_NAME = '{}'
618            AND COLUMN_KEY = 'PRI'
619            ORDER BY ORDINAL_POSITION",
620            database, table
621        );
622
623        let rs = conn.query::<mysql_async::Row, _>(sql).await?;
624
625        let mut column_infos = Vec::new();
626        for row in &rs {
627            let column_name: String = row.get(0).unwrap();
628            let column_type: String = row.get(1).unwrap();
629            let column_type =
630                type_name_to_mysql_type(&column_type).unwrap_or(ColumnType::Unknown(column_type));
631            column_infos.push((column_name, column_type));
632        }
633
634        drop(conn);
635
636        Ok(column_infos)
637    }
638
639    /// Check whether a column is `BIGINT UNSIGNED`.
640    ///
641    /// Frontend up-casts narrower unsigned integer types, and non-integer unsigned types
642    /// (`FLOAT`/`DOUBLE`/`DECIMAL UNSIGNED`) keep their own comparison semantics. Only
643    /// `BIGINT UNSIGNED` can be represented as a negative `i64` in RisingWave and needs
644    /// unsigned `u64` comparison/conversion.
645    fn needs_unsigned_i64_compare(&self, column_name: &str) -> ConnectorResult<bool> {
646        self.upstream_mysql_pk_infos
647            .iter()
648            .find(|(col_name, _)| col_name.eq_ignore_ascii_case(column_name))
649            .map(|(_, col_type)| mysql_type_is_unsigned_bigint(col_type))
650            .ok_or_else(|| {
651                anyhow!(
652                    "primary key column `{column_name}` not found in upstream MySQL primary key info"
653                )
654                .into()
655            })
656    }
657
658    /// For each given primary key column (by name), whether it needs unsigned `i64` comparison.
659    pub(crate) fn pk_column_unsigned_i64_compare_flags(
660        &self,
661        pk_names: &[String],
662    ) -> ConnectorResult<Vec<bool>> {
663        pk_names
664            .iter()
665            .map(|name| self.needs_unsigned_i64_compare(name))
666            .collect()
667    }
668
669    /// Convert negative i64 to unsigned u64 based on column type
670    fn convert_negative_to_unsigned(&self, negative_val: i64) -> u64 {
671        negative_val as u64
672    }
673
674    #[try_stream(boxed, ok = OwnedRow, error = ConnectorError)]
675    async fn snapshot_read_inner(
676        &self,
677        table_name: SchemaTableName,
678        start_pk_row: Option<OwnedRow>,
679        primary_keys: Vec<String>,
680        limit: u32,
681    ) {
682        let order_key = primary_keys
683            .iter()
684            .map(|col| Self::quote_column(col))
685            .join(",");
686        let sql = if start_pk_row.is_none() {
687            format!(
688                "SELECT {} FROM {} ORDER BY {} LIMIT {limit}",
689                self.field_names,
690                Self::get_normalized_table_name(&table_name),
691                order_key,
692            )
693        } else {
694            let filter_expr = Self::filter_expression(&primary_keys);
695            format!(
696                "SELECT {} FROM {} WHERE {} ORDER BY {} LIMIT {limit}",
697                self.field_names,
698                Self::get_normalized_table_name(&table_name),
699                filter_expr,
700                order_key,
701            )
702        };
703        let mut conn = self.pool.get_conn().await?;
704        // Set session timezone to UTC
705        conn.exec_drop("SET time_zone = \"+00:00\"", ()).await?;
706
707        if let Some(start_pk_row) = start_pk_row {
708            let field_map = self
709                .rw_schema
710                .fields
711                .iter()
712                .map(|f| (f.name.as_str(), f.data_type.clone()))
713                .collect::<HashMap<_, _>>();
714
715            // fill in start primary key params
716            let params: Vec<_> = primary_keys
717                .iter()
718                .zip_eq_fast(start_pk_row.into_iter())
719                .map(|(pk, datum)| {
720                    if let Some(value) = datum {
721                        let ty = field_map.get(pk.as_str()).unwrap();
722                        let val = match ty {
723                            DataType::Boolean => Value::from(value.into_bool()),
724                            DataType::Int16 => Value::from(value.into_int16()),
725                            DataType::Int32 => Value::from(value.into_int32()),
726                            DataType::Int64 => {
727                                let int64_val = value.into_int64();
728                                if int64_val < 0 && self.needs_unsigned_i64_compare(pk.as_str())? {
729                                    Value::from(self.convert_negative_to_unsigned(int64_val))
730                                } else {
731                                    Value::from(int64_val)
732                                }
733                            }
734                            DataType::Float32 => Value::from(value.into_float32().into_inner()),
735                            DataType::Float64 => Value::from(value.into_float64().into_inner()),
736                            DataType::Varchar => Value::from(String::from(value.into_utf8())),
737                            DataType::Date => Value::from(value.into_date().0),
738                            DataType::Time => Value::from(value.into_time().0),
739                            DataType::Timestamp => Value::from(value.into_timestamp().0),
740                            DataType::Decimal => Value::from(value.into_decimal().to_string()),
741                            DataType::Timestamptz => {
742                                // Convert timestamptz to NaiveDateTime for MySQL TIMESTAMP comparison
743                                // MySQL expects NaiveDateTime for TIMESTAMP parameters
744                                let ts = value.into_timestamptz();
745                                let datetime_utc = ts.to_datetime_utc();
746                                let naive_datetime = datetime_utc.naive_utc();
747                                Value::from(naive_datetime)
748                            }
749                            _ => bail!("unsupported primary key data type: {}", ty),
750                        };
751                        ConnectorResult::Ok((pk.to_lowercase(), val))
752                    } else {
753                        bail!("primary key {} cannot be null", pk);
754                    }
755                })
756                .try_collect::<_, _, ConnectorError>()?;
757
758            tracing::debug!("snapshot read params: {:?}", &params);
759            let rs_stream = sql
760                .with(Params::from(params))
761                .stream::<mysql_async::Row, _>(&mut conn)
762                .await?;
763
764            let row_stream = rs_stream.map(|row| {
765                // convert mysql row into OwnedRow
766                let mut row = row?;
767                mysql_row_to_owned_row_with_strict_pk(&mut row, &self.rw_schema, &self.pk_indices)
768                    .map_err(ConnectorError::from)
769            });
770            pin_mut!(row_stream);
771            #[for_await]
772            for row in row_stream {
773                let row = row?;
774                yield row;
775            }
776        } else {
777            let rs_stream = sql.stream::<mysql_async::Row, _>(&mut conn).await?;
778            let row_stream = rs_stream.map(|row| {
779                // convert mysql row into OwnedRow
780                let mut row = row?;
781                mysql_row_to_owned_row_with_strict_pk(&mut row, &self.rw_schema, &self.pk_indices)
782                    .map_err(ConnectorError::from)
783            });
784            pin_mut!(row_stream);
785            #[for_await]
786            for row in row_stream {
787                let row = row?;
788                yield row;
789            }
790        }
791        drop(conn);
792    }
793
794    // mysql cannot leverage the given key to narrow down the range of scan,
795    // we need to rewrite the comparison conditions by our own.
796    // (a, b) > (x, y) => (`a` > x) OR ((`a` = x) AND (`b` > y))
797    fn filter_expression(columns: &[String]) -> String {
798        let mut conditions = vec![];
799        // push the first condition
800        conditions.push(format!(
801            "({} > :{})",
802            Self::quote_column(&columns[0]),
803            columns[0].to_lowercase()
804        ));
805        for i in 2..=columns.len() {
806            // '=' condition
807            let mut condition = String::new();
808            for (j, col) in columns.iter().enumerate().take(i - 1) {
809                if j == 0 {
810                    condition.push_str(&format!(
811                        "({} = :{})",
812                        Self::quote_column(col),
813                        col.to_lowercase()
814                    ));
815                } else {
816                    condition.push_str(&format!(
817                        " AND ({} = :{})",
818                        Self::quote_column(col),
819                        col.to_lowercase()
820                    ));
821                }
822            }
823            // '>' condition
824            condition.push_str(&format!(
825                " AND ({} > :{})",
826                Self::quote_column(&columns[i - 1]),
827                columns[i - 1].to_lowercase()
828            ));
829            conditions.push(format!("({})", condition));
830        }
831        if columns.len() > 1 {
832            conditions.join(" OR ")
833        } else {
834            conditions.join("")
835        }
836    }
837
838    fn quote_column(column: &str) -> String {
839        format!("`{}`", column)
840    }
841}
842
843#[cfg(test)]
844mod tests {
845    use std::collections::HashMap;
846
847    use futures::pin_mut;
848    use futures_async_stream::for_await;
849    use maplit::{convert_args, hashmap};
850    use risingwave_common::catalog::{ColumnDesc, ColumnId, Field, Schema};
851    use risingwave_common::types::DataType;
852    use sea_schema::mysql::def::ColumnType;
853
854    use super::{mysql_type_is_unsigned_bigint, mysql_type_to_rw_type, type_name_to_mysql_type};
855    use crate::source::cdc::external::mysql::MySqlExternalTable;
856    use crate::source::cdc::external::{
857        CdcOffset, ExternalTableConfig, ExternalTableReader, MySqlExternalTableReader, MySqlOffset,
858        SchemaTableName,
859    };
860
861    fn parse_mysql_type_name(ty_name: &str) -> ColumnType {
862        type_name_to_mysql_type(ty_name).unwrap()
863    }
864
865    #[test]
866    fn test_mysql_unsigned_bigint_type_detection() {
867        for ty_name in [
868            "SERIAL",
869            "BIGINT UNSIGNED",
870            "BIGINT(20) UNSIGNED",
871            "BIGINT UNSIGNED ZEROFILL",
872            "INT8 UNSIGNED",
873        ] {
874            assert!(
875                mysql_type_is_unsigned_bigint(&parse_mysql_type_name(ty_name)),
876                "{ty_name}"
877            );
878        }
879
880        for ty_name in [
881            "BIGINT",
882            "INTEGER UNSIGNED",
883            "INT4 UNSIGNED",
884            "MEDIUMINT UNSIGNED",
885            "DECIMAL UNSIGNED",
886            "FLOAT8 UNSIGNED",
887        ] {
888            assert!(
889                !mysql_type_is_unsigned_bigint(&parse_mysql_type_name(ty_name)),
890                "{ty_name}"
891            );
892        }
893    }
894
895    #[test]
896    fn test_mysql_type_aliases() {
897        assert!(matches!(
898            parse_mysql_type_name("INTEGER UNSIGNED"),
899            ColumnType::Int(attr) if attr.unsigned == Some(true)
900        ));
901        assert!(matches!(
902            parse_mysql_type_name("INT1 UNSIGNED"),
903            ColumnType::TinyInt(attr) if attr.unsigned == Some(true)
904        ));
905        assert!(matches!(
906            parse_mysql_type_name("INT2 UNSIGNED"),
907            ColumnType::SmallInt(attr) if attr.unsigned == Some(true)
908        ));
909        assert!(matches!(
910            parse_mysql_type_name("INT3 UNSIGNED"),
911            ColumnType::MediumInt(attr) if attr.unsigned == Some(true)
912        ));
913        assert!(matches!(
914            parse_mysql_type_name("INT4 UNSIGNED"),
915            ColumnType::Int(attr) if attr.unsigned == Some(true)
916        ));
917        assert!(matches!(
918            parse_mysql_type_name("INT8 UNSIGNED"),
919            ColumnType::BigInt(attr) if attr.unsigned == Some(true)
920        ));
921        assert!(matches!(
922            parse_mysql_type_name("MIDDLEINT"),
923            ColumnType::MediumInt(_)
924        ));
925        assert!(matches!(
926            parse_mysql_type_name("NUMERIC"),
927            ColumnType::Decimal(_)
928        ));
929        assert!(matches!(
930            parse_mysql_type_name("CHARACTER VARYING(64)"),
931            ColumnType::Varchar(_)
932        ));
933        assert!(matches!(
934            parse_mysql_type_name("LONG VARBINARY"),
935            ColumnType::MediumBlob
936        ));
937    }
938
939    #[test]
940    fn test_mysql_serial_maps_as_unsigned_bigint() {
941        let col_type = parse_mysql_type_name("SERIAL");
942        assert!(mysql_type_is_unsigned_bigint(&col_type));
943
944        let serial_type = mysql_type_to_rw_type(&col_type).unwrap();
945        let unsigned_bigint_type =
946            mysql_type_to_rw_type(&parse_mysql_type_name("BIGINT UNSIGNED")).unwrap();
947        assert_eq!(serial_type, DataType::Decimal);
948        assert_eq!(serial_type, unsigned_bigint_type);
949    }
950
951    #[ignore]
952    #[tokio::test]
953    async fn test_mysql_schema() {
954        let config = ExternalTableConfig {
955            connector: "mysql-cdc".to_owned(),
956            host: "localhost".to_owned(),
957            port: "8306".to_owned(),
958            username: "root".to_owned(),
959            password: "123456".to_owned(),
960            database: "mydb".to_owned(),
961            schema: "".to_owned(),
962            table: "part".to_owned(),
963            ssl_mode: Default::default(),
964            ssl_root_cert: None,
965            encrypt: "false".to_owned(),
966        };
967
968        let table = MySqlExternalTable::connect(config).await.unwrap();
969        println!("columns: {:?}", table.column_descs);
970        println!("primary keys: {:?}", table.pk_names);
971    }
972
973    #[test]
974    fn test_mysql_filter_expr() {
975        let cols = vec!["id".to_owned()];
976        let expr = MySqlExternalTableReader::filter_expression(&cols);
977        assert_eq!(expr, "(`id` > :id)");
978
979        let cols = vec!["aa".to_owned(), "bb".to_owned(), "cc".to_owned()];
980        let expr = MySqlExternalTableReader::filter_expression(&cols);
981        assert_eq!(
982            expr,
983            "(`aa` > :aa) OR ((`aa` = :aa) AND (`bb` > :bb)) OR ((`aa` = :aa) AND (`bb` = :bb) AND (`cc` > :cc))"
984        );
985    }
986
987    #[test]
988    fn test_mysql_binlog_offset() {
989        let off0_str = r#"{ "sourcePartition": { "server": "test" }, "sourceOffset": { "ts_sec": 1670876905, "file": "binlog.000001", "pos": 105622, "snapshot": true }, "isHeartbeat": false }"#;
990        let off1_str = r#"{ "sourcePartition": { "server": "test" }, "sourceOffset": { "ts_sec": 1670876905, "file": "binlog.000007", "pos": 1062363217, "snapshot": true }, "isHeartbeat": false }"#;
991        let off2_str = r#"{ "sourcePartition": { "server": "test" }, "sourceOffset": { "ts_sec": 1670876905, "file": "binlog.000007", "pos": 659687560, "snapshot": true }, "isHeartbeat": false }"#;
992        let off3_str = r#"{ "sourcePartition": { "server": "test" }, "sourceOffset": { "ts_sec": 1670876905, "file": "binlog.000008", "pos": 7665875, "snapshot": true }, "isHeartbeat": false }"#;
993        let off4_str = r#"{ "sourcePartition": { "server": "test" }, "sourceOffset": { "ts_sec": 1670876905, "file": "binlog.000008", "pos": 7665875, "snapshot": true }, "isHeartbeat": false }"#;
994
995        let off0 = CdcOffset::MySql(MySqlOffset::parse_debezium_offset(off0_str).unwrap());
996        let off1 = CdcOffset::MySql(MySqlOffset::parse_debezium_offset(off1_str).unwrap());
997        let off2 = CdcOffset::MySql(MySqlOffset::parse_debezium_offset(off2_str).unwrap());
998        let off3 = CdcOffset::MySql(MySqlOffset::parse_debezium_offset(off3_str).unwrap());
999        let off4 = CdcOffset::MySql(MySqlOffset::parse_debezium_offset(off4_str).unwrap());
1000
1001        assert!(off0 <= off1);
1002        assert!(off1 > off2);
1003        assert!(off2 < off3);
1004        assert_eq!(off3, off4);
1005    }
1006
1007    // manual test case
1008    #[ignore]
1009    #[tokio::test]
1010    async fn test_mysql_table_reader() {
1011        let columns = [
1012            ColumnDesc::named("v1", ColumnId::new(1), DataType::Int32),
1013            ColumnDesc::named("v2", ColumnId::new(2), DataType::Decimal),
1014            ColumnDesc::named("v3", ColumnId::new(3), DataType::Varchar),
1015            ColumnDesc::named("v4", ColumnId::new(4), DataType::Date),
1016        ];
1017        let rw_schema = Schema {
1018            fields: columns.iter().map(Field::from).collect(),
1019        };
1020        let props: HashMap<String, String> = convert_args!(hashmap!(
1021                "hostname" => "localhost",
1022                "port" => "8306",
1023                "username" => "root",
1024                "password" => "123456",
1025                "database.name" => "mytest",
1026                "table.name" => "t1"));
1027
1028        let config =
1029            serde_json::from_value::<ExternalTableConfig>(serde_json::to_value(props).unwrap())
1030                .unwrap();
1031        let reader = MySqlExternalTableReader::new(config, rw_schema, vec![0])
1032            .await
1033            .unwrap();
1034        let offset = reader.current_cdc_offset().await.unwrap();
1035        println!("BinlogOffset: {:?}", offset);
1036
1037        let off0_str = r#"{ "sourcePartition": { "server": "test" }, "sourceOffset": { "ts_sec": 1670876905, "file": "binlog.000001", "pos": 105622, "snapshot": true }, "isHeartbeat": false }"#;
1038        let parser = MySqlExternalTableReader::get_cdc_offset_parser();
1039        println!("parsed offset: {:?}", parser(off0_str).unwrap());
1040        let table_name = SchemaTableName {
1041            schema_name: "mytest".to_owned(),
1042            table_name: "t1".to_owned(),
1043        };
1044
1045        let stream = reader.snapshot_read(table_name, None, vec!["v1".to_owned()], 1000);
1046        pin_mut!(stream);
1047        #[for_await]
1048        for row in stream {
1049            println!("OwnedRow: {:?}", row);
1050        }
1051    }
1052}