use std::collections::HashSet;
use torm::db::db_types::SqlValue;
use torm::db::database::{Database, DbError};
#[derive(Debug, Clone)]
struct Config {
p_host: String,
p_port: u16,
p_db: String,
p_user: String,
p_pass: String,
m_host: String,
m_port: u16,
m_db: String,
m_user: String,
m_pass: String,
tables: Option<Vec<String>>, batch_size: usize, create_only: bool, data_only: bool, }
impl Default for Config {
fn default() -> Self {
Self {
p_host: "localhost".to_string(),
p_port: 5432,
p_db: "mydb".to_string(),
p_user: "postgres".to_string(),
p_pass: "".to_string(),
m_host: "localhost".to_string(),
m_port: 3306,
m_db: "mydb".to_string(),
m_user: "root".to_string(),
m_pass: "".to_string(),
tables: None,
batch_size: 1000,
create_only: false,
data_only: false,
}
}
}
#[derive(Debug, Clone)]
struct PgColumn {
name: String,
data_type: String, udt_name: String, is_nullable: bool,
column_default: Option<String>,
is_identity: bool, is_primary: bool,
is_unique: bool,
}
impl PgColumn {
fn is_auto_increment(&self) -> bool {
self.is_identity
|| self
.column_default
.as_deref()
.map(|d| {
d.contains("nextval(")
|| d.contains("nextval (")
})
.unwrap_or(false)
}
}
fn parse_args() -> Config {
let mut cfg = Config::default();
let mut seen: HashSet<String> = HashSet::new();
let args: Vec<String> = std::env::args().skip(1).collect();
if args.is_empty() {
println!("{}", USAGE);
std::process::exit(0);
}
let mut i = 0;
while i < args.len() {
let flag = args[i].clone();
let take = |idx: &mut usize| -> Option<String> {
if *idx + 1 < args.len() {
*idx += 1;
Some(args[*idx].clone())
} else {
None
}
};
match flag.as_str() {
"--phost" => { seen.insert("phost".into()); if let Some(v) = take(&mut i) { cfg.p_host = v; } },
"--pport" => { seen.insert("pport".into()); if let Some(v) = take(&mut i) { cfg.p_port = v.parse().unwrap_or(cfg.p_port); } },
"--pdb" => { seen.insert("pdb".into()); if let Some(v) = take(&mut i) { cfg.p_db = v; } },
"--puser" => { seen.insert("puser".into()); if let Some(v) = take(&mut i) { cfg.p_user = v; } },
"--ppass" => { seen.insert("ppass".into()); if let Some(v) = take(&mut i) { cfg.p_pass = v; } },
"--mhost" => { seen.insert("mhost".into()); if let Some(v) = take(&mut i) { cfg.m_host = v; } },
"--mport" => { seen.insert("mport".into()); if let Some(v) = take(&mut i) { cfg.m_port = v.parse().unwrap_or(cfg.m_port); } },
"--mdb" => { seen.insert("mdb".into()); if let Some(v) = take(&mut i) { cfg.m_db = v; } },
"--muser" => { seen.insert("muser".into()); if let Some(v) = take(&mut i) { cfg.m_user = v; } },
"--mpass" => { seen.insert("mpass".into()); if let Some(v) = take(&mut i) { cfg.m_pass = v; } },
"--tables" => if let Some(v) = take(&mut i) {
cfg.tables = Some(v.split(',').map(|s| s.trim().to_string()).filter(|s| !s.is_empty()).collect());
},
"--batch" => if let Some(v) = take(&mut i) { cfg.batch_size = v.parse().unwrap_or(cfg.batch_size); },
"--create-only" => cfg.create_only = true,
"--data-only" => cfg.data_only = true,
"--help" | "-h" => {
println!("{}", USAGE);
std::process::exit(0);
}
_ => {
eprintln!("未知参数: {}", flag);
eprintln!("{}", USAGE);
std::process::exit(1);
}
}
i += 1;
}
let required = ["phost", "pport", "pdb", "puser", "mhost", "mport", "mdb", "muser"];
let missing: Vec<&str> = required.iter().copied().filter(|k| !seen.contains(*k)).collect();
if !missing.is_empty() {
eprintln!("缺少必要连接参数: {}", missing.join(", "));
eprintln!("{}", USAGE);
std::process::exit(1);
}
cfg
}
const USAGE: &str = r#"
用法:
postgresql2mysql --phost <host> --pport <port> --pdb <db> --puser <user> --ppass <pass>
--mhost <host> --mport <port> --mdb <db> --muser <user> --mpass <pass>
[--tables t1,t2] [--batch 1000] [--create-only] [--data-only]
选项:
--phost/--pport/--pdb/--puser/--ppass PostgreSQL 源库连接
--mhost/--mport/--mdb/--muser/--mpass MySQL 目标库连接
--tables <a,b,c> 只迁移指定表(默认全部)
--batch <n> 每批迁移行数(默认 1000)
--create-only 只创建表结构,不迁移数据
--data-only 只迁移数据,跳过建表
"#;
#[tokio::main]
async fn main() -> std::result::Result<(), Box<dyn std::error::Error>> {
let cfg = parse_args();
println!(
"PostgreSQL {}:{} @ {} → MySQL {}:{} @ {}",
cfg.p_host, cfg.p_port, cfg.p_db, cfg.m_host, cfg.m_port, cfg.m_db
);
println!("连接 PostgreSQL...");
let pg = Database::postgresql(&cfg.p_host, cfg.p_port, &cfg.p_db, &cfg.p_user, &cfg.p_pass).await?;
pg.ping().await?;
println!(" ✅ PostgreSQL 连接成功 ({})", pg.db_type());
println!("连接 MySQL...");
let mysql = Database::mysql(&cfg.m_host, cfg.m_port, &cfg.m_db, &cfg.m_user, &cfg.m_pass).await?;
mysql.ping().await?;
println!(" ✅ MySQL 连接成功 ({})", mysql.db_type());
let schema = resolve_schema(&pg).await?;
println!(" PostgreSQL schema: {}", schema);
let tables = resolve_tables(&pg, &cfg, &schema).await?;
println!("\n待迁移表 ({}): {:?}\n", tables.len(), tables);
for table in &tables {
let (columns, unique_groups) = fetch_columns(&pg, &schema, table).await?;
if columns.is_empty() {
eprintln!(" ⚠️ 表 {} 无列定义,跳过", table);
continue;
}
if !cfg.data_only {
match create_mysql_table(&mysql, table, &columns, &unique_groups).await {
Ok(_) => println!(" ✅ 已创建表结构: {}", table),
Err(e) => {
eprintln!(" ❌ 建表失败 {}: {}", table, e);
continue;
}
}
}
if cfg.create_only {
continue;
}
let migrated = migrate_table(&pg, &mysql, table, &columns, cfg.batch_size).await?;
println!(" ✅ 完成迁移 {}:{} 行", table, migrated);
}
pg.close().await?;
mysql.close().await?;
println!("\n🎉 迁移完成!");
Ok(())
}
async fn resolve_schema(pg: &Database) -> std::result::Result<String, DbError> {
let result = pg.query("SELECT current_schema() AS schema", &[]).await?;
if let Some(row) = result.rows.first() {
if let Some(SqlValue::String(s)) = row.get("schema") {
if !s.is_empty() {
return Ok(s.to_string());
}
}
}
Ok("public".to_string())
}
async fn resolve_tables(pg: &Database, cfg: &Config, schema: &str) -> std::result::Result<Vec<String>, DbError> {
if let Some(tables) = &cfg.tables {
return Ok(tables.clone());
}
let sql = format!(
"SELECT tablename AS table_name \
FROM pg_tables \
WHERE schemaname = '{}' \
ORDER BY tablename",
schema
);
let result = pg.query(&sql, &[]).await?;
let mut tables = Vec::new();
for row in &result.rows {
if let Some(SqlValue::String(name)) = row.get("table_name") {
tables.push(name.clone());
}
}
Ok(tables)
}
async fn fetch_columns(
pg: &Database,
db_name: &str,
table: &str,
) -> std::result::Result<(Vec<PgColumn>, Vec<Vec<String>>), DbError> {
let sql = format!(
"SELECT c.column_name, c.data_type, c.udt_name, c.is_nullable, c.column_default \
FROM information_schema.columns c \
WHERE c.table_schema = '{}' AND c.table_name = '{}' \
ORDER BY c.ordinal_position",
db_name, table
);
let result = pg.query(&sql, &[]).await?;
let mut columns: Vec<PgColumn> = Vec::new();
for row in &result.rows {
let get = |k: &str| -> Option<String> {
row.get(k).and_then(|v| v.as_str()).map(|s| s.to_string())
};
columns.push(PgColumn {
name: get("column_name").unwrap_or_default(),
data_type: get("data_type").unwrap_or_default(),
udt_name: get("udt_name").unwrap_or_default(),
is_nullable: get("is_nullable").map(|s| s.eq_ignore_ascii_case("YES")).unwrap_or(true),
column_default: get("column_default"),
is_identity: false,
is_primary: false,
is_unique: false,
});
}
for col in &mut columns {
col.is_identity = col
.column_default
.as_deref()
.map(|d| d.contains("nextval(") || d.contains("nextval ("))
.unwrap_or(false);
}
let pk_sql = format!(
"SELECT kcu.column_name \
FROM information_schema.table_constraints tc \
JOIN information_schema.key_column_usage kcu \
ON tc.constraint_name = kcu.constraint_name \
AND tc.table_schema = kcu.table_schema \
AND tc.table_name = kcu.table_name \
WHERE tc.constraint_type = 'PRIMARY KEY' \
AND tc.table_schema = '{}' AND tc.table_name = '{}' \
ORDER BY kcu.ordinal_position",
db_name, table
);
let pk_result = pg.query(&pk_sql, &[]).await?;
let pk_names: Vec<String> = pk_result
.rows
.iter()
.filter_map(|row| row.get("column_name").and_then(|v| v.as_str()).map(|s| s.to_string()))
.collect();
for col in &mut columns {
col.is_primary = pk_names.contains(&col.name);
}
let uq_sql = format!(
"SELECT tc.constraint_name, kcu.column_name \
FROM information_schema.table_constraints tc \
JOIN information_schema.key_column_usage kcu \
ON tc.constraint_name = kcu.constraint_name \
AND tc.table_schema = kcu.table_schema \
AND tc.table_name = kcu.table_name \
WHERE tc.constraint_type = 'UNIQUE' \
AND tc.table_schema = '{}' AND tc.table_name = '{}' \
ORDER BY tc.constraint_name, kcu.ordinal_position",
db_name, table
);
let uq_result = pg.query(&uq_sql, &[]).await?;
let mut groups: Vec<Vec<String>> = Vec::new();
let mut cur_name: Option<String> = None;
let mut cur_cols: Vec<String> = Vec::new();
for row in &uq_result.rows {
let cname = row
.get("constraint_name")
.and_then(|v| v.as_str())
.map(|s| s.to_string());
let col = row
.get("column_name")
.and_then(|v| v.as_str())
.map(|s| s.to_string())
.unwrap_or_default();
if cur_name.is_some() && cur_name.as_deref() != cname.as_deref() {
groups.push(std::mem::take(&mut cur_cols));
}
cur_name = cname;
cur_cols.push(col);
}
if !cur_cols.is_empty() {
groups.push(cur_cols);
}
let mut unique_groups: Vec<Vec<String>> = Vec::new();
for g in groups {
if g.len() == 1 {
for col in &mut columns {
if !col.is_primary && col.name == g[0] {
col.is_unique = true;
}
}
} else {
unique_groups.push(g);
}
}
Ok((columns, unique_groups))
}
async fn create_mysql_table(
mysql: &Database,
table: &str,
columns: &[PgColumn],
unique_groups: &[Vec<String>],
) -> std::result::Result<(), DbError> {
let _ = mysql.execute(&format!("DROP TABLE IF EXISTS `{}`", table), &[]).await;
let mut defs: Vec<String> = Vec::new();
let mut primary_keys: Vec<String> = Vec::new();
let mut key_cols: HashSet<String> = HashSet::new();
for col in columns {
if col.is_primary || col.is_unique {
key_cols.insert(col.name.clone());
}
}
for group in unique_groups {
for c in group {
key_cols.insert(c.clone());
}
}
for col in columns {
let is_key = key_cols.contains(&col.name);
let col_type = mysql_type_for_col(col, is_key);
let mut def = format!(" `{}` {}", col.name, col_type);
if col.is_auto_increment() {
def.push_str(" AUTO_INCREMENT");
}
if !col.is_nullable {
def.push_str(" NOT NULL");
}
if let Some(default) = normalize_default(&col.column_default) {
if !col.is_auto_increment() {
def.push_str(&format!(" DEFAULT {}", default));
}
}
if col.is_unique {
def.push_str(" UNIQUE");
}
if col.is_primary {
primary_keys.push(format!("`{}`", col.name));
}
defs.push(def);
}
if !primary_keys.is_empty() {
defs.push(format!(" PRIMARY KEY ({})", primary_keys.join(", ")));
}
for group in unique_groups {
if group.len() > 1 {
let cols: Vec<String> = group.iter().map(|c| format!("`{}`", c)).collect();
defs.push(format!(" UNIQUE ({})", cols.join(", ")));
}
}
let sql = format!(
"CREATE TABLE IF NOT EXISTS `{}` (\n{}\n) ENGINE=InnoDB DEFAULT CHARSET=utf8mb4 COLLATE=utf8mb4_bin",
table,
defs.join(",\n")
);
mysql.execute(&sql, &[]).await?;
Ok(())
}
async fn migrate_table(
pg: &Database,
mysql: &Database,
table: &str,
columns: &[PgColumn],
batch_size: usize,
) -> std::result::Result<u64, DbError> {
let mysql_cols: Vec<String> = columns.iter().map(|c| format!("`{}`", c.name)).collect();
let placeholders: Vec<&str> = vec!["?"; columns.len()];
let insert_sql = format!(
"INSERT INTO `{}` ({}) VALUES ({})",
table,
mysql_cols.join(", "),
placeholders.join(", ")
);
let pg_cols: Vec<String> = columns.iter().map(|c| format!("\"{}\"", c.name)).collect();
let order_by: String = if let Some(pk) = columns.iter().find(|c| c.is_primary) {
format!("\"{}\"", pk.name)
} else {
pg_cols.join(", ")
};
let select_sql = format!(
"SELECT {} FROM \"{}\" ORDER BY {}",
pg_cols.join(", "),
table,
order_by
);
let mut offset: i64 = 0;
let mut total: u64 = 0;
loop {
let page_sql = format!("{} LIMIT {} OFFSET {}", select_sql, batch_size, offset);
let result = pg.query(&page_sql, &[]).await?;
if result.rows.is_empty() {
break;
}
let mut tx = mysql.begin_transaction().await?;
for row in &result.rows {
let mut params: Vec<SqlValue> = Vec::with_capacity(columns.len());
for col in columns {
params.push(row.get(&col.name).cloned().unwrap_or(SqlValue::Null));
}
tx.execute(&insert_sql, ¶ms).await?;
}
tx.commit().await?;
total += result.rows.len() as u64;
println!(" · {}: 已迁移 {} 行", table, total);
offset += result.rows.len() as i64;
}
Ok(total)
}
fn mysql_type_for_col(col: &PgColumn, is_key: bool) -> String {
let base = pg_type_to_mysql(&col.data_type, &col.udt_name);
let upper = base.to_ascii_uppercase();
if is_key && (upper.contains("TEXT") || upper.contains("BLOB") || upper.contains("JSON")) {
"VARCHAR(255)".to_string()
} else {
base.to_string()
}
}
fn pg_type_to_mysql(data_type: &str, udt_name: &str) -> &'static str {
let t = data_type.to_ascii_lowercase();
let u = udt_name.to_ascii_lowercase();
let t = t.trim();
let u = u.trim();
match (t, u) {
("smallint", "int2") => "SMALLINT",
("integer", "int4") => "INT",
("bigint", "int8") => "BIGINT",
("real", "float4") => "FLOAT",
("double precision", "float8") => "DOUBLE",
("numeric", "numeric") => "DECIMAL",
("money", "money") => "DECIMAL(19,2)",
("boolean", "bool") => "TINYINT(1)",
("character varying", "varchar") => "LONGTEXT",
("character", "bpchar") => "CHAR(1)",
("text", "text") => "LONGTEXT",
("name", "name") => "VARCHAR(64)",
("citext", "citext") => "VARCHAR(255)",
("bytea", "bytea") => "BLOB",
("bit", "bit") => "BIT",
("bit varying", "varbit") => "VARBINARY(255)",
("timestamp without time zone", "timestamp") => "DATETIME",
("timestamp with time zone", "timestamptz") => "DATETIME",
("date", "date") => "DATE",
("time without time zone", "time") => "TIME",
("time with time zone", "timetz") => "TIME",
("json", "json") => "LONGTEXT",
("jsonb", "jsonb") => "LONGTEXT",
("uuid", "uuid") => "CHAR(36)",
("inet", "inet") => "VARCHAR(45)",
("macaddr", "macaddr") => "VARCHAR(17)",
("interval", "interval") => "VARCHAR(64)",
("smallserial", "int2") => "SMALLINT",
("serial", "int4") => "INT",
("bigserial", "int8") => "BIGINT",
("array", _) => "LONGTEXT",
_ => "TEXT", }
}
fn normalize_default(default: &Option<String>) -> Option<String> {
let d = default.as_ref()?.trim();
if d.is_empty() {
return None;
}
let lower = d.to_ascii_lowercase();
if lower.contains("nextval(") || lower.contains("nextval (") {
return None;
}
if lower == "now()"
|| lower == "current_timestamp"
|| lower == "current_timestamp()"
|| lower.contains("timezone('utc'")
|| (lower.contains("timezone(") && lower.contains("now()"))
{
return Some("CURRENT_TIMESTAMP".to_string());
}
if lower == "current_date" || lower == "current_date()" {
return Some("CURRENT_DATE".to_string());
}
if lower == "current_user" {
return Some("CURRENT_USER".to_string());
}
if lower == "true" {
return Some("1".to_string());
}
if lower == "false" {
return Some("0".to_string());
}
if d.eq_ignore_ascii_case("null") {
return Some("NULL".to_string());
}
if d.parse::<f64>().is_ok() {
return Some(d.to_string());
}
let cleaned = strip_pg_type_cast(d);
Some(format!("'{}'", cleaned.trim_matches(|c| c == '\'' || c == '"')))
}
fn strip_pg_type_cast(s: &str) -> String {
if let Some(pos) = s.find("::") {
let prefix = s[..pos].trim();
prefix.to_string()
} else {
s.to_string()
}
}