use std::collections::{HashMap, HashSet};
use sea_query::{Alias, Expr};
use sqlx::Row;
use crate::db::DbPool;
use crate::migrate::ModelMeta;
use crate::orm::SqlType;
use crate::orm::dynamic::DynQuerySet;
const STATE_TABLE: &str = "umbral_transfer_state";
#[derive(Debug, Clone, PartialEq, Eq, Default)]
pub enum TransferMap {
#[default]
None,
Django,
Rails,
Laravel,
Prisma,
Custom(CustomMap),
}
#[derive(Debug, Clone, PartialEq, Eq, Default, serde::Deserialize)]
#[serde(deny_unknown_fields)]
pub struct CustomMap {
#[serde(default)]
pub columns: std::collections::BTreeMap<String, String>,
#[serde(default)]
pub tables: std::collections::BTreeMap<String, std::collections::BTreeMap<String, String>>,
}
impl TransferMap {
pub fn parse(s: &str) -> Option<TransferMap> {
match s.to_ascii_lowercase().as_str() {
"django" => Some(TransferMap::Django),
"rails" | "activerecord" => Some(TransferMap::Rails),
"laravel" | "eloquent" => Some(TransferMap::Laravel),
"prisma" | "typeorm" => Some(TransferMap::Prisma),
"none" | "" => Some(TransferMap::None),
_ => None,
}
}
pub fn from_cli_arg(s: &str) -> Result<TransferMap, String> {
if let Some(preset) = Self::parse(s) {
return Ok(preset);
}
let path = std::path::Path::new(s);
if !path.is_file() {
return Err(format!(
"unknown --map `{s}`: not a framework (django / rails / laravel / prisma) \
and not a readable JSON file"
));
}
let text =
std::fs::read_to_string(path).map_err(|e| format!("--map: cannot read `{s}`: {e}"))?;
let custom: CustomMap = serde_json::from_str(&text)
.map_err(|e| format!("--map: `{s}` is not a valid mapping file: {e}"))?;
Ok(TransferMap::Custom(custom))
}
fn source_column(&self, table: &str, field: &str, is_fk: bool) -> Option<String> {
match self {
TransferMap::None => None,
TransferMap::Django | TransferMap::Rails | TransferMap::Laravel => {
if is_fk && !field.ends_with("_id") {
Some(format!("{field}_id"))
} else {
None
}
}
TransferMap::Prisma => {
let base = to_lower_camel(field);
let col = if is_fk { format!("{base}Id") } else { base };
(col != field).then_some(col)
}
TransferMap::Custom(m) => m
.tables
.get(table)
.and_then(|t| t.get(field))
.or_else(|| m.columns.get(field))
.cloned(),
}
}
}
fn to_lower_camel(s: &str) -> String {
let snake = umbral_casing::to_snake_case(s);
let mut out = String::new();
for (i, part) in snake.split('_').filter(|p| !p.is_empty()).enumerate() {
if i == 0 {
out.push_str(part);
} else {
let mut chars = part.chars();
if let Some(first) = chars.next() {
out.extend(first.to_uppercase());
out.push_str(chars.as_str());
}
}
}
out
}
#[derive(Debug, Clone)]
pub struct TransferOptions {
pub batch_size: u64,
pub only: Option<Vec<String>>,
pub dry_run: bool,
pub map: TransferMap,
pub workers: usize,
}
impl Default for TransferOptions {
fn default() -> Self {
Self {
batch_size: 1000,
only: None,
dry_run: false,
map: TransferMap::None,
workers: 1,
}
}
}
#[derive(Debug, Default)]
pub struct TransferReport {
pub per_table: Vec<(String, u64)>,
pub rows: u64,
}
#[derive(Debug)]
pub enum TransferError {
Db(sqlx::Error),
Write(String),
Read(String),
NoPrimaryKey(String),
}
impl std::fmt::Display for TransferError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
TransferError::Db(e) => write!(f, "database error: {e}"),
TransferError::Write(e) => write!(f, "write error: {e}"),
TransferError::Read(e) => write!(f, "read error: {e}"),
TransferError::NoPrimaryKey(t) => {
write!(f, "table `{t}` has no primary key; cannot stream it")
}
}
}
}
impl std::error::Error for TransferError {}
impl From<sqlx::Error> for TransferError {
fn from(e: sqlx::Error) -> Self {
TransferError::Db(e)
}
}
pub fn fk_topo_order(models: Vec<ModelMeta>) -> Vec<ModelMeta> {
fk_topo_levels(models).into_iter().flatten().collect()
}
pub fn fk_topo_levels(models: Vec<ModelMeta>) -> Vec<Vec<ModelMeta>> {
let (mut levels, cyclic) = fk_topo_plan(models);
if !cyclic.is_empty() {
levels.push(cyclic);
}
levels
}
pub fn fk_topo_plan(models: Vec<ModelMeta>) -> (Vec<Vec<ModelMeta>>, Vec<ModelMeta>) {
let tables: HashSet<String> = models.iter().map(|m| m.table.clone()).collect();
let mut deps: HashMap<String, HashSet<String>> = HashMap::new();
for m in &models {
let mut d = HashSet::new();
for col in &m.fields {
if let Some(target) = &col.fk_target {
if target != &m.table && tables.contains(target) {
d.insert(target.clone());
}
}
}
deps.insert(m.table.clone(), d);
}
let mut by_table: HashMap<String, ModelMeta> =
models.into_iter().map(|m| (m.table.clone(), m)).collect();
let mut levels: Vec<Vec<ModelMeta>> = Vec::new();
let mut placed: HashSet<String> = HashSet::new();
loop {
let mut ready: Vec<String> = by_table
.keys()
.filter(|t| !placed.contains(*t))
.filter(|t| deps[*t].iter().all(|d| placed.contains(d)))
.cloned()
.collect();
if ready.is_empty() {
break;
}
ready.sort();
let level: Vec<ModelMeta> = ready.iter().map(|t| by_table.remove(t).unwrap()).collect();
for t in ready {
placed.insert(t);
}
levels.push(level);
}
let mut cyclic: Vec<ModelMeta> = by_table.into_values().collect();
cyclic.sort_by(|a, b| a.table.cmp(&b.table));
(levels, cyclic)
}
#[derive(Debug, Clone)]
struct Junction {
table: String,
source_table: String,
owner_table: String,
child_table: String,
parent_ty: SqlType,
child_ty: SqlType,
}
fn collect_junctions(models: &[ModelMeta], map: &TransferMap) -> Vec<Junction> {
let pk_ty = |table: &str| -> SqlType {
models
.iter()
.find(|m| m.table == table)
.and_then(|m| m.pk_column())
.map(|c| c.ty)
.unwrap_or(SqlType::BigInt)
};
let mut out = Vec::new();
for m in models {
let parent_ty = m.pk_column().map(|c| c.ty).unwrap_or(SqlType::BigInt);
for rel in &m.m2m_relations {
let table = format!("{}_{}", m.table, rel.field_name);
let source_table = match map {
TransferMap::Prisma => {
let mut ends = [m.name.as_str(), rel.target_name.as_str()];
ends.sort_unstable();
format!("_{}To{}", ends[0], ends[1])
}
_ => table.clone(),
};
out.push(Junction {
table,
source_table,
owner_table: m.table.clone(),
child_table: rel.target_table.clone(),
parent_ty,
child_ty: pk_ty(&rel.target_table),
});
}
}
out
}
async fn resolve_junction_columns(
source: &DbPool,
jn: &Junction,
) -> Result<(String, String), TransferError> {
let fks: Vec<(String, String)> = match source {
DbPool::Sqlite(pool) => {
let jt = jn.source_table.replace('"', "\"\"");
let rows = sqlx::query(&format!("PRAGMA foreign_key_list(\"{jt}\")"))
.fetch_all(pool)
.await?;
rows.iter()
.map(|r| {
Ok::<_, TransferError>((
r.try_get::<String, _>("from")?,
r.try_get::<String, _>("table")?,
))
})
.collect::<Result<_, _>>()?
}
DbPool::Postgres(pool) => {
sqlx::query_as(
"SELECT kcu.column_name, ccu.table_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 \
JOIN information_schema.constraint_column_usage ccu \
ON ccu.constraint_name = tc.constraint_name AND ccu.table_schema = tc.table_schema \
WHERE tc.constraint_type = 'FOREIGN KEY' AND tc.table_schema = 'public' \
AND tc.table_name = $1 ORDER BY kcu.ordinal_position",
)
.bind(&jn.source_table)
.fetch_all(pool)
.await?
}
};
if jn.owner_table == jn.child_table {
let mut cols = fks.into_iter().map(|(c, _)| c);
return Ok((
cols.next().unwrap_or_else(|| "parent_id".to_string()),
cols.next().unwrap_or_else(|| "child_id".to_string()),
));
}
let parent = fks
.iter()
.find(|(_, t)| *t == jn.owner_table)
.map(|(c, _)| c.clone())
.unwrap_or_else(|| "parent_id".to_string());
let child = fks
.iter()
.find(|(_, t)| *t == jn.child_table)
.map(|(c, _)| c.clone())
.unwrap_or_else(|| "child_id".to_string());
Ok((parent, child))
}
async fn copy_one_junction(
source: &DbPool,
target: &DbPool,
jn: &Junction,
start_last: Option<(serde_json::Value, serde_json::Value)>,
batch: u64,
) -> Result<u64, TransferError> {
let (pcol, ccol) = resolve_junction_columns(source, jn).await?;
let mut last = start_last;
let mut copied: u64 = 0;
loop {
let rows = read_junction_batch(source, jn, &pcol, &ccol, last.as_ref(), batch).await?;
if rows.is_empty() {
break;
}
let batch_len = rows.len();
let new_last = rows.last().cloned();
let mut tx = begin_on(target).await?;
for (p, c) in &rows {
insert_junction_in_tx(&mut tx, jn, p, c).await?;
}
let checkpoint = new_last
.as_ref()
.map(|(p, c)| serde_json::Value::Array(vec![p.clone(), c.clone()]));
upsert_state_in_tx(&mut tx, &jn.table, checkpoint.as_ref(), false).await?;
tx.commit().await?;
copied += batch_len as u64;
last = new_last;
if batch_len < batch as usize {
break;
}
}
let checkpoint = last
.as_ref()
.map(|(p, c)| serde_json::Value::Array(vec![p.clone(), c.clone()]));
let mut tx = begin_on(target).await?;
upsert_state_in_tx(&mut tx, &jn.table, checkpoint.as_ref(), true).await?;
tx.commit().await?;
Ok(copied)
}
#[derive(Clone, Copy, PartialEq)]
enum IdKind {
Int,
Uuid,
Text,
}
fn id_kind(ty: SqlType) -> IdKind {
match ty {
SqlType::Integer | SqlType::BigInt | SqlType::SmallInt => IdKind::Int,
SqlType::Uuid => IdKind::Uuid,
_ => IdKind::Text,
}
}
fn read_id_sqlite(
row: &sqlx::sqlite::SqliteRow,
idx: usize,
ty: SqlType,
) -> Result<serde_json::Value, TransferError> {
Ok(match id_kind(ty) {
IdKind::Int => serde_json::Value::from(row.try_get::<i64, _>(idx)?),
IdKind::Uuid => serde_json::Value::from(row.try_get::<uuid::Uuid, _>(idx)?.to_string()),
IdKind::Text => serde_json::Value::from(row.try_get::<String, _>(idx)?),
})
}
fn read_id_pg(
row: &sqlx::postgres::PgRow,
idx: usize,
ty: SqlType,
) -> Result<serde_json::Value, TransferError> {
Ok(match id_kind(ty) {
IdKind::Int => serde_json::Value::from(row.try_get::<i64, _>(idx)?),
IdKind::Uuid => serde_json::Value::from(row.try_get::<uuid::Uuid, _>(idx)?.to_string()),
IdKind::Text => serde_json::Value::from(row.try_get::<String, _>(idx)?),
})
}
fn pk_gt_condition(pk_col: &str, last: &serde_json::Value) -> sea_query::SimpleExpr {
let col = Expr::col(Alias::new(pk_col));
match last {
serde_json::Value::Number(n) if n.is_i64() => col.gt(n.as_i64().unwrap()),
serde_json::Value::Number(n) if n.is_u64() => col.gt(n.as_u64().unwrap() as i64),
serde_json::Value::String(s) => col.gt(s.clone()),
_ => col.gt(last.to_string()),
}
}
fn source_meta_for(meta: &ModelMeta, map: &TransferMap) -> (ModelMeta, HashMap<String, String>) {
let mut rename = HashMap::new();
if matches!(map, TransferMap::None) {
return (meta.clone(), rename);
}
let mut src = meta.clone();
for col in &mut src.fields {
if let Some(source_col) = map.source_column(&meta.table, &col.name, col.fk_target.is_some())
{
rename.insert(source_col.clone(), col.name.clone());
col.name = source_col;
}
}
(src, rename)
}
fn apply_key_rename(
row: &serde_json::Map<String, serde_json::Value>,
rename: &HashMap<String, String>,
) -> serde_json::Map<String, serde_json::Value> {
if rename.is_empty() {
return row.clone();
}
row.iter()
.map(|(k, v)| {
(
rename.get(k).cloned().unwrap_or_else(|| k.clone()),
v.clone(),
)
})
.collect()
}
#[allow(clippy::too_many_arguments)]
async fn copy_one_model(
source: &DbPool,
target: &DbPool,
read_meta: &ModelMeta,
write_meta: &ModelMeta,
key_rename: &HashMap<String, String>,
pk_col: &str,
start_last: Option<serde_json::Value>,
batch: u64,
) -> Result<u64, TransferError> {
let mut last = start_last;
let mut copied: u64 = 0;
loop {
let mut qs = DynQuerySet::for_meta(read_meta).unredacted_for_backup();
if let Some(l) = &last {
qs = qs.filter_condition(sea_query::Condition::all().add(pk_gt_condition(pk_col, l)));
}
let rows = qs
.order_by_col(pk_col, false)
.limit(batch)
.fetch_as_json_on(source)
.await
.map_err(|e| TransferError::Read(e.to_string()))?;
if rows.is_empty() {
break;
}
let batch_len = rows.len();
let new_last = rows.last().and_then(|r| r.get(pk_col)).cloned();
let mut tx = begin_on(target).await?;
for row in &rows {
let mapped = apply_key_rename(row, key_rename);
DynQuerySet::for_meta(write_meta)
.presealed()
.trusted()
.insert_json_in_tx(&mapped, &mut tx)
.await
.map_err(|e| TransferError::Write(e.to_string()))?;
}
upsert_state_in_tx(&mut tx, &write_meta.table, new_last.as_ref(), false).await?;
tx.commit().await?;
copied += batch_len as u64;
last = new_last;
if batch_len < batch as usize {
break;
}
}
let mut tx = begin_on(target).await?;
upsert_state_in_tx(&mut tx, &write_meta.table, last.as_ref(), true).await?;
tx.commit().await?;
reset_sequence(target, &write_meta.table, pk_col).await?;
Ok(copied)
}
fn has_self_fk(meta: &ModelMeta) -> bool {
meta.fields
.iter()
.any(|c| c.fk_target.as_deref() == Some(meta.table.as_str()))
}
async fn defer_fk_in_tx(tx: &mut crate::db::Transaction) -> Result<(), TransferError> {
match tx.backend_name() {
"sqlite" => {
let inner = tx.as_sqlite_mut().expect("sqlite backend");
sqlx::query("PRAGMA defer_foreign_keys = ON")
.execute(&mut **inner)
.await?;
}
_ => {
let inner = tx.as_pg_mut().expect("postgres backend");
sqlx::query("SET CONSTRAINTS ALL DEFERRED")
.execute(&mut **inner)
.await?;
}
}
Ok(())
}
async fn copy_cyclic_group(
source: &DbPool,
target: &DbPool,
group: &[&ModelMeta],
map: &TransferMap,
batch: u64,
) -> Result<Vec<(String, u64)>, TransferError> {
let mut tx = begin_on(target).await?;
defer_fk_in_tx(&mut tx).await?;
let mut results = Vec::new();
for meta in group {
let pk_col = meta
.pk_column()
.ok_or_else(|| TransferError::NoPrimaryKey(meta.table.clone()))?
.name
.clone();
let (read_meta, key_rename) = source_meta_for(meta, map);
let mut last: Option<serde_json::Value> = None;
let mut copied: u64 = 0;
loop {
let mut qs = DynQuerySet::for_meta(&read_meta).unredacted_for_backup();
if let Some(l) = &last {
qs = qs
.filter_condition(sea_query::Condition::all().add(pk_gt_condition(&pk_col, l)));
}
let rows = qs
.order_by_col(&pk_col, false)
.limit(batch)
.fetch_as_json_on(source)
.await
.map_err(|e| TransferError::Read(e.to_string()))?;
if rows.is_empty() {
break;
}
let batch_len = rows.len();
last = rows.last().and_then(|r| r.get(&pk_col)).cloned();
for row in &rows {
let mapped = apply_key_rename(row, &key_rename);
DynQuerySet::for_meta(meta)
.presealed()
.trusted()
.insert_json_in_tx(&mapped, &mut tx)
.await
.map_err(|e| TransferError::Write(e.to_string()))?;
}
copied += batch_len as u64;
if batch_len < batch as usize {
break;
}
}
upsert_state_in_tx(&mut tx, &meta.table, last.as_ref(), true).await?;
results.push((meta.table.clone(), copied));
}
tx.commit().await?;
for meta in group {
if let Some(pk) = meta.pk_column() {
reset_sequence(target, &meta.table, &pk.name).await?;
}
}
Ok(results)
}
pub async fn transfer(
source: &DbPool,
target: &DbPool,
models: Vec<ModelMeta>,
opts: &TransferOptions,
) -> Result<TransferReport, TransferError> {
use futures_util::stream::{StreamExt, TryStreamExt};
let (levels, cyclic) = fk_topo_plan(models);
let flat: Vec<ModelMeta> = levels
.iter()
.flatten()
.chain(cyclic.iter())
.cloned()
.collect();
let only: Option<HashSet<String>> = opts.only.as_ref().map(|v| v.iter().cloned().collect());
let excluded = |t: &str| only.as_ref().is_some_and(|s| !s.contains(t));
let workers = opts.workers.max(1);
let mut report = TransferReport::default();
if opts.dry_run {
for meta in &flat {
if excluded(&meta.table) {
continue;
}
let n = count_rows(source, &meta.table).await?;
report.per_table.push((meta.table.clone(), n));
report.rows += n;
}
for jn in collect_junctions(&flat, &opts.map) {
if excluded(&jn.table) {
continue;
}
let n = count_rows(source, &jn.source_table).await?;
report.per_table.push((jn.table.clone(), n));
report.rows += n;
}
return Ok(report);
}
ensure_state_table(target).await?;
let state = read_state(target).await?;
let done = |t: &str| state.get(t).is_some_and(|(_, d)| *d);
let state_ref = &state;
for level in &levels {
let tasks = level
.iter()
.filter(|m| !excluded(&m.table) && !done(&m.table))
.map(|meta| {
let map = opts.map.clone();
let batch = opts.batch_size;
async move {
if has_self_fk(meta) {
return copy_cyclic_group(source, target, &[meta], &map, batch).await;
}
let pk_col = meta
.pk_column()
.ok_or_else(|| TransferError::NoPrimaryKey(meta.table.clone()))?
.name
.clone();
let (read_meta, key_rename) = source_meta_for(meta, &map);
let start_last = state_ref.get(&meta.table).and_then(|(pk, _)| pk.clone());
let copied = copy_one_model(
source,
target,
&read_meta,
meta,
&key_rename,
&pk_col,
start_last,
batch,
)
.await?;
Ok::<_, TransferError>(vec![(meta.table.clone(), copied)])
}
});
let results: Vec<Vec<(String, u64)>> = futures_util::stream::iter(tasks)
.buffer_unordered(workers)
.try_collect()
.await?;
for (t, c) in results.into_iter().flatten() {
report.per_table.push((t, c));
report.rows += c;
}
}
let cyclic_group: Vec<&ModelMeta> = cyclic
.iter()
.filter(|m| !excluded(&m.table) && !done(&m.table))
.collect();
if !cyclic_group.is_empty() {
let results =
copy_cyclic_group(source, target, &cyclic_group, &opts.map, opts.batch_size).await?;
for (t, c) in results {
report.per_table.push((t, c));
report.rows += c;
}
}
let junctions = collect_junctions(&flat, &opts.map);
let jtasks = junctions
.iter()
.filter(|jn| !excluded(&jn.table) && !done(&jn.table))
.map(|jn| {
let start = state
.get(&jn.table)
.and_then(|(pk, _)| pk.clone())
.and_then(decode_pair);
async move {
let copied = copy_one_junction(source, target, jn, start, opts.batch_size).await?;
Ok::<_, TransferError>((jn.table.clone(), copied))
}
});
let jresults: Vec<(String, u64)> = futures_util::stream::iter(jtasks)
.buffer_unordered(workers)
.try_collect()
.await?;
for (t, c) in jresults {
report.per_table.push((t, c));
report.rows += c;
}
Ok(report)
}
async fn begin_on(pool: &DbPool) -> Result<crate::db::Transaction, TransferError> {
Ok(match pool {
DbPool::Sqlite(p) => crate::db::begin_sqlite(p).await?,
DbPool::Postgres(p) => crate::db::begin_pg(p).await?,
})
}
async fn count_rows(pool: &DbPool, table: &str) -> Result<u64, TransferError> {
let sql = format!("SELECT COUNT(*) FROM \"{}\"", table.replace('"', "\"\""));
let n: i64 = match pool {
DbPool::Sqlite(p) => sqlx::query_scalar(&sql).fetch_one(p).await?,
DbPool::Postgres(p) => sqlx::query_scalar(&sql).fetch_one(p).await?,
};
Ok(n.max(0) as u64)
}
async fn ensure_state_table(target: &DbPool) -> Result<(), TransferError> {
let sql = format!(
"CREATE TABLE IF NOT EXISTS {STATE_TABLE} \
(table_name TEXT PRIMARY KEY, last_pk TEXT, done INTEGER NOT NULL DEFAULT 0)"
);
match target {
DbPool::Sqlite(p) => {
sqlx::query(&sql).execute(p).await?;
}
DbPool::Postgres(p) => {
sqlx::query(&sql).execute(p).await?;
}
}
Ok(())
}
async fn read_state(
target: &DbPool,
) -> Result<HashMap<String, (Option<serde_json::Value>, bool)>, TransferError> {
let sql = format!("SELECT table_name, last_pk, done FROM {STATE_TABLE}");
let mut out = HashMap::new();
let rows: Vec<(String, Option<String>, i64)> = match target {
DbPool::Sqlite(p) => sqlx::query_as(&sql).fetch_all(p).await?,
DbPool::Postgres(p) => sqlx::query_as(&sql).fetch_all(p).await?,
};
for (table, last, done) in rows {
let pk = last.and_then(|s| serde_json::from_str::<serde_json::Value>(&s).ok());
out.insert(table, (pk, done != 0));
}
Ok(out)
}
async fn upsert_state_in_tx(
tx: &mut crate::db::Transaction,
table: &str,
last_pk: Option<&serde_json::Value>,
done: bool,
) -> Result<(), TransferError> {
let last_str = last_pk.map(|v| v.to_string());
let done_int = i64::from(done);
match tx.backend_name() {
"sqlite" => {
let sql = format!(
"INSERT INTO {STATE_TABLE} (table_name, last_pk, done) VALUES (?, ?, ?) \
ON CONFLICT(table_name) DO UPDATE SET last_pk = excluded.last_pk, done = excluded.done"
);
let inner = tx.as_sqlite_mut().expect("sqlite backend");
sqlx::query(&sql)
.bind(table)
.bind(last_str)
.bind(done_int)
.execute(&mut **inner)
.await?;
}
_ => {
let sql = format!(
"INSERT INTO {STATE_TABLE} (table_name, last_pk, done) VALUES ($1, $2, $3) \
ON CONFLICT(table_name) DO UPDATE SET last_pk = excluded.last_pk, done = excluded.done"
);
let inner = tx.as_pg_mut().expect("postgres backend");
sqlx::query(&sql)
.bind(table)
.bind(last_str)
.bind(done_int)
.execute(&mut **inner)
.await?;
}
}
Ok(())
}
async fn reset_sequence(target: &DbPool, table: &str, pk_col: &str) -> Result<(), TransferError> {
if let DbPool::Postgres(p) = target {
let sql = format!(
"SELECT setval(pg_get_serial_sequence('{t}', '{c}'), \
COALESCE((SELECT MAX(\"{c}\") FROM \"{t}\"), 1)) \
WHERE pg_get_serial_sequence('{t}', '{c}') IS NOT NULL",
t = table.replace('\'', "''"),
c = pk_col.replace('\'', "''"),
);
let _ = sqlx::query(&sql).execute(p).await;
}
Ok(())
}
fn decode_pair(v: serde_json::Value) -> Option<(serde_json::Value, serde_json::Value)> {
match v {
serde_json::Value::Array(a) if a.len() == 2 => Some((a[0].clone(), a[1].clone())),
_ => None,
}
}
async fn read_junction_batch(
source: &DbPool,
jn: &Junction,
parent_col: &str,
child_col: &str,
last: Option<&(serde_json::Value, serde_json::Value)>,
limit: u64,
) -> Result<Vec<(serde_json::Value, serde_json::Value)>, TransferError> {
let jt = jn.source_table.replace('"', "\"\"");
let pcol = parent_col.replace('"', "\"\"");
let ccol = child_col.replace('"', "\"\"");
let mut out = Vec::new();
match source {
DbPool::Sqlite(pool) => {
let where_sql = if last.is_some() {
format!("WHERE (\"{pcol}\", \"{ccol}\") > (?, ?)")
} else {
String::new()
};
let sql = format!(
"SELECT \"{pcol}\", \"{ccol}\" FROM \"{jt}\" {where_sql} \
ORDER BY \"{pcol}\", \"{ccol}\" LIMIT {limit}"
);
let mut q = sqlx::query(&sql);
if let Some((p, c)) = last {
q = bind_id_sqlite(q, p, jn.parent_ty);
q = bind_id_sqlite(q, c, jn.child_ty);
}
for row in q.fetch_all(pool).await? {
out.push((
read_id_sqlite(&row, 0, jn.parent_ty)?,
read_id_sqlite(&row, 1, jn.child_ty)?,
));
}
}
DbPool::Postgres(pool) => {
let where_sql = if last.is_some() {
format!("WHERE (\"{pcol}\", \"{ccol}\") > ($1, $2)")
} else {
String::new()
};
let sql = format!(
"SELECT \"{pcol}\", \"{ccol}\" FROM \"{jt}\" {where_sql} \
ORDER BY \"{pcol}\", \"{ccol}\" LIMIT {limit}"
);
let mut q = sqlx::query(&sql);
if let Some((p, c)) = last {
q = bind_id_pg(q, p, jn.parent_ty);
q = bind_id_pg(q, c, jn.child_ty);
}
for row in q.fetch_all(pool).await? {
out.push((
read_id_pg(&row, 0, jn.parent_ty)?,
read_id_pg(&row, 1, jn.child_ty)?,
));
}
}
}
Ok(out)
}
fn bind_id_sqlite<'q>(
q: sqlx::query::Query<'q, sqlx::Sqlite, sqlx::sqlite::SqliteArguments<'q>>,
v: &serde_json::Value,
ty: SqlType,
) -> sqlx::query::Query<'q, sqlx::Sqlite, sqlx::sqlite::SqliteArguments<'q>> {
match id_kind(ty) {
IdKind::Int => q.bind(v.as_i64()),
_ => q.bind(v.as_str().map(str::to_string)),
}
}
fn bind_id_pg<'q>(
q: sqlx::query::Query<'q, sqlx::Postgres, sqlx::postgres::PgArguments>,
v: &serde_json::Value,
ty: SqlType,
) -> sqlx::query::Query<'q, sqlx::Postgres, sqlx::postgres::PgArguments> {
match id_kind(ty) {
IdKind::Int => q.bind(v.as_i64()),
IdKind::Uuid => q.bind(v.as_str().and_then(|s| uuid::Uuid::parse_str(s).ok())),
IdKind::Text => q.bind(v.as_str().map(str::to_string)),
}
}
async fn insert_junction_in_tx(
tx: &mut crate::db::Transaction,
jn: &Junction,
p: &serde_json::Value,
c: &serde_json::Value,
) -> Result<(), TransferError> {
let jt = jn.table.replace('"', "\"\"");
match tx.backend_name() {
"sqlite" => {
let sql = format!(
"INSERT INTO \"{jt}\" (parent_id, child_id) VALUES (?, ?) \
ON CONFLICT (parent_id, child_id) DO NOTHING"
);
let inner = tx.as_sqlite_mut().expect("sqlite backend");
let mut q = sqlx::query(&sql);
q = bind_id_sqlite(q, p, jn.parent_ty);
q = bind_id_sqlite(q, c, jn.child_ty);
q.execute(&mut **inner).await?;
}
_ => {
let sql = format!(
"INSERT INTO \"{jt}\" (parent_id, child_id) VALUES ($1, $2) \
ON CONFLICT (parent_id, child_id) DO NOTHING"
);
let inner = tx.as_pg_mut().expect("postgres backend");
let mut q = sqlx::query(&sql);
q = bind_id_pg(q, p, jn.parent_ty);
q = bind_id_pg(q, c, jn.child_ty);
q.execute(&mut **inner).await?;
}
}
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn custom_map_prefers_table_over_global_and_falls_back() {
let mut columns = std::collections::BTreeMap::new();
columns.insert("created_at".to_string(), "createdAt".to_string());
let mut users = std::collections::BTreeMap::new();
users.insert("created_at".to_string(), "user_created".to_string());
let mut tables = std::collections::BTreeMap::new();
tables.insert("users".to_string(), users);
let map = TransferMap::Custom(CustomMap { columns, tables });
assert_eq!(
map.source_column("users", "created_at", false).as_deref(),
Some("user_created")
);
assert_eq!(
map.source_column("posts", "created_at", false).as_deref(),
Some("createdAt")
);
assert_eq!(map.source_column("users", "email", false), None);
}
#[test]
fn from_cli_arg_loads_a_json_map_and_rejects_garbage() {
assert_eq!(
TransferMap::from_cli_arg("prisma").unwrap(),
TransferMap::Prisma
);
let dir = tempfile::TempDir::new().unwrap();
let path = dir.path().join("map.json");
std::fs::write(
&path,
r#"{ "tables": { "witness": { "witness_merkle_root": "Witness_merkle_root" } } }"#,
)
.unwrap();
let map = TransferMap::from_cli_arg(path.to_str().unwrap()).unwrap();
assert_eq!(
map.source_column("witness", "witness_merkle_root", false)
.as_deref(),
Some("Witness_merkle_root")
);
let err = TransferMap::from_cli_arg("not_a_framework_or_file").unwrap_err();
assert!(err.contains("unknown --map"), "got: {err}");
let bad = dir.path().join("bad.json");
std::fs::write(&bad, "{ not json").unwrap();
assert!(TransferMap::from_cli_arg(bad.to_str().unwrap()).is_err());
}
}