use crate::db_type::DbType;
use crate::error::DbError;
use std::fmt;
pub const MAX_IDENTIFIER_LEN: usize = 63;
pub trait Dialect: Send + Sync {
fn db_type(&self) -> DbType;
fn quote(&self, identifier: &str) -> String;
fn quote_checked(&self, identifier: &str) -> Result<String, DbError> {
crate::sql_safety::validate_identifier(identifier, "identifier")?;
Ok(self.quote(identifier))
}
fn escape_string(&self, s: &str) -> String;
fn supports_returning(&self) -> bool;
fn build_pagination(&self, sql: &str, page: u64, limit: u64) -> String;
fn json_type(&self) -> &'static str;
fn json_extract(&self, column: &str, path: &str) -> String;
fn full_text_search(&self, columns: &[&str], keyword: &str) -> String;
fn bool_to_int(&self, expr: &str) -> String;
fn concat(&self, parts: &[&str]) -> String;
fn supports_if_exists(&self) -> bool;
fn supports_if_not_exists(&self) -> bool;
fn auto_increment_keyword(&self) -> &'static str;
fn last_insert_id_sql(&self) -> Option<&'static str>;
fn build_create_table(&self, table: &str, columns: &[ColumnDef]) -> String;
fn build_alter_table(&self, table: &str, changes: &[TableChange]) -> String;
fn build_drop_table(&self, table: &str, if_exists: bool) -> String {
if if_exists && self.supports_if_exists() {
format!("DROP TABLE IF EXISTS {}", self.quote(table))
} else {
format!("DROP TABLE {}", self.quote(table))
}
}
}
#[derive(Debug, Clone)]
pub struct ColumnDef {
pub name: String,
pub sql_type: String,
pub nullable: bool,
pub default: Option<String>,
pub auto_increment: bool,
pub primary_key: bool,
}
#[derive(Debug, Clone)]
pub enum TableChange {
AddColumn(ColumnDef),
DropColumn(String),
ModifyColumn(ColumnDef),
AddIndex(String, Vec<String>),
DropIndex(String),
AddForeignKey {
columns: Vec<String>,
reference_table: String,
reference_columns: Vec<String>,
},
}
pub struct MySqlDialect;
impl Dialect for MySqlDialect {
fn db_type(&self) -> DbType {
DbType::MySQL
}
fn quote(&self, identifier: &str) -> String {
format!("`{}`", identifier.replace('`', "``"))
}
fn escape_string(&self, s: &str) -> String {
let mut escaped = String::with_capacity(s.len() * 2);
for c in s.chars() {
match c {
'\\' => escaped.push_str("\\\\"),
'\'' => escaped.push_str("\\'"),
'\0' => escaped.push_str("\\0"),
'\n' => escaped.push_str("\\n"),
'\r' => escaped.push_str("\\r"),
'\t' => escaped.push_str("\\t"),
'\x1a' => escaped.push_str("\\Z"),
_ => escaped.push(c),
}
}
escaped
}
fn supports_returning(&self) -> bool {
false
}
fn build_pagination(&self, sql: &str, page: u64, limit: u64) -> String {
let offset = page.saturating_sub(1).saturating_mul(limit);
format!("{} LIMIT {} OFFSET {}", sql, limit, offset)
}
fn json_type(&self) -> &'static str {
"JSON"
}
fn json_extract(&self, column: &str, path: &str) -> String {
let normalized = if path.starts_with('$') {
path.to_string()
} else {
format!("$.{}", path)
};
format!(
"JSON_EXTRACT({}, '{}')",
column,
self.escape_string(&normalized)
)
}
fn full_text_search(&self, columns: &[&str], keyword: &str) -> String {
let cols = columns.join(", ");
let escaped = self.escape_string(keyword);
format!(
"MATCH({}) AGAINST('{}' IN NATURAL LANGUAGE MODE)",
cols, escaped
)
}
fn bool_to_int(&self, expr: &str) -> String {
format!("IF({}, 1, 0)", expr)
}
fn concat(&self, parts: &[&str]) -> String {
if parts.is_empty() {
return "NULL".to_string();
}
let concat_parts: Vec<String> = parts
.iter()
.map(|p| format!("CAST({} AS CHAR)", p))
.collect();
format!("CONCAT({})", concat_parts.join(", "))
}
fn supports_if_exists(&self) -> bool {
true
}
fn supports_if_not_exists(&self) -> bool {
true
}
fn auto_increment_keyword(&self) -> &'static str {
"AUTO_INCREMENT"
}
fn last_insert_id_sql(&self) -> Option<&'static str> {
Some("LAST_INSERT_ID()")
}
fn build_create_table(&self, table: &str, columns: &[ColumnDef]) -> String {
let cols: Vec<String> = columns
.iter()
.map(|col| {
let mut sql = format!("{} {}", self.quote(&col.name), col.sql_type);
if !col.nullable {
sql.push_str(" NOT NULL");
}
if let Some(default) = &col.default {
sql.push_str(&format!(" DEFAULT {}", default));
}
if col.auto_increment {
sql.push_str(&format!(" {}", self.auto_increment_keyword()));
}
if col.primary_key {
sql.push_str(" PRIMARY KEY");
}
sql
})
.collect();
format!("CREATE TABLE {} ({})", self.quote(table), cols.join(", "))
}
fn build_alter_table(&self, table: &str, changes: &[TableChange]) -> String {
let stmts: Vec<String> = changes.iter().map(|change| {
match change {
TableChange::AddColumn(col) => {
let mut sql = format!("ALTER TABLE {} ADD {}", self.quote(table), self.quote(&col.name));
sql.push_str(&format!(" {}", col.sql_type));
if !col.nullable {
sql.push_str(" NOT NULL");
}
if let Some(default) = &col.default {
sql.push_str(&format!(" DEFAULT {}", default));
}
sql
}
TableChange::DropColumn(name) => {
format!("ALTER TABLE {} DROP COLUMN {}", self.quote(table), self.quote(name))
}
TableChange::ModifyColumn(col) => {
let mut sql = format!("ALTER TABLE {} MODIFY COLUMN {} {}", self.quote(table), self.quote(&col.name), col.sql_type);
if !col.nullable {
sql.push_str(" NOT NULL");
}
if let Some(default) = &col.default {
sql.push_str(&format!(" DEFAULT {}", default));
}
sql
}
TableChange::AddIndex(name, cols) => {
format!("ALTER TABLE {} ADD INDEX {} ({})", self.quote(table), name, cols.join(", "))
}
TableChange::DropIndex(name) => {
format!("ALTER TABLE {} DROP INDEX {}", self.quote(table), name)
}
TableChange::AddForeignKey { columns, reference_table, reference_columns } => {
format!("ALTER TABLE {} ADD CONSTRAINT fk_{}_{} FOREIGN KEY ({}) REFERENCES {} ({})",
self.quote(table),
table,
columns.join("_"),
columns.iter().map(|c| self.quote(c)).collect::<Vec<_>>().join(", "),
self.quote(reference_table),
reference_columns.iter().map(|c| self.quote(c)).collect::<Vec<_>>().join(", "))
}
}
}).collect();
stmts.join("; ")
}
}
pub struct PostgreSqlDialect;
impl Dialect for PostgreSqlDialect {
fn db_type(&self) -> DbType {
DbType::PostgreSQL
}
fn quote(&self, identifier: &str) -> String {
format!("\"{}\"", identifier.replace('"', "\"\""))
}
fn escape_string(&self, s: &str) -> String {
let mut escaped = String::with_capacity(s.len() * 2);
for c in s.chars() {
match c {
'\'' => escaped.push_str("''"),
_ => escaped.push(c),
}
}
escaped
}
fn supports_returning(&self) -> bool {
true
}
fn build_pagination(&self, sql: &str, page: u64, limit: u64) -> String {
let offset = page.saturating_sub(1).saturating_mul(limit);
format!("{} LIMIT {} OFFSET {}", sql, limit, offset)
}
fn json_type(&self) -> &'static str {
"JSONB"
}
fn json_extract(&self, column: &str, path: &str) -> String {
let normalized = path.trim_start_matches("$.");
let parts: Vec<&str> = normalized.split('.').filter(|s| !s.is_empty()).collect();
let path_lit = parts
.iter()
.map(|p| {
let needs_quoting = p.chars().any(|c| matches!(c, ',' | '{' | '}' | '"' | '\\'));
if needs_quoting {
let escaped = p.replace('\\', "\\\\").replace('"', "\\\"");
format!("\"{}\"", escaped)
} else {
p.to_string()
}
})
.collect::<Vec<_>>()
.join(",");
let path_lit_escaped = path_lit.replace('\'', "''");
format!("{}#>>'{{{}}}'", column, path_lit_escaped)
}
fn full_text_search(&self, columns: &[&str], keyword: &str) -> String {
let cols = columns
.iter()
.map(|c| format!("{}::text", c))
.collect::<Vec<_>>()
.join(" || ' ' || ");
let escaped = self.escape_string(keyword);
format!("to_tsvector({}) @@ to_tsquery('{}')", cols, escaped)
}
fn bool_to_int(&self, expr: &str) -> String {
format!("(CASE WHEN {} THEN 1 ELSE 0 END)", expr)
}
fn concat(&self, parts: &[&str]) -> String {
if parts.is_empty() {
return "NULL".to_string();
}
format!("CONCAT({})", parts.join(", "))
}
fn supports_if_exists(&self) -> bool {
true
}
fn supports_if_not_exists(&self) -> bool {
true
}
fn auto_increment_keyword(&self) -> &'static str {
"GENERATED BY DEFAULT AS IDENTITY"
}
fn last_insert_id_sql(&self) -> Option<&'static str> {
Some("lastval()")
}
fn build_create_table(&self, table: &str, columns: &[ColumnDef]) -> String {
let cols: Vec<String> = columns
.iter()
.map(|col| {
let mut sql = format!("{} {}", self.quote(&col.name), col.sql_type);
if !col.nullable {
sql.push_str(" NOT NULL");
}
if let Some(default) = &col.default {
sql.push_str(&format!(" DEFAULT {}", default));
}
if col.primary_key {
sql.push_str(" PRIMARY KEY");
}
sql
})
.collect();
format!("CREATE TABLE {} ({})", self.quote(table), cols.join(", "))
}
fn build_alter_table(&self, table: &str, changes: &[TableChange]) -> String {
let stmts: Vec<String> = changes.iter().map(|change| {
match change {
TableChange::AddColumn(col) => {
let mut sql = format!("ALTER TABLE {} ADD COLUMN {} {}", self.quote(table), self.quote(&col.name), col.sql_type);
if !col.nullable {
sql.push_str(" NOT NULL");
}
if let Some(default) = &col.default {
sql.push_str(&format!(" DEFAULT {}", default));
}
sql
}
TableChange::DropColumn(name) => {
format!("ALTER TABLE {} DROP COLUMN {}", self.quote(table), self.quote(name))
}
TableChange::ModifyColumn(col) => {
let mut sql = format!("ALTER TABLE {} ALTER COLUMN {} TYPE {}", self.quote(table), self.quote(&col.name), col.sql_type);
if !col.nullable {
sql.push_str(&format!(", ALTER COLUMN {} SET NOT NULL", self.quote(&col.name)));
}
if let Some(default) = &col.default {
sql.push_str(&format!(", ALTER COLUMN {} SET DEFAULT {}", self.quote(&col.name), default));
}
sql
}
TableChange::AddIndex(name, cols) => {
format!("CREATE INDEX {} ON {} ({})", name, self.quote(table), cols.join(", "))
}
TableChange::DropIndex(name) => {
format!("DROP INDEX {}", name)
}
TableChange::AddForeignKey { columns, reference_table, reference_columns } => {
format!("ALTER TABLE {} ADD CONSTRAINT fk_{}_{} FOREIGN KEY ({}) REFERENCES {} ({})",
self.quote(table),
table,
columns.join("_"),
columns.iter().map(|c| self.quote(c)).collect::<Vec<_>>().join(", "),
self.quote(reference_table),
reference_columns.iter().map(|c| self.quote(c)).collect::<Vec<_>>().join(", "))
}
}
}).collect();
stmts.join("; ")
}
}
pub struct SqliteDialect;
impl Dialect for SqliteDialect {
fn db_type(&self) -> DbType {
DbType::Sqlite
}
fn quote(&self, identifier: &str) -> String {
format!("\"{}\"", identifier.replace('"', "\"\""))
}
fn escape_string(&self, s: &str) -> String {
let mut escaped = String::with_capacity(s.len() * 2);
for c in s.chars() {
match c {
'\'' => escaped.push_str("''"),
_ => escaped.push(c),
}
}
escaped
}
fn supports_returning(&self) -> bool {
true
}
fn build_pagination(&self, sql: &str, page: u64, limit: u64) -> String {
let offset = page.saturating_sub(1).saturating_mul(limit);
format!("{} LIMIT {} OFFSET {}", sql, limit, offset)
}
fn json_type(&self) -> &'static str {
"TEXT"
}
fn json_extract(&self, column: &str, path: &str) -> String {
let normalized = if path.starts_with('$') {
path.to_string()
} else {
format!("$.{}", path)
};
format!(
"json_extract({}, '{}')",
column,
self.escape_string(&normalized)
)
}
fn full_text_search(&self, columns: &[&str], keyword: &str) -> String {
if columns.is_empty() {
return "0".to_string();
}
let escaped = self.escape_string(keyword);
columns
.iter()
.map(|c| format!("{} LIKE '%{}%'", c.trim(), escaped))
.collect::<Vec<_>>()
.join(" OR ")
}
fn bool_to_int(&self, expr: &str) -> String {
expr.to_string()
}
fn concat(&self, parts: &[&str]) -> String {
if parts.is_empty() {
return "NULL".to_string();
}
let coalesced: Vec<String> = parts
.iter()
.map(|p| format!("COALESCE({}, '')", p))
.collect();
coalesced.join(" || ")
}
fn supports_if_exists(&self) -> bool {
true
}
fn supports_if_not_exists(&self) -> bool {
true
}
fn auto_increment_keyword(&self) -> &'static str {
"AUTOINCREMENT"
}
fn last_insert_id_sql(&self) -> Option<&'static str> {
Some("last_insert_rowid()")
}
fn build_create_table(&self, table: &str, columns: &[ColumnDef]) -> String {
let cols: Vec<String> = columns
.iter()
.map(|col| {
let mut sql = format!("{} {}", self.quote(&col.name), col.sql_type);
if !col.nullable {
sql.push_str(" NOT NULL");
}
if let Some(default) = &col.default {
sql.push_str(&format!(" DEFAULT {}", default));
}
if col.auto_increment {
sql.push_str(" PRIMARY KEY AUTOINCREMENT");
} else if col.primary_key {
sql.push_str(" PRIMARY KEY");
}
sql
})
.collect();
format!("CREATE TABLE {} ({})", self.quote(table), cols.join(", "))
}
fn build_alter_table(&self, table: &str, changes: &[TableChange]) -> String {
let stmts: Vec<String> = changes
.iter()
.map(|change| {
match change {
TableChange::AddColumn(col) => {
let mut sql = format!(
"ALTER TABLE {} ADD COLUMN {} {}",
self.quote(table),
self.quote(&col.name),
col.sql_type
);
if !col.nullable {
sql.push_str(" NOT NULL");
}
if let Some(default) = &col.default {
sql.push_str(&format!(" DEFAULT {}", default));
}
sql
}
TableChange::DropColumn(name) => {
format!(
"ALTER TABLE {} DROP COLUMN {}",
self.quote(table),
self.quote(name)
)
}
TableChange::ModifyColumn(col) => {
format!(
"-- SQLite 不支持 MODIFY COLUMN({} {}),需重建表",
col.name, col.sql_type
)
}
TableChange::AddIndex(name, cols) => {
format!(
"CREATE INDEX {} ON {} ({})",
name,
self.quote(table),
cols.join(", ")
)
}
TableChange::DropIndex(name) => {
format!("DROP INDEX {}", name)
}
TableChange::AddForeignKey {
columns,
reference_table,
reference_columns: _,
} => {
format!(
"-- SQLite 不支持 ADD FOREIGN KEY({} -> {}),需重建表",
columns.join(","),
reference_table
)
}
}
})
.collect();
stmts.join("; ")
}
}
fn map_to_oracle_type(sql_type: &str) -> String {
let upper = sql_type.to_uppercase();
let trimmed = upper.trim();
if trimmed.starts_with("BIGINT") {
sql_type.replacen("BIGINT", "NUMBER(19)", 1)
} else if trimmed.starts_with("VARCHAR2") {
sql_type.to_string()
} else if trimmed.starts_with("VARCHAR") {
sql_type.replacen("VARCHAR", "VARCHAR2", 1)
} else if matches!(trimmed, "TEXT" | "MEDIUMTEXT" | "LONGTEXT" | "TINYTEXT") {
"CLOB".to_string()
} else if matches!(trimmed, "BOOLEAN" | "BOOL") {
"NUMBER(1)".to_string()
} else if trimmed == "INTEGER" {
"NUMBER(10)".to_string()
} else if trimmed.starts_with("INT") {
sql_type.replacen("INT", "NUMBER(10)", 1)
} else {
sql_type.to_string()
}
}
pub struct OracleDialect;
impl Dialect for OracleDialect {
fn db_type(&self) -> DbType {
DbType::Oracle
}
fn quote(&self, identifier: &str) -> String {
format!("\"{}\"", identifier.replace('"', "\"\""))
}
fn escape_string(&self, s: &str) -> String {
let mut escaped = String::with_capacity(s.len() * 2);
for c in s.chars() {
match c {
'\'' => escaped.push_str("''"),
_ => escaped.push(c),
}
}
escaped
}
fn supports_returning(&self) -> bool {
true
}
fn build_pagination(&self, sql: &str, page: u64, limit: u64) -> String {
let offset = page.saturating_sub(1).saturating_mul(limit);
format!(
"{} OFFSET {} ROWS FETCH NEXT {} ROWS ONLY",
sql, offset, limit
)
}
fn json_type(&self) -> &'static str {
"JSON"
}
fn json_extract(&self, column: &str, path: &str) -> String {
let normalized = if path.starts_with('$') {
path.to_string()
} else {
format!("$.{}", path)
};
format!(
"JSON_VALUE({}, '{}')",
column,
self.escape_string(&normalized)
)
}
fn full_text_search(&self, columns: &[&str], keyword: &str) -> String {
if columns.is_empty() {
return "0".to_string();
}
let escaped = self.escape_string(keyword);
let parts: Vec<String> = columns
.iter()
.map(|c| format!("CONTAINS({}, '{}', 1) > 0", c, escaped))
.collect();
parts.join(" OR ")
}
fn bool_to_int(&self, expr: &str) -> String {
format!("(CASE WHEN {} THEN 1 ELSE 0 END)", expr)
}
fn concat(&self, parts: &[&str]) -> String {
if parts.is_empty() {
return "NULL".to_string();
}
parts.join(" || ")
}
fn supports_if_exists(&self) -> bool {
true
}
fn supports_if_not_exists(&self) -> bool {
true
}
fn auto_increment_keyword(&self) -> &'static str {
"GENERATED BY DEFAULT AS IDENTITY"
}
fn last_insert_id_sql(&self) -> Option<&'static str> {
None
}
fn build_create_table(&self, table: &str, columns: &[ColumnDef]) -> String {
let cols: Vec<String> = columns
.iter()
.map(|col| {
let oracle_type = map_to_oracle_type(&col.sql_type);
let mut sql = format!("{} {}", self.quote(&col.name), oracle_type);
if !col.nullable && !col.auto_increment {
sql.push_str(" NOT NULL");
}
if let Some(default) = &col.default {
sql.push_str(&format!(" DEFAULT {}", default));
}
if col.auto_increment {
sql.push_str(&format!(" {}", self.auto_increment_keyword()));
}
if col.primary_key {
sql.push_str(" PRIMARY KEY");
}
sql
})
.collect();
format!("CREATE TABLE {} ({})", self.quote(table), cols.join(", "))
}
fn build_alter_table(&self, table: &str, changes: &[TableChange]) -> String {
let stmts: Vec<String> = changes
.iter()
.map(|change| match change {
TableChange::AddColumn(col) => {
let oracle_type = map_to_oracle_type(&col.sql_type);
let mut sql = format!(
"ALTER TABLE {} ADD {} {}",
self.quote(table),
self.quote(&col.name),
oracle_type
);
if !col.nullable {
sql.push_str(" NOT NULL");
}
if let Some(default) = &col.default {
sql.push_str(&format!(" DEFAULT {}", default));
}
sql
}
TableChange::DropColumn(name) => {
format!(
"ALTER TABLE {} DROP COLUMN {}",
self.quote(table),
self.quote(name)
)
}
TableChange::ModifyColumn(col) => {
let oracle_type = map_to_oracle_type(&col.sql_type);
let mut sql = format!(
"ALTER TABLE {} MODIFY {} {}",
self.quote(table),
self.quote(&col.name),
oracle_type
);
if !col.nullable {
sql.push_str(" NOT NULL");
}
if let Some(default) = &col.default {
sql.push_str(&format!(" DEFAULT {}", default));
}
sql
}
TableChange::AddIndex(name, cols) => {
format!(
"CREATE INDEX {} ON {} ({})",
name,
self.quote(table),
cols.join(", ")
)
}
TableChange::DropIndex(name) => {
format!("DROP INDEX {}", name)
}
TableChange::AddForeignKey {
columns,
reference_table,
reference_columns,
} => {
format!(
"ALTER TABLE {} ADD CONSTRAINT fk_{}_{} FOREIGN KEY ({}) REFERENCES {} ({})",
self.quote(table),
table,
columns.join("_"),
columns.iter().map(|c| self.quote(c)).collect::<Vec<_>>().join(", "),
self.quote(reference_table),
reference_columns.iter().map(|c| self.quote(c)).collect::<Vec<_>>().join(", ")
)
}
})
.collect();
stmts.join("; ")
}
}
fn map_to_sqlserver_type(sql_type: &str) -> String {
let upper = sql_type.to_uppercase();
let trimmed = upper.trim();
if trimmed.starts_with("BIGINT") {
sql_type.to_string()
} else if matches!(trimmed, "INT" | "INTEGER") {
"INT".to_string()
} else if trimmed.starts_with("NVARCHAR") {
sql_type.to_string()
} else if trimmed.starts_with("VARCHAR") {
sql_type.replacen("VARCHAR", "NVARCHAR", 1)
} else if matches!(trimmed, "TEXT" | "MEDIUMTEXT" | "LONGTEXT" | "TINYTEXT") {
"NVARCHAR(MAX)".to_string()
} else if matches!(trimmed, "BOOLEAN" | "BOOL") {
"BIT".to_string()
} else {
sql_type.to_string()
}
}
pub struct SqlServerDialect;
impl Dialect for SqlServerDialect {
fn db_type(&self) -> DbType {
DbType::SqlServer
}
fn quote(&self, identifier: &str) -> String {
format!("[{}]", identifier.replace(']', "]]"))
}
fn escape_string(&self, s: &str) -> String {
let mut escaped = String::with_capacity(s.len() * 2);
for c in s.chars() {
match c {
'\'' => escaped.push_str("''"),
_ => escaped.push(c),
}
}
escaped
}
fn supports_returning(&self) -> bool {
true
}
fn build_pagination(&self, sql: &str, page: u64, limit: u64) -> String {
let offset = page.saturating_sub(1).saturating_mul(limit);
format!(
"{} OFFSET {} ROWS FETCH NEXT {} ROWS ONLY",
sql, offset, limit
)
}
fn json_type(&self) -> &'static str {
"NVARCHAR(MAX)"
}
fn json_extract(&self, column: &str, path: &str) -> String {
let normalized = if path.starts_with('$') {
path.to_string()
} else {
format!("$.{}", path)
};
format!(
"JSON_VALUE({}, '{}')",
column,
self.escape_string(&normalized)
)
}
fn full_text_search(&self, columns: &[&str], keyword: &str) -> String {
if columns.is_empty() {
return "0".to_string();
}
let escaped = self.escape_string(keyword);
let cols = columns.join(", ");
format!("CONTAINS({}, '{}')", cols, escaped)
}
fn bool_to_int(&self, expr: &str) -> String {
format!("(CASE WHEN {} THEN 1 ELSE 0 END)", expr)
}
fn concat(&self, parts: &[&str]) -> String {
if parts.is_empty() {
return "NULL".to_string();
}
format!("CONCAT({})", parts.join(", "))
}
fn supports_if_exists(&self) -> bool {
true
}
fn supports_if_not_exists(&self) -> bool {
true
}
fn auto_increment_keyword(&self) -> &'static str {
"IDENTITY(1,1)"
}
fn last_insert_id_sql(&self) -> Option<&'static str> {
Some("SCOPE_IDENTITY()")
}
fn build_create_table(&self, table: &str, columns: &[ColumnDef]) -> String {
let cols: Vec<String> = columns
.iter()
.map(|col| {
let sqlserver_type = map_to_sqlserver_type(&col.sql_type);
let mut sql = format!("{} {}", self.quote(&col.name), sqlserver_type);
if !col.nullable {
sql.push_str(" NOT NULL");
}
if let Some(default) = &col.default {
sql.push_str(&format!(" DEFAULT {}", default));
}
if col.auto_increment {
sql.push_str(&format!(" {}", self.auto_increment_keyword()));
}
if col.primary_key {
sql.push_str(" PRIMARY KEY");
}
sql
})
.collect();
format!("CREATE TABLE {} ({})", self.quote(table), cols.join(", "))
}
fn build_alter_table(&self, table: &str, changes: &[TableChange]) -> String {
let stmts: Vec<String> = changes
.iter()
.map(|change| match change {
TableChange::AddColumn(col) => {
let sqlserver_type = map_to_sqlserver_type(&col.sql_type);
let mut sql = format!(
"ALTER TABLE {} ADD {} {}",
self.quote(table),
self.quote(&col.name),
sqlserver_type
);
if !col.nullable {
sql.push_str(" NOT NULL");
}
if let Some(default) = &col.default {
sql.push_str(&format!(" DEFAULT {}", default));
}
sql
}
TableChange::DropColumn(name) => {
format!(
"ALTER TABLE {} DROP COLUMN {}",
self.quote(table),
self.quote(name)
)
}
TableChange::ModifyColumn(col) => {
let sqlserver_type = map_to_sqlserver_type(&col.sql_type);
let mut sql = format!(
"ALTER TABLE {} ALTER COLUMN {} {}",
self.quote(table),
self.quote(&col.name),
sqlserver_type
);
if !col.nullable {
sql.push_str(" NOT NULL");
}
if let Some(default) = &col.default {
sql.push_str(&format!(" DEFAULT {}", default));
}
sql
}
TableChange::AddIndex(name, cols) => {
format!(
"CREATE INDEX {} ON {} ({})",
name,
self.quote(table),
cols.join(", ")
)
}
TableChange::DropIndex(name) => {
format!("DROP INDEX {} ON {}", name, self.quote(table))
}
TableChange::AddForeignKey {
columns,
reference_table,
reference_columns,
} => {
format!(
"ALTER TABLE {} ADD CONSTRAINT fk_{}_{} FOREIGN KEY ({}) REFERENCES {} ({})",
self.quote(table),
table,
columns.join("_"),
columns.iter().map(|c| self.quote(c)).collect::<Vec<_>>().join(", "),
self.quote(reference_table),
reference_columns.iter().map(|c| self.quote(c)).collect::<Vec<_>>().join(", ")
)
}
})
.collect();
stmts.join("; ")
}
}
macro_rules! delegate_dialect_to {
($wrapper:ident, $base:ident, $db_type:expr) => {
pub struct $wrapper;
impl Dialect for $wrapper {
fn db_type(&self) -> DbType {
$db_type
}
fn quote(&self, identifier: &str) -> String {
$base.quote(identifier)
}
fn escape_string(&self, s: &str) -> String {
$base.escape_string(s)
}
fn supports_returning(&self) -> bool {
$base.supports_returning()
}
fn build_pagination(&self, sql: &str, page: u64, limit: u64) -> String {
$base.build_pagination(sql, page, limit)
}
fn json_type(&self) -> &'static str {
$base.json_type()
}
fn json_extract(&self, column: &str, path: &str) -> String {
$base.json_extract(column, path)
}
fn full_text_search(&self, columns: &[&str], keyword: &str) -> String {
$base.full_text_search(columns, keyword)
}
fn bool_to_int(&self, expr: &str) -> String {
$base.bool_to_int(expr)
}
fn concat(&self, parts: &[&str]) -> String {
$base.concat(parts)
}
fn supports_if_exists(&self) -> bool {
$base.supports_if_exists()
}
fn supports_if_not_exists(&self) -> bool {
$base.supports_if_not_exists()
}
fn auto_increment_keyword(&self) -> &'static str {
$base.auto_increment_keyword()
}
fn last_insert_id_sql(&self) -> Option<&'static str> {
$base.last_insert_id_sql()
}
fn build_create_table(&self, table: &str, columns: &[ColumnDef]) -> String {
$base.build_create_table(table, columns)
}
fn build_alter_table(&self, table: &str, changes: &[TableChange]) -> String {
$base.build_alter_table(table, changes)
}
fn build_drop_table(&self, table: &str, if_exists: bool) -> String {
$base.build_drop_table(table, if_exists)
}
}
};
}
delegate_dialect_to!(MariaDbDialect, MySqlDialect, DbType::MariaDB);
delegate_dialect_to!(TiDbDialect, MySqlDialect, DbType::TiDB);
delegate_dialect_to!(KingbaseDialect, PostgreSqlDialect, DbType::Kingbase);
delegate_dialect_to!(PolarDbDialect, PostgreSqlDialect, DbType::PolarDB);
delegate_dialect_to!(GaussDbDialect, PostgreSqlDialect, DbType::GaussDB);
delegate_dialect_to!(DamengDialect, OracleDialect, DbType::Dameng);
delegate_dialect_to!(SybaseDialect, SqlServerDialect, DbType::Sybase);
delegate_dialect_to!(GBaseDialect, SqlServerDialect, DbType::GBase);
pub struct ClickHouseDialect;
impl Dialect for ClickHouseDialect {
fn db_type(&self) -> DbType {
DbType::ClickHouse
}
fn quote(&self, identifier: &str) -> String {
format!("`{}`", identifier.replace('`', "``"))
}
fn escape_string(&self, s: &str) -> String {
let mut escaped = String::with_capacity(s.len() * 2);
for c in s.chars() {
match c {
'\'' => escaped.push_str("\\'"),
'\\' => escaped.push_str("\\\\"),
'\n' => escaped.push_str("\\n"),
'\r' => escaped.push_str("\\r"),
'\t' => escaped.push_str("\\t"),
_ => escaped.push(c),
}
}
escaped
}
fn supports_returning(&self) -> bool {
false
}
fn build_pagination(&self, sql: &str, page: u64, limit: u64) -> String {
let offset = page.saturating_sub(1).saturating_mul(limit);
format!("{} LIMIT {}, {}", sql, offset, limit)
}
fn json_type(&self) -> &'static str {
"String"
}
fn json_extract(&self, column: &str, path: &str) -> String {
let normalized = if path.starts_with('$') {
path.to_string()
} else {
format!("$.{}", path)
};
format!(
"JSONExtractString({}, '{}')",
column,
self.escape_string(&normalized)
)
}
fn full_text_search(&self, columns: &[&str], keyword: &str) -> String {
if columns.is_empty() {
return "0".to_string();
}
let escaped = self.escape_string(keyword);
let parts: Vec<String> = columns
.iter()
.map(|c| format!("position({}, '{}') > 0", c, escaped))
.collect();
parts.join(" OR ")
}
fn bool_to_int(&self, expr: &str) -> String {
format!("toUInt8({})", expr)
}
fn concat(&self, parts: &[&str]) -> String {
if parts.is_empty() {
return "''".to_string();
}
format!("concat({})", parts.join(", "))
}
fn supports_if_exists(&self) -> bool {
true
}
fn supports_if_not_exists(&self) -> bool {
true
}
fn auto_increment_keyword(&self) -> &'static str {
""
}
fn last_insert_id_sql(&self) -> Option<&'static str> {
None
}
fn build_create_table(&self, table: &str, columns: &[ColumnDef]) -> String {
let cols: Vec<String> = columns
.iter()
.map(|col| {
let ch_type = map_to_clickhouse_type(&col.sql_type);
let mut sql = format!("{} {}", self.quote(&col.name), ch_type);
if let Some(default) = &col.default {
sql.push_str(&format!(" DEFAULT {}", default));
}
if col.primary_key {
sql.push_str(" PRIMARY KEY");
}
sql
})
.collect();
format!(
"CREATE TABLE {} ({}) ENGINE = MergeTree()",
self.quote(table),
cols.join(", ")
)
}
fn build_alter_table(&self, table: &str, changes: &[TableChange]) -> String {
let stmts: Vec<String> = changes
.iter()
.map(|change| match change {
TableChange::AddColumn(col) => {
let ch_type = map_to_clickhouse_type(&col.sql_type);
format!(
"ALTER TABLE {} ADD COLUMN {} {}",
self.quote(table),
self.quote(&col.name),
ch_type
)
}
TableChange::DropColumn(name) => {
format!(
"ALTER TABLE {} DROP COLUMN {}",
self.quote(table),
self.quote(name)
)
}
TableChange::ModifyColumn(col) => {
let ch_type = map_to_clickhouse_type(&col.sql_type);
format!(
"ALTER TABLE {} MODIFY COLUMN {} {}",
self.quote(table),
self.quote(&col.name),
ch_type
)
}
TableChange::AddIndex(name, cols) => {
format!(
"ALTER TABLE {} ADD INDEX {} ({})",
self.quote(table),
name,
cols.join(", ")
)
}
TableChange::DropIndex(name) => {
format!("ALTER TABLE {} DROP INDEX {}", self.quote(table), name)
}
TableChange::AddForeignKey { .. } => {
String::new()
}
})
.filter(|s| !s.is_empty())
.collect();
stmts.join("; ")
}
}
fn map_to_clickhouse_type(sql_type: &str) -> String {
let upper = sql_type.to_uppercase();
let trimmed = upper.trim();
if trimmed.starts_with("BIGINT") {
"Int64".to_string()
} else if matches!(trimmed, "INT" | "INTEGER") {
"Int32".to_string()
} else if matches!(trimmed, "TINYINT" | "SMALLINT") {
"Int16".to_string()
} else if trimmed.starts_with("VARCHAR")
|| trimmed.starts_with("CHAR")
|| matches!(trimmed, "TEXT" | "MEDIUMTEXT" | "LONGTEXT" | "TINYTEXT")
{
"String".to_string()
} else if matches!(trimmed, "BOOLEAN" | "BOOL") {
"UInt8".to_string()
} else if matches!(trimmed, "FLOAT" | "REAL") {
"Float32".to_string()
} else if matches!(trimmed, "DOUBLE" | "DOUBLE PRECISION") {
"Float64".to_string()
} else if matches!(trimmed, "DATETIME" | "TIMESTAMP") {
"DateTime".to_string()
} else if matches!(trimmed, "DATE") {
"Date".to_string()
} else if trimmed.starts_with("DECIMAL") || trimmed.starts_with("NUMERIC") {
"Decimal(38, 4)".to_string()
} else {
sql_type.to_string()
}
}
pub struct Db2Dialect;
impl Dialect for Db2Dialect {
fn db_type(&self) -> DbType {
DbType::Db2
}
fn quote(&self, identifier: &str) -> String {
format!("\"{}\"", identifier.replace('"', "\"\""))
}
fn escape_string(&self, s: &str) -> String {
let mut escaped = String::with_capacity(s.len() * 2);
for c in s.chars() {
match c {
'\'' => escaped.push_str("''"),
_ => escaped.push(c),
}
}
escaped
}
fn supports_returning(&self) -> bool {
false
}
fn build_pagination(&self, sql: &str, page: u64, limit: u64) -> String {
let offset = page.saturating_sub(1).saturating_mul(limit);
format!(
"{} OFFSET {} ROWS FETCH NEXT {} ROWS ONLY",
sql, offset, limit
)
}
fn json_type(&self) -> &'static str {
"JSON"
}
fn json_extract(&self, column: &str, path: &str) -> String {
let normalized = if path.starts_with('$') {
path.to_string()
} else {
format!("$.{}", path)
};
format!(
"JSON_VALUE({}, '{}')",
column,
self.escape_string(&normalized)
)
}
fn full_text_search(&self, columns: &[&str], keyword: &str) -> String {
if columns.is_empty() {
return "0".to_string();
}
let escaped = self.escape_string(keyword);
let parts: Vec<String> = columns
.iter()
.map(|c| format!("CONTAINS({}, '{}') > 0", c, escaped))
.collect();
parts.join(" OR ")
}
fn bool_to_int(&self, expr: &str) -> String {
format!("(CASE WHEN {} THEN 1 ELSE 0 END)", expr)
}
fn concat(&self, parts: &[&str]) -> String {
if parts.is_empty() {
return "''".to_string();
}
parts.join(" || ")
}
fn supports_if_exists(&self) -> bool {
false
}
fn supports_if_not_exists(&self) -> bool {
false
}
fn auto_increment_keyword(&self) -> &'static str {
"GENERATED ALWAYS AS IDENTITY"
}
fn last_insert_id_sql(&self) -> Option<&'static str> {
Some("SELECT IDENTITY_VAL_LOCAL() FROM SYSIBM.SYSDUMMY1")
}
fn build_create_table(&self, table: &str, columns: &[ColumnDef]) -> String {
let cols: Vec<String> = columns
.iter()
.map(|col| {
let db2_type = map_to_db2_type(&col.sql_type);
let mut sql = format!("{} {}", self.quote(&col.name), db2_type);
if !col.nullable && !col.auto_increment {
sql.push_str(" NOT NULL");
}
if let Some(default) = &col.default {
sql.push_str(&format!(" DEFAULT {}", default));
}
if col.auto_increment {
sql.push_str(&format!(" {}", self.auto_increment_keyword()));
}
if col.primary_key {
sql.push_str(" PRIMARY KEY");
}
sql
})
.collect();
format!("CREATE TABLE {} ({})", self.quote(table), cols.join(", "))
}
fn build_alter_table(&self, table: &str, changes: &[TableChange]) -> String {
let stmts: Vec<String> = changes
.iter()
.map(|change| match change {
TableChange::AddColumn(col) => {
let db2_type = map_to_db2_type(&col.sql_type);
let mut sql = format!(
"ALTER TABLE {} ADD COLUMN {} {}",
self.quote(table),
self.quote(&col.name),
db2_type
);
if !col.nullable {
sql.push_str(" NOT NULL");
}
if let Some(default) = &col.default {
sql.push_str(&format!(" DEFAULT {}", default));
}
sql
}
TableChange::DropColumn(name) => {
format!(
"ALTER TABLE {} DROP COLUMN {}",
self.quote(table),
self.quote(name)
)
}
TableChange::ModifyColumn(col) => {
let db2_type = map_to_db2_type(&col.sql_type);
format!(
"ALTER TABLE {} ALTER COLUMN {} SET DATA TYPE {}",
self.quote(table),
self.quote(&col.name),
db2_type
)
}
TableChange::AddIndex(name, cols) => {
format!(
"CREATE INDEX {} ON {} ({})",
name,
self.quote(table),
cols.join(", ")
)
}
TableChange::DropIndex(name) => {
format!("DROP INDEX {}", name)
}
TableChange::AddForeignKey {
columns,
reference_table,
reference_columns,
} => {
format!(
"ALTER TABLE {} ADD CONSTRAINT fk_{}_{} FOREIGN KEY ({}) REFERENCES {} ({})",
self.quote(table),
table,
columns.join("_"),
columns.iter().map(|c| self.quote(c)).collect::<Vec<_>>().join(", "),
self.quote(reference_table),
reference_columns.iter().map(|c| self.quote(c)).collect::<Vec<_>>().join(", ")
)
}
})
.collect();
stmts.join("; ")
}
fn build_drop_table(&self, table: &str, if_exists: bool) -> String {
let _ = if_exists;
format!("DROP TABLE {}", self.quote(table))
}
}
fn map_to_db2_type(sql_type: &str) -> String {
let upper = sql_type.to_uppercase();
let trimmed = upper.trim();
if trimmed.starts_with("BIGINT") {
"BIGINT".to_string()
} else if matches!(trimmed, "INT" | "INTEGER") {
"INTEGER".to_string()
} else if matches!(trimmed, "TINYINT" | "SMALLINT") {
"SMALLINT".to_string()
} else if trimmed.starts_with("VARCHAR") || trimmed.starts_with("CHAR") {
sql_type.to_string()
} else if matches!(trimmed, "TEXT" | "MEDIUMTEXT" | "LONGTEXT" | "TINYTEXT") {
"CLOB(2G)".to_string()
} else if matches!(trimmed, "BOOLEAN" | "BOOL") {
"SMALLINT".to_string()
} else if matches!(trimmed, "FLOAT" | "REAL") {
"REAL".to_string()
} else if matches!(trimmed, "DOUBLE" | "DOUBLE PRECISION") {
"DOUBLE".to_string()
} else if matches!(trimmed, "DATETIME" | "TIMESTAMP") {
"TIMESTAMP".to_string()
} else if matches!(trimmed, "DATE") {
"DATE".to_string()
} else {
sql_type.to_string()
}
}
pub fn get_dialect(db_type: DbType) -> Result<Box<dyn Dialect>, DbError> {
match db_type {
DbType::MySQL => Ok(Box::new(MySqlDialect)),
DbType::PostgreSQL => Ok(Box::new(PostgreSqlDialect)),
DbType::Sqlite => Ok(Box::new(SqliteDialect)),
DbType::Redis => Err(DbError::Unsupported(
"Redis does not support standard SQL dialect".to_string(),
)),
DbType::MongoDB => Err(DbError::Unsupported(
"MongoDB uses different query syntax".to_string(),
)),
DbType::ClickHouse => Ok(Box::new(ClickHouseDialect)),
DbType::Oracle => Ok(Box::new(OracleDialect)),
DbType::OceanBase => Ok(Box::new(MySqlDialect)),
DbType::SqlServer => Ok(Box::new(SqlServerDialect)),
DbType::VectorDb => Err(DbError::Unsupported(
"Vector databases have specific APIs".to_string(),
)),
DbType::PureJsDb => Err(DbError::Unsupported(
"PureJS database uses JavaScript".to_string(),
)),
DbType::Dameng => Ok(Box::new(DamengDialect)),
DbType::Kingbase => Ok(Box::new(KingbaseDialect)),
DbType::Db2 => Ok(Box::new(Db2Dialect)),
DbType::MariaDB => Ok(Box::new(MariaDbDialect)),
DbType::TiDB => Ok(Box::new(TiDbDialect)),
DbType::PolarDB => Ok(Box::new(PolarDbDialect)),
DbType::GaussDB => Ok(Box::new(GaussDbDialect)),
DbType::GBase => Ok(Box::new(GBaseDialect)),
DbType::Sybase => Ok(Box::new(SybaseDialect)),
}
}
impl fmt::Display for dyn Dialect {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
write!(f, "Dialect({})", self.db_type())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_mysql_quote() {
let dialect = MySqlDialect;
assert_eq!(dialect.quote("users"), "`users`");
assert_eq!(dialect.quote("user`id"), "`user``id`");
}
#[test]
fn test_mysql_escape() {
let dialect = MySqlDialect;
assert_eq!(dialect.escape_string("hello"), "hello");
assert_eq!(dialect.escape_string("it's"), "it\\'s");
assert_eq!(dialect.escape_string("line\nbreak"), "line\\nbreak");
}
#[test]
fn test_mysql_pagination() {
let dialect = MySqlDialect;
let sql = dialect.build_pagination("SELECT * FROM users", 2, 10);
assert_eq!(sql, "SELECT * FROM users LIMIT 10 OFFSET 10");
}
#[test]
fn test_postgres_quote() {
let dialect = PostgreSqlDialect;
assert_eq!(dialect.quote("users"), "\"users\"");
assert_eq!(dialect.quote("user\"id"), "\"user\"\"id\"");
}
#[test]
fn test_postgres_pagination() {
let dialect = PostgreSqlDialect;
let sql = dialect.build_pagination("SELECT * FROM users", 3, 20);
assert_eq!(sql, "SELECT * FROM users LIMIT 20 OFFSET 40");
}
#[test]
fn test_postgres_returning() {
let dialect = PostgreSqlDialect;
assert!(dialect.supports_returning());
}
#[test]
fn test_sqlite_quote() {
let dialect = SqliteDialect;
assert_eq!(dialect.quote("users"), "\"users\"");
assert_eq!(dialect.quote("user\"id"), "\"user\"\"id\"");
}
#[test]
fn test_sqlite_escape() {
let dialect = SqliteDialect;
assert_eq!(dialect.escape_string("hello"), "hello");
assert_eq!(dialect.escape_string("it's"), "it''s");
}
#[test]
fn test_get_dialect() {
let dialect = get_dialect(DbType::MySQL);
assert!(dialect.is_ok());
let dialect = get_dialect(DbType::Redis);
assert!(dialect.is_err());
}
#[test]
fn test_bool_to_int() {
let mysql = MySqlDialect;
assert_eq!(mysql.bool_to_int("active"), "IF(active, 1, 0)");
let pg = PostgreSqlDialect;
assert_eq!(
pg.bool_to_int("active"),
"(CASE WHEN active THEN 1 ELSE 0 END)"
);
}
#[test]
fn test_json_extract_with_path() {
let mysql = MySqlDialect;
let sql = mysql.json_extract("data", "$.user.name");
assert!(sql.contains("$.user.name"));
assert!(sql.contains("JSON_EXTRACT"));
let pg = PostgreSqlDialect;
let sql = pg.json_extract("data", "user.name");
assert!(sql.contains("#>>"));
let sqlite = SqliteDialect;
let sql = sqlite.json_extract("data", "$.user.name");
assert!(sql.contains("$.user.name"));
assert!(sql.contains("json_extract"));
}
#[test]
fn test_sqlite_full_text_search() {
let sqlite = SqliteDialect;
let sql = sqlite.full_text_search(&["title", "content"], "hello");
assert!(sql.contains("LIKE"));
assert!(sql.contains("title LIKE '%hello%'"));
assert!(sql.contains("content LIKE '%hello%'"));
assert!(sql.contains(" OR "));
assert_eq!(sqlite.full_text_search(&[], "hello"), "0");
let sql = sqlite.full_text_search(&["title"], "it's");
assert!(sql.contains("title LIKE '%it''s%'"));
}
#[test]
fn test_alter_table_modify_column() {
let mysql = MySqlDialect;
let col = ColumnDef {
name: "name".to_string(),
sql_type: "VARCHAR(255)".to_string(),
nullable: false,
default: None,
auto_increment: false,
primary_key: false,
};
let sql = mysql.build_alter_table("users", &[TableChange::ModifyColumn(col)]);
assert!(sql.contains("MODIFY COLUMN"));
let pg = PostgreSqlDialect;
let col = ColumnDef {
name: "name".to_string(),
sql_type: "VARCHAR(255)".to_string(),
nullable: false,
default: None,
auto_increment: false,
primary_key: false,
};
let sql = pg.build_alter_table("users", &[TableChange::ModifyColumn(col)]);
assert!(sql.contains("ALTER COLUMN"));
assert!(sql.contains("TYPE"));
}
#[test]
fn test_alter_table_add_foreign_key() {
let mysql = MySqlDialect;
let sql = mysql.build_alter_table(
"orders",
&[TableChange::AddForeignKey {
columns: vec!["user_id".to_string()],
reference_table: "users".to_string(),
reference_columns: vec!["id".to_string()],
}],
);
assert!(sql.contains("FOREIGN KEY"));
assert!(sql.contains("REFERENCES"));
let sqlite = SqliteDialect;
let sql = sqlite.build_alter_table(
"orders",
&[TableChange::AddForeignKey {
columns: vec!["user_id".to_string()],
reference_table: "users".to_string(),
reference_columns: vec!["id".to_string()],
}],
);
assert!(sql.starts_with("--"));
}
#[test]
fn test_sqlite_alter_table_add_column() {
let sqlite = SqliteDialect;
let col = ColumnDef {
name: "email".to_string(),
sql_type: "TEXT".to_string(),
nullable: true,
default: None,
auto_increment: false,
primary_key: false,
};
let sql = sqlite.build_alter_table("users", &[TableChange::AddColumn(col)]);
assert!(sql.contains("ADD COLUMN"));
assert!(sql.contains("email"));
}
#[test]
fn test_oracle_quote_and_escape() {
let dialect = OracleDialect;
assert_eq!(dialect.quote("users"), "\"users\"");
assert_eq!(dialect.quote("user\"id"), "\"user\"\"id\"");
assert_eq!(dialect.quote("column_name"), "\"column_name\"");
assert_eq!(dialect.escape_string("hello"), "hello");
assert_eq!(dialect.escape_string("it's"), "it''s");
assert_eq!(dialect.escape_string("O'Brien"), "O''Brien");
assert_eq!(dialect.escape_string("a'b'c"), "a''b''c");
assert_eq!(dialect.escape_string("path\\to"), "path\\to");
}
#[test]
fn test_oracle_pagination() {
let dialect = OracleDialect;
let sql = dialect.build_pagination("SELECT * FROM users", 1, 10);
assert_eq!(
sql,
"SELECT * FROM users OFFSET 0 ROWS FETCH NEXT 10 ROWS ONLY"
);
let sql = dialect.build_pagination("SELECT * FROM users", 3, 20);
assert_eq!(
sql,
"SELECT * FROM users OFFSET 40 ROWS FETCH NEXT 20 ROWS ONLY"
);
let sql = dialect.build_pagination("SELECT * FROM users", 0, 10);
assert_eq!(
sql,
"SELECT * FROM users OFFSET 0 ROWS FETCH NEXT 10 ROWS ONLY"
);
}
#[test]
fn test_oracle_json_extract() {
let dialect = OracleDialect;
let sql = dialect.json_extract("data", "$.user.name");
assert!(sql.contains("JSON_VALUE"));
assert!(sql.contains("$.user.name"));
assert!(sql.starts_with("JSON_VALUE(data, '$.user.name')"));
let sql = dialect.json_extract("data", "user.name");
assert!(sql.contains("$.user.name"));
assert!(sql.contains("JSON_VALUE"));
let sql = dialect.json_extract("data", "$.key's");
assert!(sql.contains("$.key''s"));
}
#[test]
fn test_oracle_create_table() {
let dialect = OracleDialect;
let columns = vec![
ColumnDef {
name: "id".to_string(),
sql_type: "BIGINT".to_string(),
nullable: false,
default: None,
auto_increment: true,
primary_key: true,
},
ColumnDef {
name: "name".to_string(),
sql_type: "VARCHAR(255)".to_string(),
nullable: false,
default: None,
auto_increment: false,
primary_key: false,
},
ColumnDef {
name: "bio".to_string(),
sql_type: "TEXT".to_string(),
nullable: true,
default: None,
auto_increment: false,
primary_key: false,
},
ColumnDef {
name: "is_active".to_string(),
sql_type: "BOOLEAN".to_string(),
nullable: false,
default: Some("1".to_string()),
auto_increment: false,
primary_key: false,
},
];
let sql = dialect.build_create_table("users", &columns);
assert!(
sql.contains("NUMBER(19)"),
"BIGINT should map to NUMBER(19): {}",
sql
);
assert!(
sql.contains("VARCHAR2(255)"),
"VARCHAR should map to VARCHAR2: {}",
sql
);
assert!(sql.contains("CLOB"), "TEXT should map to CLOB: {}", sql);
assert!(
sql.contains("NUMBER(1)"),
"BOOLEAN should map to NUMBER(1): {}",
sql
);
assert!(sql.contains("GENERATED BY DEFAULT AS IDENTITY"));
assert!(sql.contains("PRIMARY KEY"));
assert!(sql.contains("NOT NULL"));
assert!(sql.contains("DEFAULT 1"));
assert!(sql.contains("\"users\""));
assert!(sql.contains("\"id\""));
}
#[test]
fn test_oracle_bool_to_int_and_concat() {
let dialect = OracleDialect;
assert_eq!(
dialect.bool_to_int("active"),
"(CASE WHEN active THEN 1 ELSE 0 END)"
);
assert_eq!(
dialect.bool_to_int("x > 0"),
"(CASE WHEN x > 0 THEN 1 ELSE 0 END)"
);
assert_eq!(dialect.concat(&["a", "b", "c"]), "a || b || c");
assert_eq!(
dialect.concat(&["first_name", "last_name"]),
"first_name || last_name"
);
assert_eq!(dialect.concat(&[]), "NULL");
}
#[test]
fn test_oracle_misc_dialect_methods() {
let dialect = OracleDialect;
assert_eq!(dialect.db_type(), DbType::Oracle);
assert!(dialect.supports_returning());
assert!(dialect.supports_if_exists());
assert!(dialect.supports_if_not_exists());
assert_eq!(
dialect.auto_increment_keyword(),
"GENERATED BY DEFAULT AS IDENTITY"
);
assert_eq!(dialect.last_insert_id_sql(), None);
assert_eq!(dialect.json_type(), "JSON");
}
#[test]
fn test_oracle_get_dialect() {
let dialect = get_dialect(DbType::Oracle);
assert!(dialect.is_ok(), "Oracle dialect should be available");
let dialect = dialect.unwrap();
assert_eq!(dialect.db_type(), DbType::Oracle);
assert_eq!(dialect.quote("users"), "\"users\"");
assert!(dialect.supports_returning());
assert_eq!(dialect.last_insert_id_sql(), None);
}
#[test]
fn test_oracle_drop_table() {
let dialect = OracleDialect;
let sql = dialect.build_drop_table("users", true);
assert_eq!(sql, "DROP TABLE IF EXISTS \"users\"");
let sql = dialect.build_drop_table("users", false);
assert_eq!(sql, "DROP TABLE \"users\"");
}
#[test]
fn test_oracle_alter_table() {
let dialect = OracleDialect;
let col = ColumnDef {
name: "name".to_string(),
sql_type: "VARCHAR(255)".to_string(),
nullable: false,
default: None,
auto_increment: false,
primary_key: false,
};
let sql = dialect.build_alter_table("users", &[TableChange::ModifyColumn(col)]);
assert!(sql.contains("MODIFY"));
assert!(sql.contains("VARCHAR2(255)"));
assert!(!sql.contains("MODIFY COLUMN"));
let col = ColumnDef {
name: "email".to_string(),
sql_type: "VARCHAR(255)".to_string(),
nullable: true,
default: None,
auto_increment: false,
primary_key: false,
};
let sql = dialect.build_alter_table("users", &[TableChange::AddColumn(col)]);
assert!(sql.contains("ADD \"email\""));
assert!(sql.contains("VARCHAR2(255)"));
let sql =
dialect.build_alter_table("users", &[TableChange::DropColumn("email".to_string())]);
assert!(sql.contains("DROP COLUMN"));
assert!(sql.contains("\"email\""));
}
#[test]
fn test_sqlite_concat_handles_null() {
let sqlite = SqliteDialect;
let sql = sqlite.concat(&["a", "b"]);
assert_eq!(sql, "COALESCE(a, '') || COALESCE(b, '')");
let sql = sqlite.concat(&["a"]);
assert_eq!(sql, "COALESCE(a, '')");
assert_eq!(sqlite.concat(&[]), "NULL");
}
#[test]
fn test_sqlserver_quote_and_escape() {
let dialect = SqlServerDialect;
assert_eq!(dialect.quote("users"), "[users]");
assert_eq!(dialect.quote("col]name"), "[col]]name]");
assert_eq!(dialect.escape_string("hello"), "hello");
assert_eq!(dialect.escape_string("it's"), "it''s");
assert_eq!(dialect.escape_string("O'Brien"), "O''Brien");
assert_eq!(dialect.escape_string("path\\to"), "path\\to");
}
#[test]
fn test_sqlserver_pagination() {
let dialect = SqlServerDialect;
let sql = dialect.build_pagination("SELECT * FROM users", 1, 10);
assert_eq!(
sql,
"SELECT * FROM users OFFSET 0 ROWS FETCH NEXT 10 ROWS ONLY"
);
let sql = dialect.build_pagination("SELECT * FROM users", 3, 20);
assert_eq!(
sql,
"SELECT * FROM users OFFSET 40 ROWS FETCH NEXT 20 ROWS ONLY"
);
let sql = dialect.build_pagination("SELECT * FROM users", 0, 10);
assert_eq!(
sql,
"SELECT * FROM users OFFSET 0 ROWS FETCH NEXT 10 ROWS ONLY"
);
}
#[test]
fn test_sqlserver_misc_dialect_methods() {
let dialect = SqlServerDialect;
assert_eq!(dialect.db_type(), DbType::SqlServer);
assert!(dialect.supports_returning());
assert!(dialect.supports_if_exists());
assert!(dialect.supports_if_not_exists());
assert_eq!(dialect.auto_increment_keyword(), "IDENTITY(1,1)");
assert_eq!(dialect.last_insert_id_sql(), Some("SCOPE_IDENTITY()"));
assert_eq!(dialect.json_type(), "NVARCHAR(MAX)");
}
#[test]
fn test_sqlserver_json_extract() {
let dialect = SqlServerDialect;
let sql = dialect.json_extract("data", "$.user.name");
assert!(sql.starts_with("JSON_VALUE(data, '$.user.name')"));
let sql = dialect.json_extract("data", "user.name");
assert!(sql.contains("$.user.name"));
assert!(sql.contains("JSON_VALUE"));
let sql = dialect.json_extract("data", "$.key's");
assert!(sql.contains("$.key''s"));
}
#[test]
fn test_sqlserver_full_text_search() {
let dialect = SqlServerDialect;
let sql = dialect.full_text_search(&["title", "content"], "hello");
assert!(sql.starts_with("CONTAINS(title, content, 'hello')"));
assert_eq!(dialect.full_text_search(&[], "hello"), "0");
let sql = dialect.full_text_search(&["title"], "it's");
assert!(sql.contains("it''s"));
}
#[test]
fn test_sqlserver_bool_to_int_and_concat() {
let dialect = SqlServerDialect;
assert_eq!(
dialect.bool_to_int("active"),
"(CASE WHEN active THEN 1 ELSE 0 END)"
);
assert_eq!(dialect.concat(&["a", "b", "c"]), "CONCAT(a, b, c)");
assert_eq!(dialect.concat(&[]), "NULL");
}
#[test]
fn test_sqlserver_create_table() {
let dialect = SqlServerDialect;
let columns = vec![
ColumnDef {
name: "id".to_string(),
sql_type: "BIGINT".to_string(),
nullable: false,
default: None,
auto_increment: true,
primary_key: true,
},
ColumnDef {
name: "name".to_string(),
sql_type: "VARCHAR(255)".to_string(),
nullable: false,
default: None,
auto_increment: false,
primary_key: false,
},
ColumnDef {
name: "bio".to_string(),
sql_type: "TEXT".to_string(),
nullable: true,
default: None,
auto_increment: false,
primary_key: false,
},
ColumnDef {
name: "is_active".to_string(),
sql_type: "BOOLEAN".to_string(),
nullable: false,
default: Some("1".to_string()),
auto_increment: false,
primary_key: false,
},
];
let sql = dialect.build_create_table("users", &columns);
assert!(sql.contains("[users]"));
assert!(sql.contains("[id]"));
assert!(sql.contains("IDENTITY(1,1)"));
assert!(
sql.contains("NVARCHAR(255)"),
"VARCHAR should map to NVARCHAR: {}",
sql
);
assert!(
sql.contains("NVARCHAR(MAX)"),
"TEXT should map to NVARCHAR(MAX): {}",
sql
);
assert!(sql.contains("BIT"), "BOOLEAN should map to BIT: {}", sql);
assert!(sql.contains("PRIMARY KEY"));
assert!(sql.contains("NOT NULL"));
assert!(sql.contains("DEFAULT 1"));
}
#[test]
fn test_sqlserver_drop_table() {
let dialect = SqlServerDialect;
assert_eq!(
dialect.build_drop_table("users", true),
"DROP TABLE IF EXISTS [users]"
);
assert_eq!(
dialect.build_drop_table("users", false),
"DROP TABLE [users]"
);
}
#[test]
fn test_sqlserver_alter_table() {
let dialect = SqlServerDialect;
let col = ColumnDef {
name: "name".to_string(),
sql_type: "VARCHAR(255)".to_string(),
nullable: false,
default: None,
auto_increment: false,
primary_key: false,
};
let sql = dialect.build_alter_table("users", &[TableChange::ModifyColumn(col)]);
assert!(sql.contains("ALTER COLUMN"));
assert!(sql.contains("NVARCHAR(255)"));
assert!(!sql.contains("MODIFY"));
let col = ColumnDef {
name: "email".to_string(),
sql_type: "VARCHAR(255)".to_string(),
nullable: true,
default: None,
auto_increment: false,
primary_key: false,
};
let sql = dialect.build_alter_table("users", &[TableChange::AddColumn(col)]);
assert!(sql.contains("ADD [email]"));
assert!(sql.contains("NVARCHAR(255)"));
let sql =
dialect.build_alter_table("users", &[TableChange::DropColumn("email".to_string())]);
assert!(sql.contains("DROP COLUMN"));
assert!(sql.contains("[email]"));
let sql =
dialect.build_alter_table("users", &[TableChange::DropIndex("idx_name".to_string())]);
assert!(sql.contains("DROP INDEX idx_name ON [users]"));
}
#[test]
fn test_sqlserver_get_dialect() {
let dialect = get_dialect(DbType::SqlServer);
assert!(dialect.is_ok(), "SqlServer dialect should be available");
let dialect = dialect.unwrap();
assert_eq!(dialect.db_type(), DbType::SqlServer);
assert_eq!(dialect.quote("users"), "[users]");
assert_eq!(dialect.last_insert_id_sql(), Some("SCOPE_IDENTITY()"));
assert_eq!(dialect.auto_increment_keyword(), "IDENTITY(1,1)");
}
#[test]
fn test_clickhouse_get_dialect_unsupported() {
let dialect = get_dialect(DbType::ClickHouse);
assert!(dialect.is_ok(), "ClickHouse should be supported");
let dialect = dialect.unwrap();
assert_eq!(dialect.db_type(), DbType::ClickHouse);
assert_eq!(dialect.quote("users"), "`users`");
assert!(!dialect.supports_returning());
let sql = dialect.build_pagination("SELECT * FROM t", 2, 10);
assert_eq!(sql, "SELECT * FROM t LIMIT 10, 10");
assert_eq!(dialect.auto_increment_keyword(), "");
}
#[test]
fn test_get_dialect_all_supported_types() {
assert!(get_dialect(DbType::MySQL).is_ok());
assert!(get_dialect(DbType::PostgreSQL).is_ok());
assert!(get_dialect(DbType::Sqlite).is_ok());
assert!(get_dialect(DbType::Oracle).is_ok());
assert!(get_dialect(DbType::SqlServer).is_ok());
assert!(get_dialect(DbType::OceanBase).is_ok());
assert!(get_dialect(DbType::ClickHouse).is_ok());
assert!(get_dialect(DbType::Dameng).is_ok());
assert!(get_dialect(DbType::Kingbase).is_ok());
assert!(get_dialect(DbType::Db2).is_ok());
assert!(get_dialect(DbType::MariaDB).is_ok());
assert!(get_dialect(DbType::TiDB).is_ok());
assert!(get_dialect(DbType::PolarDB).is_ok());
assert!(get_dialect(DbType::GaussDB).is_ok());
assert!(get_dialect(DbType::GBase).is_ok());
assert!(get_dialect(DbType::Sybase).is_ok());
assert!(get_dialect(DbType::Redis).is_err());
assert!(get_dialect(DbType::MongoDB).is_err());
assert!(get_dialect(DbType::VectorDb).is_err());
assert!(get_dialect(DbType::PureJsDb).is_err());
}
#[test]
fn test_mariadb_dialect() {
let dialect = get_dialect(DbType::MariaDB).unwrap();
assert_eq!(dialect.db_type(), DbType::MariaDB);
assert_eq!(dialect.quote("users"), "`users`");
assert_eq!(dialect.escape_string("it's"), "it\\'s");
assert_eq!(dialect.auto_increment_keyword(), "AUTO_INCREMENT");
assert!(!dialect.supports_returning());
}
#[test]
fn test_tidb_dialect() {
let dialect = get_dialect(DbType::TiDB).unwrap();
assert_eq!(dialect.db_type(), DbType::TiDB);
assert_eq!(dialect.quote("users"), "`users`");
assert_eq!(dialect.escape_string("it's"), "it\\'s");
assert_eq!(dialect.auto_increment_keyword(), "AUTO_INCREMENT");
}
#[test]
fn test_dameng_dialect() {
let dialect = get_dialect(DbType::Dameng).unwrap();
assert_eq!(dialect.db_type(), DbType::Dameng);
assert_eq!(dialect.quote("users"), "\"users\"");
assert_eq!(dialect.escape_string("it's"), "it''s");
assert_eq!(
dialect.auto_increment_keyword(),
"GENERATED BY DEFAULT AS IDENTITY"
);
assert!(dialect.supports_returning());
}
#[test]
fn test_kingbase_dialect() {
let dialect = get_dialect(DbType::Kingbase).unwrap();
assert_eq!(dialect.db_type(), DbType::Kingbase);
assert_eq!(dialect.quote("users"), "\"users\"");
assert_eq!(dialect.escape_string("it's"), "it''s");
assert!(dialect.supports_returning());
assert_eq!(
dialect.auto_increment_keyword(),
"GENERATED BY DEFAULT AS IDENTITY"
);
}
#[test]
fn test_polardb_dialect() {
let dialect = get_dialect(DbType::PolarDB).unwrap();
assert_eq!(dialect.db_type(), DbType::PolarDB);
assert_eq!(dialect.quote("users"), "\"users\"");
assert!(dialect.supports_returning());
}
#[test]
fn test_gaussdb_dialect() {
let dialect = get_dialect(DbType::GaussDB).unwrap();
assert_eq!(dialect.db_type(), DbType::GaussDB);
assert_eq!(dialect.quote("users"), "\"users\"");
assert!(dialect.supports_returning());
}
#[test]
fn test_gbase_dialect() {
let dialect = get_dialect(DbType::GBase).unwrap();
assert_eq!(dialect.db_type(), DbType::GBase);
assert_eq!(dialect.quote("users"), "[users]");
}
#[test]
fn test_sybase_dialect() {
let dialect = get_dialect(DbType::Sybase).unwrap();
assert_eq!(dialect.db_type(), DbType::Sybase);
assert_eq!(dialect.quote("users"), "[users]");
}
#[test]
fn test_db2_dialect_basic() {
let dialect = get_dialect(DbType::Db2).unwrap();
assert_eq!(dialect.db_type(), DbType::Db2);
assert_eq!(dialect.quote("users"), "\"users\"");
assert_eq!(dialect.escape_string("it's"), "it''s");
assert_eq!(
dialect.auto_increment_keyword(),
"GENERATED ALWAYS AS IDENTITY"
);
assert!(!dialect.supports_if_exists());
assert!(!dialect.supports_if_not_exists());
assert!(!dialect.supports_returning());
}
#[test]
fn test_db2_pagination() {
let dialect = Db2Dialect;
let sql = dialect.build_pagination("SELECT * FROM users", 2, 10);
assert_eq!(
sql,
"SELECT * FROM users OFFSET 10 ROWS FETCH NEXT 10 ROWS ONLY"
);
}
#[test]
fn test_db2_last_insert_id() {
let dialect = Db2Dialect;
assert_eq!(
dialect.last_insert_id_sql(),
Some("SELECT IDENTITY_VAL_LOCAL() FROM SYSIBM.SYSDUMMY1")
);
}
#[test]
fn test_db2_concat() {
let dialect = Db2Dialect;
assert_eq!(dialect.concat(&["a", "b", "c"]), "a || b || c");
assert_eq!(dialect.concat(&[]), "''");
}
#[test]
fn test_db2_create_table() {
let dialect = Db2Dialect;
let cols = vec![ColumnDef {
name: "id".to_string(),
sql_type: "BIGINT".to_string(),
nullable: false,
default: None,
auto_increment: true,
primary_key: true,
}];
let sql = dialect.build_create_table("users", &cols);
assert!(sql.contains("\"id\" BIGINT"));
assert!(sql.contains("GENERATED ALWAYS AS IDENTITY"));
assert!(sql.contains("PRIMARY KEY"));
}
#[test]
fn test_db2_type_mapping() {
assert_eq!(map_to_db2_type("BIGINT"), "BIGINT");
assert_eq!(map_to_db2_type("INT"), "INTEGER");
assert_eq!(map_to_db2_type("INTEGER"), "INTEGER");
assert_eq!(map_to_db2_type("TINYINT"), "SMALLINT");
assert_eq!(map_to_db2_type("SMALLINT"), "SMALLINT");
assert_eq!(map_to_db2_type("TEXT"), "CLOB(2G)");
assert_eq!(map_to_db2_type("LONGTEXT"), "CLOB(2G)");
assert_eq!(map_to_db2_type("BOOLEAN"), "SMALLINT");
assert_eq!(map_to_db2_type("BOOL"), "SMALLINT");
assert_eq!(map_to_db2_type("DATETIME"), "TIMESTAMP");
assert_eq!(map_to_db2_type("TIMESTAMP"), "TIMESTAMP");
assert_eq!(map_to_db2_type("DATE"), "DATE");
assert_eq!(map_to_db2_type("VARCHAR(255)"), "VARCHAR(255)");
}
#[test]
fn test_clickhouse_dialect_basic() {
let dialect = get_dialect(DbType::ClickHouse).unwrap();
assert_eq!(dialect.db_type(), DbType::ClickHouse);
assert_eq!(dialect.quote("users"), "`users`");
assert_eq!(dialect.escape_string("it's"), "it\\'s");
assert!(!dialect.supports_returning());
assert_eq!(dialect.auto_increment_keyword(), "");
assert!(dialect.supports_if_exists());
assert!(dialect.supports_if_not_exists());
}
#[test]
fn test_clickhouse_type_mapping() {
assert_eq!(map_to_clickhouse_type("BIGINT"), "Int64");
assert_eq!(map_to_clickhouse_type("INT"), "Int32");
assert_eq!(map_to_clickhouse_type("INTEGER"), "Int32");
assert_eq!(map_to_clickhouse_type("TINYINT"), "Int16");
assert_eq!(map_to_clickhouse_type("SMALLINT"), "Int16");
assert_eq!(map_to_clickhouse_type("VARCHAR(255)"), "String");
assert_eq!(map_to_clickhouse_type("TEXT"), "String");
assert_eq!(map_to_clickhouse_type("BOOLEAN"), "UInt8");
assert_eq!(map_to_clickhouse_type("BOOL"), "UInt8");
assert_eq!(map_to_clickhouse_type("FLOAT"), "Float32");
assert_eq!(map_to_clickhouse_type("DOUBLE"), "Float64");
assert_eq!(map_to_clickhouse_type("DATETIME"), "DateTime");
assert_eq!(map_to_clickhouse_type("TIMESTAMP"), "DateTime");
assert_eq!(map_to_clickhouse_type("DATE"), "Date");
}
#[test]
fn test_clickhouse_create_table() {
let dialect = ClickHouseDialect;
let cols = vec![ColumnDef {
name: "id".to_string(),
sql_type: "BIGINT".to_string(),
nullable: false,
default: None,
auto_increment: false, primary_key: true,
}];
let sql = dialect.build_create_table("users", &cols);
assert!(
sql.contains("ENGINE = MergeTree()"),
"ClickHouse CREATE TABLE 必须指定 ENGINE: {}",
sql
);
assert!(sql.contains("`id` Int64"));
assert!(sql.contains("PRIMARY KEY"));
}
#[test]
fn test_clickhouse_json_extract() {
let dialect = ClickHouseDialect;
let sql = dialect.json_extract("data", "$.name");
assert!(
sql.contains("JSONExtractString"),
"ClickHouse 应使用 JSONExtractString: {}",
sql
);
}
#[test]
fn test_clickhouse_concat() {
let dialect = ClickHouseDialect;
assert_eq!(dialect.concat(&["a", "b", "c"]), "concat(a, b, c)");
assert_eq!(dialect.concat(&[]), "''");
}
#[test]
fn test_db_type_dameng_str() {
assert_eq!(DbType::Dameng.as_str(), "dameng");
assert_eq!(DbType::from_str("dameng"), Some(DbType::Dameng));
assert_eq!(DbType::from_str("DM"), Some(DbType::Dameng));
assert_eq!(DbType::from_str("dm8"), Some(DbType::Dameng));
assert_eq!(DbType::Dameng.default_port(), 5236);
}
#[test]
fn test_db_type_kingbase_str() {
assert_eq!(DbType::Kingbase.as_str(), "kingbase");
assert_eq!(DbType::from_str("kingbase"), Some(DbType::Kingbase));
assert_eq!(DbType::Kingbase.default_port(), 54321);
}
#[test]
fn test_db_type_db2_str() {
assert_eq!(DbType::Db2.as_str(), "db2");
assert_eq!(DbType::from_str("db2"), Some(DbType::Db2));
assert_eq!(DbType::Db2.default_port(), 50000);
}
#[test]
fn test_db_type_mariadb_str() {
assert_eq!(DbType::MariaDB.as_str(), "mariadb");
assert_eq!(DbType::from_str("mariadb"), Some(DbType::MariaDB));
assert_eq!(DbType::MariaDB.default_port(), 3306);
}
#[test]
fn test_db_type_tidb_str() {
assert_eq!(DbType::TiDB.as_str(), "tidb");
assert_eq!(DbType::from_str("tidb"), Some(DbType::TiDB));
assert_eq!(DbType::TiDB.default_port(), 4000);
}
#[test]
fn test_db_type_polardb_str() {
assert_eq!(DbType::PolarDB.as_str(), "polardb");
assert_eq!(DbType::from_str("polardb"), Some(DbType::PolarDB));
assert_eq!(DbType::PolarDB.default_port(), 5432);
}
#[test]
fn test_db_type_gaussdb_str() {
assert_eq!(DbType::GaussDB.as_str(), "gaussdb");
assert_eq!(DbType::from_str("gaussdb"), Some(DbType::GaussDB));
assert_eq!(DbType::GaussDB.default_port(), 25308);
}
#[test]
fn test_db_type_gbase_str() {
assert_eq!(DbType::GBase.as_str(), "gbase");
assert_eq!(DbType::from_str("gbase"), Some(DbType::GBase));
assert_eq!(DbType::GBase.default_port(), 9088);
}
#[test]
fn test_db_type_sybase_str() {
assert_eq!(DbType::Sybase.as_str(), "sybase");
assert_eq!(DbType::from_str("sybase"), Some(DbType::Sybase));
assert_eq!(DbType::Sybase.default_port(), 5000);
}
#[test]
fn test_db_type_family_classification() {
assert!(DbType::MySQL.is_mysql_family());
assert!(DbType::MariaDB.is_mysql_family());
assert!(DbType::TiDB.is_mysql_family());
assert!(DbType::OceanBase.is_mysql_family());
assert!(!DbType::PostgreSQL.is_mysql_family());
assert!(DbType::PostgreSQL.is_postgres_family());
assert!(DbType::Kingbase.is_postgres_family());
assert!(DbType::GaussDB.is_postgres_family());
assert!(!DbType::MySQL.is_postgres_family());
assert!(DbType::Oracle.is_oracle_family());
assert!(DbType::Dameng.is_oracle_family());
assert!(!DbType::MySQL.is_oracle_family());
}
#[test]
fn test_db_type_supports_stored_procedure_extended() {
assert!(DbType::Dameng.supports_stored_procedure());
assert!(DbType::Kingbase.supports_stored_procedure());
assert!(DbType::Db2.supports_stored_procedure());
assert!(DbType::MariaDB.supports_stored_procedure());
assert!(DbType::TiDB.supports_stored_procedure());
assert!(DbType::PolarDB.supports_stored_procedure());
assert!(DbType::GaussDB.supports_stored_procedure());
assert!(DbType::GBase.supports_stored_procedure());
assert!(DbType::Sybase.supports_stored_procedure());
}
#[test]
fn test_l4_max_identifier_len_constant() {
assert_eq!(MAX_IDENTIFIER_LEN, 63);
}
#[test]
fn test_l4_quote_checked_valid_identifier() {
let dialect = MySqlDialect;
assert_eq!(dialect.quote_checked("users").unwrap(), "`users`");
assert_eq!(dialect.quote_checked("user_id").unwrap(), "`user_id`");
let name_63 = "a".repeat(63);
assert!(dialect.quote_checked(&name_63).is_ok());
}
#[test]
fn test_l4_quote_checked_rejects_too_long() {
let dialect = MySqlDialect;
let long_name = "a".repeat(64); let result = dialect.quote_checked(&long_name);
assert!(result.is_err());
match result {
Err(DbError::InvalidInput(msg)) => {
assert!(
msg.contains("too long"),
"expected 'too long' error, got: {}",
msg
);
}
_ => panic!("Expected DbError::InvalidInput"),
}
}
#[test]
fn test_l4_quote_checked_rejects_empty() {
let dialect = MySqlDialect;
let result = dialect.quote_checked("");
assert!(result.is_err());
}
#[test]
fn test_l4_quote_checked_rejects_sql_injection() {
let dialect = MySqlDialect;
assert!(dialect.quote_checked("users; DROP TABLE users").is_err());
assert!(dialect.quote_checked("user'name").is_err());
assert!(dialect.quote_checked("user name").is_err());
assert!(dialect.quote_checked("1users").is_err());
assert!(dialect.quote_checked("schema.table").is_err());
}
#[test]
fn test_l4_quote_checked_postgres() {
let dialect = PostgreSqlDialect;
assert_eq!(dialect.quote_checked("users").unwrap(), "\"users\"");
assert!(dialect.quote_checked(&"a".repeat(64)).is_err());
}
#[test]
fn test_l4_quote_checked_sqlite() {
let dialect = SqliteDialect;
assert_eq!(dialect.quote_checked("users").unwrap(), "\"users\"");
assert!(dialect.quote_checked(&"a".repeat(64)).is_err());
}
#[test]
fn test_l4_quote_checked_oracle() {
let dialect = OracleDialect;
assert_eq!(dialect.quote_checked("users").unwrap(), "\"users\"");
assert!(dialect.quote_checked(&"a".repeat(64)).is_err());
}
#[test]
fn test_l4_quote_checked_sql_server() {
let dialect = SqlServerDialect;
assert_eq!(dialect.quote_checked("users").unwrap(), "[users]");
assert!(dialect.quote_checked(&"a".repeat(64)).is_err());
}
}