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,
sqlite_file: 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(),
sqlite_file: "migrated.db".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; } },
"--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);
}
_ if !flag.starts_with('-') && !seen.contains("sqlite_file") => {
seen.insert("sqlite_file".into());
cfg.sqlite_file = flag;
}
_ => {
eprintln!("未知参数: {}", flag);
eprintln!("{}", USAGE);
std::process::exit(1);
}
}
i += 1;
}
let mut missing: Vec<&str> = Vec::new();
if !seen.contains("sqlite_file") {
missing.push("sqlite_file");
}
for k in ["phost", "pport", "pdb", "puser"] {
if !seen.contains(k) {
missing.push(k);
}
}
if !missing.is_empty() {
eprintln!("缺少必要参数: {}", missing.join(", "));
eprintln!("{}", USAGE);
std::process::exit(1);
}
cfg
}
const USAGE: &str = r#"
用法:
postgres2sqlite <sqlite_file>
[--phost <host> --pport <port> --pdb <db> --puser <user> --ppass <pass>]
[--tables t1,t2] [--batch 1000] [--create-only] [--data-only]
选项:
<sqlite_file> 目标 SQLite 数据库文件路径(必填,不存在则创建)
--phost/--pport/--pdb/--puser/--ppass PostgreSQL 源库连接(缺省用默认配置)
--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 {}:{} @ {} → SQLite file: {}",
cfg.p_host, cfg.p_port, cfg.p_db, cfg.sqlite_file
);
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!("连接 SQLite...");
let sqlite = Database::sqlite(&cfg.sqlite_file).await?;
sqlite.ping().await?;
println!(" ✅ SQLite 连接成功 ({})", sqlite.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_sqlite_table(&sqlite, table, &columns, &unique_groups).await {
Ok(_) => println!(" ✅ 已创建表结构: {}", table),
Err(e) => {
eprintln!(" ❌ 建表失败 {}: {}", table, e);
continue;
}
}
}
if cfg.create_only {
continue;
}
let migrated = migrate_table(&pg, &sqlite, table, &columns, cfg.batch_size).await?;
println!(" ✅ 完成迁移 {}:{} 行", table, migrated);
}
pg.close().await?;
sqlite.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_sqlite_table(
sqlite: &Database,
table: &str,
columns: &[PgColumn],
unique_groups: &[Vec<String>],
) -> std::result::Result<(), DbError> {
let _ = sqlite.execute(&format!("DROP TABLE IF EXISTS \"{}\"", table), &[]).await;
let mut defs: Vec<String> = Vec::new();
let mut primary_keys: Vec<String> = Vec::new();
for col in columns {
let mut def = format!(" \"{}\" {}", col.name, pg_type_to_sqlite(&col.data_type, &col.udt_name));
if col.is_auto_increment() && col.is_primary {
def = format!(" \"{}\" INTEGER PRIMARY KEY AUTOINCREMENT", col.name);
if !col.is_nullable {
def.push_str(" NOT NULL");
}
defs.push(def);
continue;
}
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)", table, defs.join(",\n"));
sqlite.execute(&sql, &[]).await?;
Ok(())
}
async fn migrate_table(
pg: &Database,
sqlite: &Database,
table: &str,
columns: &[PgColumn],
batch_size: usize,
) -> std::result::Result<u64, DbError> {
let sqlite_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,
sqlite_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 = sqlite.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 pg_type_to_sqlite(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") => "INTEGER",
("integer", "int4") => "INTEGER",
("bigint", "int8") => "INTEGER",
("smallserial", "int2") => "INTEGER",
("serial", "int4") => "INTEGER",
("bigserial", "int8") => "INTEGER",
("real", "float4") => "REAL",
("double precision", "float8") => "REAL",
("numeric", "numeric") => "REAL",
("money", "money") => "REAL",
("boolean", "bool") => "INTEGER",
("character varying", "varchar") => "TEXT",
("character", "bpchar") => "TEXT",
("text", "text") => "TEXT",
("name", "name") => "TEXT",
("citext", "citext") => "TEXT",
("bytea", "bytea") => "BLOB",
("bit", "bit") => "BLOB",
("bit varying", "varbit") => "BLOB",
("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",
("interval", "interval") => "TEXT",
("json", "json") => "TEXT",
("jsonb", "jsonb") => "TEXT",
("uuid", "uuid") => "TEXT",
("inet", "inet") => "TEXT",
("macaddr", "macaddr") => "TEXT",
("array", _) => "TEXT",
_ => "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()" {
return Some("CURRENT_TIMESTAMP".to_string());
}
if lower == "current_date" || lower == "current_date()" {
return Some("CURRENT_DATE".to_string());
}
if lower == "current_time" || lower == "current_time()" {
return Some("CURRENT_TIME".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("::") {
s[..pos].trim().to_string()
} else {
s.to_string()
}
}