mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-07 18:07:47 +08:00
refactor(data): remove MySQL and SQLite support
Use PostgreSQL as the only database backend across runtime, schema tooling, installation, Compose, and CI. Update regression tests and reject removed drivers explicitly.
This commit is contained in:
@@ -1,25 +1,13 @@
|
||||
#[cfg(feature = "mysql")]
|
||||
mod mysql;
|
||||
#[cfg(feature = "postgres")]
|
||||
mod postgres;
|
||||
#[cfg(feature = "sqlite")]
|
||||
mod sqlite;
|
||||
mod types;
|
||||
|
||||
#[cfg(all(test, feature = "postgres", feature = "mysql", feature = "sqlite"))]
|
||||
#[cfg(all(test, feature = "postgres"))]
|
||||
mod tests;
|
||||
|
||||
#[cfg(feature = "mysql")]
|
||||
pub use mysql::{
|
||||
pending_backfills as pending_mysql_backfills, run_backfills as run_mysql_backfills,
|
||||
};
|
||||
#[cfg(feature = "postgres")]
|
||||
pub use postgres::{pending_backfills, run_backfills};
|
||||
#[cfg(feature = "sqlite")]
|
||||
pub use sqlite::{
|
||||
pending_backfills as pending_sqlite_backfills, run_backfills as run_sqlite_backfills,
|
||||
};
|
||||
pub use types::PendingBackfillInfo;
|
||||
|
||||
#[cfg(all(test, feature = "postgres", feature = "mysql", feature = "sqlite"))]
|
||||
#[cfg(all(test, feature = "postgres"))]
|
||||
use postgres::{pending_backfills_from_applied, AppliedBackfill};
|
||||
|
||||
@@ -1,263 +0,0 @@
|
||||
use std::collections::{HashMap, HashSet};
|
||||
|
||||
use sqlx::{
|
||||
migrate::{Migrate, MigrateError, Migrator},
|
||||
query, query_scalar, Connection, MySqlConnection, Row,
|
||||
};
|
||||
use tracing::{error, info, warn};
|
||||
|
||||
use super::types::PendingBackfillInfo;
|
||||
use crate::driver::mysql::MysqlPool;
|
||||
|
||||
static BACKFILL_MIGRATOR: Migrator = sqlx::migrate!("./backfills/mysql");
|
||||
|
||||
const SCHEMA_BACKFILLS_TABLE_EXISTS_SQL: &str = "SELECT COUNT(*) FROM information_schema.tables WHERE table_schema = DATABASE() AND table_name = 'schema_backfills'";
|
||||
const ENSURE_SCHEMA_BACKFILLS_TABLE_SQL: &str = r#"
|
||||
CREATE TABLE IF NOT EXISTS schema_backfills (
|
||||
version BIGINT NOT NULL,
|
||||
description TEXT NOT NULL,
|
||||
success BOOLEAN NOT NULL DEFAULT TRUE,
|
||||
checksum BLOB NOT NULL,
|
||||
execution_time BIGINT NOT NULL DEFAULT 0,
|
||||
applied_at TIMESTAMP(6) NOT NULL DEFAULT CURRENT_TIMESTAMP(6),
|
||||
PRIMARY KEY (version)
|
||||
)
|
||||
"#;
|
||||
const LIST_APPLIED_BACKFILLS_SQL: &str = r#"
|
||||
SELECT version, checksum
|
||||
FROM schema_backfills
|
||||
WHERE success IS TRUE
|
||||
ORDER BY version ASC
|
||||
"#;
|
||||
const INSERT_APPLIED_BACKFILL_SQL: &str = r#"
|
||||
INSERT INTO schema_backfills (
|
||||
version,
|
||||
description,
|
||||
success,
|
||||
checksum,
|
||||
execution_time,
|
||||
applied_at
|
||||
) VALUES (
|
||||
?,
|
||||
?,
|
||||
TRUE,
|
||||
?,
|
||||
?,
|
||||
CURRENT_TIMESTAMP(6)
|
||||
)
|
||||
ON DUPLICATE KEY UPDATE version = schema_backfills.version
|
||||
"#;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
struct AppliedBackfill {
|
||||
version: i64,
|
||||
checksum: Vec<u8>,
|
||||
}
|
||||
|
||||
pub async fn run_backfills(pool: &MysqlPool) -> Result<(), MigrateError> {
|
||||
let mut conn = pool.acquire().await?;
|
||||
|
||||
if BACKFILL_MIGRATOR.locking {
|
||||
conn.lock().await?;
|
||||
}
|
||||
|
||||
let result = run_backfills_locked(&mut conn).await;
|
||||
|
||||
if BACKFILL_MIGRATOR.locking {
|
||||
match conn.unlock().await {
|
||||
Ok(()) => {}
|
||||
Err(unlock_error) if result.is_ok() => return Err(unlock_error),
|
||||
Err(unlock_error) => {
|
||||
warn!(
|
||||
error = %unlock_error,
|
||||
"mysql database backfill lock release failed after backfill error"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
result
|
||||
}
|
||||
|
||||
pub async fn pending_backfills(pool: &MysqlPool) -> Result<Vec<PendingBackfillInfo>, MigrateError> {
|
||||
let mut conn = pool.acquire().await?;
|
||||
pending_backfills_locked(&mut conn).await
|
||||
}
|
||||
|
||||
async fn run_backfills_locked(conn: &mut MySqlConnection) -> Result<(), MigrateError> {
|
||||
ensure_schema_backfills_table(conn).await?;
|
||||
|
||||
let applied_backfills = list_applied_backfills(conn).await?;
|
||||
validate_applied_backfills(&applied_backfills)?;
|
||||
|
||||
let applied_by_version: HashMap<_, _> = applied_backfills
|
||||
.iter()
|
||||
.map(|backfill| (backfill.version, backfill))
|
||||
.collect();
|
||||
let pending_backfills: Vec<_> = BACKFILL_MIGRATOR
|
||||
.iter()
|
||||
.filter(|backfill| backfill.migration_type.is_up_migration())
|
||||
.filter(|backfill| !applied_by_version.contains_key(&backfill.version))
|
||||
.collect();
|
||||
|
||||
if pending_backfills.is_empty() {
|
||||
info!(
|
||||
driver = "mysql",
|
||||
pending_backfills = 0,
|
||||
"database backfills already up to date"
|
||||
);
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
info!(
|
||||
driver = "mysql",
|
||||
pending_backfills = pending_backfills.len(),
|
||||
"database backfills pending"
|
||||
);
|
||||
|
||||
for (index, backfill) in pending_backfills.iter().enumerate() {
|
||||
let current = index + 1;
|
||||
let total = pending_backfills.len();
|
||||
info!(
|
||||
driver = "mysql",
|
||||
current,
|
||||
total,
|
||||
version = backfill.version,
|
||||
description = %backfill.description,
|
||||
"applying database backfill"
|
||||
);
|
||||
|
||||
let mut tx = conn.begin().await?;
|
||||
let started_at = std::time::Instant::now();
|
||||
sqlx::raw_sql(&backfill.sql).execute(&mut *tx).await?;
|
||||
let elapsed_ms = i64::try_from(started_at.elapsed().as_millis()).unwrap_or(i64::MAX);
|
||||
query(INSERT_APPLIED_BACKFILL_SQL)
|
||||
.bind(backfill.version)
|
||||
.bind(backfill.description.as_ref())
|
||||
.bind(backfill.checksum.as_ref())
|
||||
.bind(elapsed_ms)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
tx.commit().await?;
|
||||
|
||||
info!(
|
||||
driver = "mysql",
|
||||
current,
|
||||
total,
|
||||
version = backfill.version,
|
||||
description = %backfill.description,
|
||||
elapsed_ms,
|
||||
"applied database backfill"
|
||||
);
|
||||
}
|
||||
|
||||
info!(
|
||||
driver = "mysql",
|
||||
pending_backfills = 0,
|
||||
"database backfills complete"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn pending_backfills_locked(
|
||||
conn: &mut MySqlConnection,
|
||||
) -> Result<Vec<PendingBackfillInfo>, MigrateError> {
|
||||
if !schema_backfills_table_exists(conn).await? {
|
||||
return Ok(pending_backfills_from_applied(&[]));
|
||||
}
|
||||
let applied_backfills = list_applied_backfills(conn).await?;
|
||||
validate_applied_backfills(&applied_backfills)?;
|
||||
Ok(pending_backfills_from_applied(&applied_backfills))
|
||||
}
|
||||
|
||||
async fn schema_backfills_table_exists(conn: &mut MySqlConnection) -> Result<bool, MigrateError> {
|
||||
let total: i64 = query_scalar(SCHEMA_BACKFILLS_TABLE_EXISTS_SQL)
|
||||
.fetch_one(&mut *conn)
|
||||
.await?;
|
||||
Ok(total > 0)
|
||||
}
|
||||
|
||||
async fn ensure_schema_backfills_table(conn: &mut MySqlConnection) -> Result<(), MigrateError> {
|
||||
query(ENSURE_SCHEMA_BACKFILLS_TABLE_SQL)
|
||||
.execute(&mut *conn)
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn list_applied_backfills(
|
||||
conn: &mut MySqlConnection,
|
||||
) -> Result<Vec<AppliedBackfill>, MigrateError> {
|
||||
let rows = query(LIST_APPLIED_BACKFILLS_SQL)
|
||||
.fetch_all(&mut *conn)
|
||||
.await?;
|
||||
rows.into_iter()
|
||||
.map(|row| {
|
||||
Ok(AppliedBackfill {
|
||||
version: row.try_get("version")?,
|
||||
checksum: row.try_get("checksum")?,
|
||||
})
|
||||
})
|
||||
.collect::<Result<Vec<_>, sqlx::Error>>()
|
||||
.map_err(MigrateError::from)
|
||||
}
|
||||
|
||||
fn validate_applied_backfills(applied_backfills: &[AppliedBackfill]) -> Result<(), MigrateError> {
|
||||
if BACKFILL_MIGRATOR.ignore_missing {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let known_versions: HashSet<_> = BACKFILL_MIGRATOR
|
||||
.iter()
|
||||
.map(|backfill| backfill.version)
|
||||
.collect();
|
||||
for applied_backfill in applied_backfills {
|
||||
if !known_versions.contains(&applied_backfill.version) {
|
||||
error!(
|
||||
driver = "mysql",
|
||||
version = applied_backfill.version,
|
||||
"applied database backfill is missing from embedded backfills"
|
||||
);
|
||||
return Err(MigrateError::VersionMissing(applied_backfill.version));
|
||||
}
|
||||
}
|
||||
|
||||
for backfill in BACKFILL_MIGRATOR
|
||||
.iter()
|
||||
.filter(|backfill| backfill.migration_type.is_up_migration())
|
||||
{
|
||||
let Some(applied) = applied_backfills
|
||||
.iter()
|
||||
.find(|applied| applied.version == backfill.version)
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
if backfill.checksum != applied.checksum {
|
||||
warn!(
|
||||
driver = "mysql",
|
||||
version = backfill.version,
|
||||
description = %backfill.description,
|
||||
"applied database backfill checksum differs from embedded backfill; skipping strict enforcement"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn pending_backfills_from_applied(
|
||||
applied_backfills: &[AppliedBackfill],
|
||||
) -> Vec<PendingBackfillInfo> {
|
||||
let applied_versions: HashSet<_> = applied_backfills
|
||||
.iter()
|
||||
.map(|backfill| backfill.version)
|
||||
.collect();
|
||||
BACKFILL_MIGRATOR
|
||||
.iter()
|
||||
.filter(|backfill| backfill.migration_type.is_up_migration())
|
||||
.filter(|backfill| !applied_versions.contains(&backfill.version))
|
||||
.map(|backfill| PendingBackfillInfo {
|
||||
version: backfill.version,
|
||||
description: backfill.description.to_string(),
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
@@ -1,265 +0,0 @@
|
||||
use std::collections::{HashMap, HashSet};
|
||||
|
||||
use sqlx::{
|
||||
migrate::{Migrate, MigrateError, Migrator},
|
||||
query, query_scalar, Connection, Row, SqliteConnection,
|
||||
};
|
||||
use tracing::{error, info, warn};
|
||||
|
||||
use super::types::PendingBackfillInfo;
|
||||
use crate::driver::sqlite::SqlitePool;
|
||||
|
||||
static BACKFILL_MIGRATOR: Migrator = sqlx::migrate!("./backfills/sqlite");
|
||||
|
||||
const SCHEMA_BACKFILLS_TABLE_EXISTS_SQL: &str =
|
||||
"SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = 'schema_backfills'";
|
||||
const ENSURE_SCHEMA_BACKFILLS_TABLE_SQL: &str = r#"
|
||||
CREATE TABLE IF NOT EXISTS schema_backfills (
|
||||
version INTEGER NOT NULL PRIMARY KEY,
|
||||
description TEXT NOT NULL,
|
||||
success INTEGER NOT NULL DEFAULT 1,
|
||||
checksum BLOB NOT NULL,
|
||||
execution_time INTEGER NOT NULL DEFAULT 0,
|
||||
applied_at INTEGER NOT NULL DEFAULT (CAST(strftime('%s', 'now') AS INTEGER))
|
||||
)
|
||||
"#;
|
||||
const LIST_APPLIED_BACKFILLS_SQL: &str = r#"
|
||||
SELECT version, checksum
|
||||
FROM schema_backfills
|
||||
WHERE success = 1
|
||||
ORDER BY version ASC
|
||||
"#;
|
||||
const INSERT_APPLIED_BACKFILL_SQL: &str = r#"
|
||||
INSERT INTO schema_backfills (
|
||||
version,
|
||||
description,
|
||||
success,
|
||||
checksum,
|
||||
execution_time,
|
||||
applied_at
|
||||
) VALUES (
|
||||
?,
|
||||
?,
|
||||
1,
|
||||
?,
|
||||
?,
|
||||
CAST(strftime('%s', 'now') AS INTEGER)
|
||||
)
|
||||
ON CONFLICT(version) DO NOTHING
|
||||
"#;
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq)]
|
||||
struct AppliedBackfill {
|
||||
version: i64,
|
||||
checksum: Vec<u8>,
|
||||
}
|
||||
|
||||
pub async fn run_backfills(pool: &SqlitePool) -> Result<(), MigrateError> {
|
||||
let mut conn = pool.acquire().await?;
|
||||
|
||||
if BACKFILL_MIGRATOR.locking {
|
||||
conn.lock().await?;
|
||||
}
|
||||
|
||||
let result = run_backfills_locked(&mut conn).await;
|
||||
|
||||
if BACKFILL_MIGRATOR.locking {
|
||||
match conn.unlock().await {
|
||||
Ok(()) => {}
|
||||
Err(unlock_error) if result.is_ok() => return Err(unlock_error),
|
||||
Err(unlock_error) => {
|
||||
warn!(
|
||||
error = %unlock_error,
|
||||
"sqlite database backfill lock release failed after backfill error"
|
||||
);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
result
|
||||
}
|
||||
|
||||
pub async fn pending_backfills(
|
||||
pool: &SqlitePool,
|
||||
) -> Result<Vec<PendingBackfillInfo>, MigrateError> {
|
||||
let mut conn = pool.acquire().await?;
|
||||
pending_backfills_locked(&mut conn).await
|
||||
}
|
||||
|
||||
async fn run_backfills_locked(conn: &mut SqliteConnection) -> Result<(), MigrateError> {
|
||||
ensure_schema_backfills_table(conn).await?;
|
||||
|
||||
let applied_backfills = list_applied_backfills(conn).await?;
|
||||
validate_applied_backfills(&applied_backfills)?;
|
||||
|
||||
let applied_by_version: HashMap<_, _> = applied_backfills
|
||||
.iter()
|
||||
.map(|backfill| (backfill.version, backfill))
|
||||
.collect();
|
||||
let pending_backfills: Vec<_> = BACKFILL_MIGRATOR
|
||||
.iter()
|
||||
.filter(|backfill| backfill.migration_type.is_up_migration())
|
||||
.filter(|backfill| !applied_by_version.contains_key(&backfill.version))
|
||||
.collect();
|
||||
|
||||
if pending_backfills.is_empty() {
|
||||
info!(
|
||||
driver = "sqlite",
|
||||
pending_backfills = 0,
|
||||
"database backfills already up to date"
|
||||
);
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
info!(
|
||||
driver = "sqlite",
|
||||
pending_backfills = pending_backfills.len(),
|
||||
"database backfills pending"
|
||||
);
|
||||
|
||||
for (index, backfill) in pending_backfills.iter().enumerate() {
|
||||
let current = index + 1;
|
||||
let total = pending_backfills.len();
|
||||
info!(
|
||||
driver = "sqlite",
|
||||
current,
|
||||
total,
|
||||
version = backfill.version,
|
||||
description = %backfill.description,
|
||||
"applying database backfill"
|
||||
);
|
||||
|
||||
let mut tx = conn.begin().await?;
|
||||
let started_at = std::time::Instant::now();
|
||||
sqlx::raw_sql(&backfill.sql).execute(&mut *tx).await?;
|
||||
let elapsed_ms = i64::try_from(started_at.elapsed().as_millis()).unwrap_or(i64::MAX);
|
||||
query(INSERT_APPLIED_BACKFILL_SQL)
|
||||
.bind(backfill.version)
|
||||
.bind(backfill.description.as_ref())
|
||||
.bind(backfill.checksum.as_ref())
|
||||
.bind(elapsed_ms)
|
||||
.execute(&mut *tx)
|
||||
.await?;
|
||||
tx.commit().await?;
|
||||
|
||||
info!(
|
||||
driver = "sqlite",
|
||||
current,
|
||||
total,
|
||||
version = backfill.version,
|
||||
description = %backfill.description,
|
||||
elapsed_ms,
|
||||
"applied database backfill"
|
||||
);
|
||||
}
|
||||
|
||||
info!(
|
||||
driver = "sqlite",
|
||||
pending_backfills = 0,
|
||||
"database backfills complete"
|
||||
);
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn pending_backfills_locked(
|
||||
conn: &mut SqliteConnection,
|
||||
) -> Result<Vec<PendingBackfillInfo>, MigrateError> {
|
||||
if !schema_backfills_table_exists(conn).await? {
|
||||
return Ok(pending_backfills_from_applied(&[]));
|
||||
}
|
||||
let applied_backfills = list_applied_backfills(conn).await?;
|
||||
validate_applied_backfills(&applied_backfills)?;
|
||||
Ok(pending_backfills_from_applied(&applied_backfills))
|
||||
}
|
||||
|
||||
async fn schema_backfills_table_exists(conn: &mut SqliteConnection) -> Result<bool, MigrateError> {
|
||||
let total: i64 = query_scalar(SCHEMA_BACKFILLS_TABLE_EXISTS_SQL)
|
||||
.fetch_one(&mut *conn)
|
||||
.await?;
|
||||
Ok(total > 0)
|
||||
}
|
||||
|
||||
async fn ensure_schema_backfills_table(conn: &mut SqliteConnection) -> Result<(), MigrateError> {
|
||||
query(ENSURE_SCHEMA_BACKFILLS_TABLE_SQL)
|
||||
.execute(&mut *conn)
|
||||
.await?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn list_applied_backfills(
|
||||
conn: &mut SqliteConnection,
|
||||
) -> Result<Vec<AppliedBackfill>, MigrateError> {
|
||||
let rows = query(LIST_APPLIED_BACKFILLS_SQL)
|
||||
.fetch_all(&mut *conn)
|
||||
.await?;
|
||||
rows.into_iter()
|
||||
.map(|row| {
|
||||
Ok(AppliedBackfill {
|
||||
version: row.try_get("version")?,
|
||||
checksum: row.try_get("checksum")?,
|
||||
})
|
||||
})
|
||||
.collect::<Result<Vec<_>, sqlx::Error>>()
|
||||
.map_err(MigrateError::from)
|
||||
}
|
||||
|
||||
fn validate_applied_backfills(applied_backfills: &[AppliedBackfill]) -> Result<(), MigrateError> {
|
||||
if BACKFILL_MIGRATOR.ignore_missing {
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let known_versions: HashSet<_> = BACKFILL_MIGRATOR
|
||||
.iter()
|
||||
.map(|backfill| backfill.version)
|
||||
.collect();
|
||||
for applied_backfill in applied_backfills {
|
||||
if !known_versions.contains(&applied_backfill.version) {
|
||||
error!(
|
||||
driver = "sqlite",
|
||||
version = applied_backfill.version,
|
||||
"applied database backfill is missing from embedded backfills"
|
||||
);
|
||||
return Err(MigrateError::VersionMissing(applied_backfill.version));
|
||||
}
|
||||
}
|
||||
|
||||
for backfill in BACKFILL_MIGRATOR
|
||||
.iter()
|
||||
.filter(|backfill| backfill.migration_type.is_up_migration())
|
||||
{
|
||||
let Some(applied) = applied_backfills
|
||||
.iter()
|
||||
.find(|applied| applied.version == backfill.version)
|
||||
else {
|
||||
continue;
|
||||
};
|
||||
if backfill.checksum != applied.checksum {
|
||||
warn!(
|
||||
driver = "sqlite",
|
||||
version = backfill.version,
|
||||
description = %backfill.description,
|
||||
"applied database backfill checksum differs from embedded backfill; skipping strict enforcement"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn pending_backfills_from_applied(
|
||||
applied_backfills: &[AppliedBackfill],
|
||||
) -> Vec<PendingBackfillInfo> {
|
||||
let applied_versions: HashSet<_> = applied_backfills
|
||||
.iter()
|
||||
.map(|backfill| backfill.version)
|
||||
.collect();
|
||||
BACKFILL_MIGRATOR
|
||||
.iter()
|
||||
.filter(|backfill| backfill.migration_type.is_up_migration())
|
||||
.filter(|backfill| !applied_versions.contains(&backfill.version))
|
||||
.map(|backfill| PendingBackfillInfo {
|
||||
version: backfill.version,
|
||||
description: backfill.description.to_string(),
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
@@ -4,14 +4,10 @@ use std::{
|
||||
time::{Duration, Instant},
|
||||
};
|
||||
|
||||
use sqlx::{query, query_as, query_scalar, Connection, PgConnection, PgPool};
|
||||
use sqlx::{query, query_scalar, Connection, PgConnection, PgPool};
|
||||
|
||||
use super::{
|
||||
pending_backfills, pending_backfills_from_applied, pending_mysql_backfills,
|
||||
pending_sqlite_backfills, run_backfills, run_mysql_backfills, run_sqlite_backfills,
|
||||
AppliedBackfill,
|
||||
};
|
||||
use crate::lifecycle::migrate::{prepare_database_for_startup, run_sqlite_migrations};
|
||||
use super::{pending_backfills, pending_backfills_from_applied, run_backfills, AppliedBackfill};
|
||||
use crate::lifecycle::migrate::prepare_database_for_startup;
|
||||
|
||||
const LEGACY_SYNC_ENABLED_ACTIVE_FLAGS_VERSION: i64 = 20260517012000;
|
||||
const LEGACY_SYNC_ENABLED_ACTIVE_FLAGS_SQL: &str =
|
||||
@@ -88,586 +84,6 @@ fn corrected_legacy_backfill_is_not_requeued_after_application() {
|
||||
assert!(!pending_versions.contains(&LEGACY_SYNC_ENABLED_ACTIVE_FLAGS_VERSION));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mysql_backfills_apply_portable_repairs_when_url_is_set() {
|
||||
let Some(database_url) = std::env::var("AETHER_TEST_MYSQL_URL")
|
||||
.ok()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
else {
|
||||
eprintln!("skipping mysql backfill test because AETHER_TEST_MYSQL_URL is unset");
|
||||
return;
|
||||
};
|
||||
|
||||
let pool = sqlx::mysql::MySqlPoolOptions::new()
|
||||
.max_connections(1)
|
||||
.connect(&database_url)
|
||||
.await
|
||||
.expect("mysql backfill test pool should connect");
|
||||
let mut conn = pool
|
||||
.acquire()
|
||||
.await
|
||||
.expect("mysql backfill test connection should acquire");
|
||||
sqlx::raw_sql(
|
||||
r#"
|
||||
CREATE TEMPORARY TABLE schema_backfills (
|
||||
version BIGINT PRIMARY KEY,
|
||||
description TEXT NOT NULL,
|
||||
success BOOLEAN NOT NULL,
|
||||
checksum BLOB NOT NULL,
|
||||
execution_time BIGINT NOT NULL,
|
||||
applied_at TIMESTAMP(6) NOT NULL DEFAULT CURRENT_TIMESTAMP(6)
|
||||
);
|
||||
CREATE TEMPORARY TABLE api_keys (
|
||||
id VARCHAR(64) PRIMARY KEY,
|
||||
total_requests BIGINT NOT NULL DEFAULT 0,
|
||||
total_tokens BIGINT NOT NULL DEFAULT 0,
|
||||
total_cost_usd DOUBLE NOT NULL DEFAULT 0,
|
||||
last_used_at BIGINT
|
||||
);
|
||||
CREATE TEMPORARY TABLE provider_api_keys (
|
||||
id VARCHAR(64) PRIMARY KEY,
|
||||
total_tokens BIGINT NOT NULL DEFAULT 0
|
||||
);
|
||||
CREATE TEMPORARY TABLE global_models (
|
||||
id VARCHAR(64) PRIMARY KEY,
|
||||
name VARCHAR(255) NOT NULL,
|
||||
usage_count BIGINT NOT NULL DEFAULT 0,
|
||||
updated_at BIGINT NOT NULL
|
||||
);
|
||||
CREATE TEMPORARY TABLE providers (
|
||||
id VARCHAR(64) PRIMARY KEY,
|
||||
enabled BOOLEAN NOT NULL,
|
||||
is_active BOOLEAN NOT NULL
|
||||
);
|
||||
CREATE TEMPORARY TABLE provider_endpoints (
|
||||
id VARCHAR(64) PRIMARY KEY,
|
||||
enabled BOOLEAN NOT NULL,
|
||||
is_active BOOLEAN NOT NULL
|
||||
);
|
||||
CREATE TEMPORARY TABLE models (
|
||||
id VARCHAR(64) PRIMARY KEY,
|
||||
enabled BOOLEAN NOT NULL,
|
||||
is_active BOOLEAN NOT NULL
|
||||
);
|
||||
CREATE TEMPORARY TABLE `usage` (
|
||||
request_id VARCHAR(128) PRIMARY KEY,
|
||||
api_key_id VARCHAR(64),
|
||||
provider_api_key_id VARCHAR(64),
|
||||
model VARCHAR(255),
|
||||
status VARCHAR(64) NOT NULL,
|
||||
total_tokens BIGINT NOT NULL DEFAULT 0,
|
||||
input_tokens BIGINT NOT NULL DEFAULT 0,
|
||||
output_tokens BIGINT NOT NULL DEFAULT 0,
|
||||
cache_creation_input_tokens BIGINT NOT NULL DEFAULT 0,
|
||||
cache_creation_input_tokens_5m BIGINT NOT NULL DEFAULT 0,
|
||||
cache_creation_input_tokens_1h BIGINT NOT NULL DEFAULT 0,
|
||||
cache_creation_ephemeral_5m_input_tokens BIGINT NOT NULL DEFAULT 0,
|
||||
cache_creation_ephemeral_1h_input_tokens BIGINT NOT NULL DEFAULT 0,
|
||||
cache_read_input_tokens BIGINT NOT NULL DEFAULT 0,
|
||||
endpoint_api_format VARCHAR(64),
|
||||
api_format VARCHAR(64),
|
||||
total_cost_usd DOUBLE NOT NULL DEFAULT 0,
|
||||
created_at BIGINT,
|
||||
created_at_unix_ms BIGINT NOT NULL DEFAULT 0,
|
||||
updated_at_unix_secs BIGINT NOT NULL DEFAULT 0
|
||||
);
|
||||
CREATE TEMPORARY TABLE usage_settlement_snapshots (
|
||||
request_id VARCHAR(128) PRIMARY KEY,
|
||||
billing_effective_input_tokens BIGINT,
|
||||
billing_output_tokens BIGINT,
|
||||
billing_cache_creation_tokens BIGINT,
|
||||
billing_cache_creation_5m_tokens BIGINT,
|
||||
billing_cache_creation_1h_tokens BIGINT,
|
||||
billing_cache_read_tokens BIGINT,
|
||||
billing_total_input_context BIGINT
|
||||
);
|
||||
INSERT INTO api_keys (id, total_requests, total_tokens, total_cost_usd)
|
||||
VALUES ('mysql-backfill-api-key', 77, 7777, 77.0);
|
||||
INSERT INTO provider_api_keys (id, total_tokens)
|
||||
VALUES ('mysql-backfill-provider-key', 7777);
|
||||
INSERT INTO global_models (id, name, usage_count, updated_at)
|
||||
VALUES ('mysql-backfill-model', 'gpt-portable', 77, 1);
|
||||
INSERT INTO providers (id, enabled, is_active)
|
||||
VALUES ('mysql-backfill-provider', TRUE, FALSE);
|
||||
INSERT INTO provider_endpoints (id, enabled, is_active)
|
||||
VALUES ('mysql-backfill-endpoint', TRUE, FALSE);
|
||||
INSERT INTO models (id, enabled, is_active)
|
||||
VALUES ('mysql-backfill-provider-model', TRUE, FALSE);
|
||||
INSERT INTO `usage` (
|
||||
request_id,
|
||||
api_key_id,
|
||||
provider_api_key_id,
|
||||
model,
|
||||
status,
|
||||
total_tokens,
|
||||
input_tokens,
|
||||
output_tokens,
|
||||
cache_read_input_tokens,
|
||||
api_format,
|
||||
total_cost_usd,
|
||||
created_at,
|
||||
created_at_unix_ms,
|
||||
updated_at_unix_secs
|
||||
) VALUES
|
||||
(
|
||||
'mysql-backfill-completed',
|
||||
'mysql-backfill-api-key',
|
||||
'mysql-backfill-provider-key',
|
||||
'gpt-portable',
|
||||
'completed',
|
||||
0,
|
||||
120,
|
||||
30,
|
||||
20,
|
||||
'openai',
|
||||
1.25,
|
||||
1714979289,
|
||||
1714979289,
|
||||
1714979289
|
||||
),
|
||||
(
|
||||
'mysql-backfill-pending',
|
||||
'mysql-backfill-api-key',
|
||||
'mysql-backfill-provider-key',
|
||||
'gpt-portable',
|
||||
'pending',
|
||||
777,
|
||||
700,
|
||||
77,
|
||||
0,
|
||||
'openai',
|
||||
0.25,
|
||||
1714979349,
|
||||
1714979349,
|
||||
1714979349
|
||||
);
|
||||
INSERT INTO usage_settlement_snapshots (
|
||||
request_id,
|
||||
billing_effective_input_tokens,
|
||||
billing_output_tokens,
|
||||
billing_cache_creation_tokens,
|
||||
billing_cache_read_tokens
|
||||
) VALUES ('mysql-backfill-completed', 100, 30, 10, 20);
|
||||
"#,
|
||||
)
|
||||
.execute(&mut *conn)
|
||||
.await
|
||||
.expect("mysql temporary backfill schema should initialize");
|
||||
drop(conn);
|
||||
|
||||
let pending_versions = pending_mysql_backfills(&pool)
|
||||
.await
|
||||
.expect("mysql pending backfills should load")
|
||||
.into_iter()
|
||||
.map(|item| item.version)
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(
|
||||
pending_versions,
|
||||
vec![
|
||||
20260422120000,
|
||||
20260505120000,
|
||||
20260517012000,
|
||||
20260716010000
|
||||
]
|
||||
);
|
||||
|
||||
run_mysql_backfills(&pool)
|
||||
.await
|
||||
.expect("mysql backfills should apply");
|
||||
assert!(pending_mysql_backfills(&pool)
|
||||
.await
|
||||
.expect("mysql pending backfills should reload")
|
||||
.is_empty());
|
||||
|
||||
let api_key_stats: (i64, i64, f64, Option<i64>) = query_as(
|
||||
"SELECT total_requests, total_tokens, total_cost_usd, last_used_at FROM api_keys WHERE id = 'mysql-backfill-api-key'",
|
||||
)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("mysql api key backfill result should load");
|
||||
assert_eq!(api_key_stats, (2, 160, 1.5, Some(1714979349)));
|
||||
let provider_total_tokens: i64 = query_scalar(
|
||||
"SELECT total_tokens FROM provider_api_keys WHERE id = 'mysql-backfill-provider-key'",
|
||||
)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("mysql provider key total should load");
|
||||
assert_eq!(provider_total_tokens, 160);
|
||||
let global_usage_count: i64 =
|
||||
query_scalar("SELECT usage_count FROM global_models WHERE id = 'mysql-backfill-model'")
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("mysql global model count should load");
|
||||
assert_eq!(global_usage_count, 1);
|
||||
for table in ["providers", "provider_endpoints", "models"] {
|
||||
let enabled: bool = query_scalar(&format!(
|
||||
"SELECT enabled FROM {table} WHERE is_active = FALSE"
|
||||
))
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.unwrap_or_else(|error| panic!("mysql {table} legacy flag should load: {error}"));
|
||||
assert!(!enabled, "mysql {table}.enabled should follow is_active");
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn pending_sqlite_backfills_does_not_create_tracking_table() {
|
||||
let pool = sqlx::sqlite::SqlitePoolOptions::new()
|
||||
.max_connections(1)
|
||||
.connect("sqlite::memory:")
|
||||
.await
|
||||
.expect("sqlite backfill status pool should connect");
|
||||
run_sqlite_migrations(&pool)
|
||||
.await
|
||||
.expect("sqlite schema should migrate");
|
||||
|
||||
let pending = pending_sqlite_backfills(&pool)
|
||||
.await
|
||||
.expect("sqlite pending backfills should load");
|
||||
assert!(!pending.is_empty());
|
||||
|
||||
let tracking_tables: i64 = query_scalar(
|
||||
"SELECT COUNT(*) FROM sqlite_master WHERE type = 'table' AND name = 'schema_backfills'",
|
||||
)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("sqlite tracking table state should load");
|
||||
assert_eq!(tracking_tables, 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn sqlite_backfills_apply_portable_repairs_and_record_versions() {
|
||||
let pool = sqlx::sqlite::SqlitePoolOptions::new()
|
||||
.max_connections(1)
|
||||
.connect("sqlite::memory:")
|
||||
.await
|
||||
.expect("sqlite backfill test pool should connect");
|
||||
run_sqlite_migrations(&pool)
|
||||
.await
|
||||
.expect("sqlite schema should migrate");
|
||||
|
||||
query(
|
||||
r#"
|
||||
INSERT INTO api_keys (
|
||||
id, user_id, key_hash, total_requests, total_tokens, total_cost_usd, created_at, updated_at
|
||||
) VALUES (
|
||||
'sqlite-backfill-api-key', 'sqlite-backfill-user', 'sqlite-backfill-hash',
|
||||
77, 7777, 77.0, 1, 1
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("sqlite api key fixture should insert");
|
||||
query(
|
||||
r#"
|
||||
INSERT INTO provider_api_keys (
|
||||
id, provider_id, name, total_tokens, created_at, updated_at
|
||||
) VALUES (
|
||||
'sqlite-backfill-provider-key', 'sqlite-backfill-provider', 'Portable key', 7777, 1, 1
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("sqlite provider key fixture should insert");
|
||||
query(
|
||||
r#"
|
||||
INSERT INTO global_models (
|
||||
id, name, display_name, usage_count, created_at, updated_at
|
||||
) VALUES (
|
||||
'sqlite-backfill-model', 'gpt-portable', 'GPT Portable', 77, 1, 1
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("sqlite global model fixture should insert");
|
||||
query(
|
||||
r#"
|
||||
INSERT INTO providers (
|
||||
id, name, provider_type, enabled, is_active, created_at, updated_at
|
||||
) VALUES (
|
||||
'sqlite-backfill-provider', 'SQLite Backfill Provider', 'openai', 1, 0, 1, 1
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("sqlite provider flag fixture should insert");
|
||||
query(
|
||||
r#"
|
||||
INSERT INTO provider_endpoints (
|
||||
id, provider_id, name, base_url, enabled, is_active, created_at, updated_at
|
||||
) VALUES (
|
||||
'sqlite-backfill-endpoint', 'sqlite-backfill-provider', 'Default',
|
||||
'https://example.invalid', 1, 0, 1, 1
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("sqlite provider endpoint flag fixture should insert");
|
||||
query(
|
||||
r#"
|
||||
INSERT INTO models (
|
||||
id, provider_id, provider_model_name, enabled, is_active, created_at, updated_at
|
||||
) VALUES (
|
||||
'sqlite-backfill-provider-model', 'sqlite-backfill-provider', 'gpt-portable',
|
||||
1, 0, 1, 1
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("sqlite model flag fixture should insert");
|
||||
query(
|
||||
r#"
|
||||
INSERT INTO "usage" (
|
||||
request_id,
|
||||
api_key_id,
|
||||
provider_api_key_id,
|
||||
model,
|
||||
status,
|
||||
total_tokens,
|
||||
input_tokens,
|
||||
output_tokens,
|
||||
cache_read_input_tokens,
|
||||
api_format,
|
||||
total_cost_usd,
|
||||
created_at,
|
||||
created_at_unix_ms,
|
||||
updated_at_unix_secs
|
||||
) VALUES
|
||||
(
|
||||
'sqlite-backfill-completed',
|
||||
'sqlite-backfill-api-key',
|
||||
'sqlite-backfill-provider-key',
|
||||
'gpt-portable',
|
||||
'completed',
|
||||
0,
|
||||
120,
|
||||
30,
|
||||
20,
|
||||
'openai',
|
||||
1.25,
|
||||
1714979289,
|
||||
1714979289,
|
||||
1714979289
|
||||
),
|
||||
(
|
||||
'sqlite-backfill-pending',
|
||||
'sqlite-backfill-api-key',
|
||||
'sqlite-backfill-provider-key',
|
||||
'gpt-portable',
|
||||
'pending',
|
||||
777,
|
||||
700,
|
||||
77,
|
||||
0,
|
||||
'openai',
|
||||
0.25,
|
||||
1714979349,
|
||||
1714979349,
|
||||
1714979349
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("sqlite usage fixtures should insert");
|
||||
query(
|
||||
r#"
|
||||
INSERT INTO usage_settlement_snapshots (
|
||||
request_id,
|
||||
billing_status,
|
||||
billing_effective_input_tokens,
|
||||
billing_output_tokens,
|
||||
billing_cache_creation_tokens,
|
||||
billing_cache_read_tokens,
|
||||
created_at,
|
||||
updated_at
|
||||
) VALUES (
|
||||
'sqlite-backfill-completed', 'settled', 100, 30, 10, 20, 1, 1
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("sqlite settlement fixture should insert");
|
||||
|
||||
let pending_versions = pending_sqlite_backfills(&pool)
|
||||
.await
|
||||
.expect("sqlite pending backfills should load")
|
||||
.into_iter()
|
||||
.map(|item| item.version)
|
||||
.collect::<Vec<_>>();
|
||||
assert_eq!(
|
||||
pending_versions,
|
||||
vec![
|
||||
20260422120000,
|
||||
20260505120000,
|
||||
20260517012000,
|
||||
20260716010000
|
||||
]
|
||||
);
|
||||
|
||||
run_sqlite_backfills(&pool)
|
||||
.await
|
||||
.expect("sqlite backfills should apply");
|
||||
assert!(pending_sqlite_backfills(&pool)
|
||||
.await
|
||||
.expect("sqlite pending backfills should reload")
|
||||
.is_empty());
|
||||
|
||||
let applied_versions: Vec<i64> =
|
||||
query_scalar("SELECT version FROM schema_backfills ORDER BY version")
|
||||
.fetch_all(&pool)
|
||||
.await
|
||||
.expect("sqlite applied backfill versions should load");
|
||||
assert_eq!(
|
||||
applied_versions,
|
||||
vec![
|
||||
20260422120000,
|
||||
20260505120000,
|
||||
20260517012000,
|
||||
20260716010000
|
||||
]
|
||||
);
|
||||
let api_key_stats: (i64, i64, f64, Option<i64>) = query_as(
|
||||
"SELECT total_requests, total_tokens, total_cost_usd, last_used_at FROM api_keys WHERE id = 'sqlite-backfill-api-key'",
|
||||
)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("sqlite api key backfill result should load");
|
||||
assert_eq!(api_key_stats, (2, 160, 1.5, Some(1714979349)));
|
||||
let provider_total_tokens: i64 = query_scalar(
|
||||
"SELECT total_tokens FROM provider_api_keys WHERE id = 'sqlite-backfill-provider-key'",
|
||||
)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("sqlite provider key total should load");
|
||||
assert_eq!(provider_total_tokens, 160);
|
||||
let global_usage_count: i64 =
|
||||
query_scalar("SELECT usage_count FROM global_models WHERE id = 'sqlite-backfill-model'")
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("sqlite global model count should load");
|
||||
assert_eq!(global_usage_count, 1);
|
||||
for table in ["providers", "provider_endpoints", "models"] {
|
||||
let enabled: i64 =
|
||||
query_scalar(&format!("SELECT enabled FROM {table} WHERE is_active = 0"))
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.unwrap_or_else(|error| panic!("sqlite {table} legacy flag should load: {error}"));
|
||||
assert_eq!(enabled, 0, "sqlite {table}.enabled should follow is_active");
|
||||
}
|
||||
|
||||
run_sqlite_backfills(&pool)
|
||||
.await
|
||||
.expect("sqlite backfills should be idempotent");
|
||||
let applied_count: i64 = query_scalar("SELECT COUNT(*) FROM schema_backfills")
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("sqlite applied backfill count should load");
|
||||
assert_eq!(applied_count, 4);
|
||||
|
||||
query("UPDATE schema_backfills SET checksum = X'00' WHERE version = 20260422120000")
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("sqlite checksum compatibility fixture should update");
|
||||
assert!(pending_sqlite_backfills(&pool)
|
||||
.await
|
||||
.expect("checksum drift should retain the postgres compatibility policy")
|
||||
.is_empty());
|
||||
|
||||
query(
|
||||
r#"
|
||||
INSERT INTO schema_backfills (
|
||||
version, description, success, checksum, execution_time
|
||||
) VALUES (
|
||||
99999999999999, 'missing embedded backfill', 1, X'', 0
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("unknown sqlite backfill fixture should insert");
|
||||
let error = pending_sqlite_backfills(&pool)
|
||||
.await
|
||||
.expect_err("unknown applied sqlite backfill should fail validation");
|
||||
assert!(matches!(
|
||||
error,
|
||||
sqlx::migrate::MigrateError::VersionMissing(99999999999999)
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn sqlite_backfill_sql_and_version_record_commit_atomically() {
|
||||
let pool = sqlx::sqlite::SqlitePoolOptions::new()
|
||||
.max_connections(1)
|
||||
.connect("sqlite::memory:")
|
||||
.await
|
||||
.expect("sqlite backfill transaction test pool should connect");
|
||||
run_sqlite_migrations(&pool)
|
||||
.await
|
||||
.expect("sqlite schema should migrate");
|
||||
query(
|
||||
r#"
|
||||
INSERT INTO global_models (
|
||||
id, name, display_name, usage_count, created_at, updated_at
|
||||
) VALUES (
|
||||
'sqlite-backfill-rollback-model', 'rollback-model', 'Rollback Model', 77, 1, 1
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("sqlite rollback global model fixture should insert");
|
||||
query(
|
||||
r#"
|
||||
CREATE TRIGGER reject_global_model_backfill
|
||||
BEFORE UPDATE OF usage_count ON global_models
|
||||
BEGIN
|
||||
SELECT RAISE(ABORT, 'forced global model backfill failure');
|
||||
END
|
||||
"#,
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("sqlite rollback trigger should create");
|
||||
|
||||
run_sqlite_backfills(&pool)
|
||||
.await
|
||||
.expect_err("forced sqlite backfill failure should propagate");
|
||||
let applied_versions: Vec<i64> =
|
||||
query_scalar("SELECT version FROM schema_backfills ORDER BY version")
|
||||
.fetch_all(&pool)
|
||||
.await
|
||||
.expect("sqlite partial applied versions should load");
|
||||
assert_eq!(applied_versions, vec![20260422120000]);
|
||||
let usage_count: i64 = query_scalar(
|
||||
"SELECT usage_count FROM global_models WHERE id = 'sqlite-backfill-rollback-model'",
|
||||
)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("sqlite rolled back global model should load");
|
||||
assert_eq!(usage_count, 77);
|
||||
|
||||
query("DROP TRIGGER reject_global_model_backfill")
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("sqlite rollback trigger should drop");
|
||||
run_sqlite_backfills(&pool)
|
||||
.await
|
||||
.expect("sqlite backfills should resume after the failed transaction");
|
||||
let applied_count: i64 = query_scalar("SELECT COUNT(*) FROM schema_backfills")
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("sqlite resumed applied backfill count should load");
|
||||
assert_eq!(applied_count, 4);
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct ManagedPostgresServer {
|
||||
child: Option<Child>,
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,914 +0,0 @@
|
||||
use super::*;
|
||||
use sqlx::Acquire;
|
||||
|
||||
#[derive(Debug, Clone, Default, PartialEq, Eq)]
|
||||
struct MysqlImportColumns {
|
||||
names: ImportColumnNames,
|
||||
data_types: BTreeMap<String, String>,
|
||||
primary_key: Vec<String>,
|
||||
}
|
||||
|
||||
pub async fn export_mysql_core_jsonl(
|
||||
pool: &crate::driver::mysql::MysqlPool,
|
||||
created_at_unix_secs: u64,
|
||||
) -> Result<String, DataLayerError> {
|
||||
export_mysql_jsonl(pool, mysql_core_export_domains(), created_at_unix_secs).await
|
||||
}
|
||||
|
||||
pub async fn export_mysql_jsonl(
|
||||
pool: &crate::driver::mysql::MysqlPool,
|
||||
domains: Vec<ExportDomain>,
|
||||
created_at_unix_secs: u64,
|
||||
) -> Result<String, DataLayerError> {
|
||||
let mut connection = pool.acquire().await.map_sql_err()?;
|
||||
sqlx::query("SET TRANSACTION ISOLATION LEVEL REPEATABLE READ")
|
||||
.execute(&mut *connection)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let mut tx = connection.begin().await.map_sql_err()?;
|
||||
let manifest = DataExportManifest::new(
|
||||
created_at_unix_secs,
|
||||
Some(DatabaseDriver::Mysql),
|
||||
domains.clone(),
|
||||
);
|
||||
let mut records = vec![DataExportRecord::manifest(manifest)];
|
||||
|
||||
for domain in domains {
|
||||
if domain == ExportDomain::Auxiliary {
|
||||
export_mysql_auxiliary_records(&mut tx, &mut records).await?;
|
||||
continue;
|
||||
}
|
||||
if domain == ExportDomain::Billing {
|
||||
export_mysql_billing_records(&mut tx, &mut records).await?;
|
||||
continue;
|
||||
}
|
||||
if domain == ExportDomain::Wallets {
|
||||
export_mysql_wallet_records(&mut tx, &mut records).await?;
|
||||
continue;
|
||||
}
|
||||
let (table_name, id_column) = mysql_domain_table(domain)?;
|
||||
let order_by = export_order_by(domain, id_column);
|
||||
let sql = format!("SELECT * FROM {table_name} ORDER BY {order_by}");
|
||||
let rows = sqlx::query(&sql).fetch_all(&mut *tx).await.map_sql_err()?;
|
||||
for row in rows {
|
||||
let id = mysql_export_row_id(domain, &row, id_column)?;
|
||||
records.push(DataExportRecord::row(domain, id, mysql_row_payload(&row)?));
|
||||
}
|
||||
}
|
||||
|
||||
tx.commit().await.map_sql_err()?;
|
||||
encode_jsonl(&records)
|
||||
}
|
||||
|
||||
pub async fn import_mysql_jsonl(
|
||||
pool: &crate::driver::mysql::MysqlPool,
|
||||
input: &str,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
let plan = build_import_plan(input)?;
|
||||
import_mysql_plan(pool, &plan).await
|
||||
}
|
||||
|
||||
pub async fn import_mysql_plan(
|
||||
pool: &crate::driver::mysql::MysqlPool,
|
||||
plan: &DataImportPlan,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
let identity_scope = IdentityImportScope::from_plan(plan)?;
|
||||
let mut tx = pool.begin().await.map_sql_err()?;
|
||||
let identity_state = capture_mysql_identity_import_state(&mut tx, &identity_scope).await?;
|
||||
let mut imported = 0usize;
|
||||
let mut column_cache = BTreeMap::<String, MysqlImportColumns>::new();
|
||||
for domain in &plan.manifest.domains {
|
||||
if *domain == ExportDomain::Auxiliary {
|
||||
for row in plan.rows(*domain) {
|
||||
import_mysql_auxiliary_row(&mut tx, row, &mut column_cache).await?;
|
||||
imported = imported.saturating_add(1);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
if *domain == ExportDomain::Billing {
|
||||
for row in plan.rows(*domain) {
|
||||
import_mysql_billing_row(&mut tx, row, &mut column_cache).await?;
|
||||
imported = imported.saturating_add(1);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
if *domain == ExportDomain::Wallets {
|
||||
for row in plan.rows(*domain) {
|
||||
import_mysql_wallet_row(&mut tx, row, &mut column_cache).await?;
|
||||
imported = imported.saturating_add(1);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
let (table_name, _id_column) = mysql_domain_table(*domain)?;
|
||||
let target_columns =
|
||||
mysql_import_columns_cached(&mut tx, &mut column_cache, table_name).await?;
|
||||
for row in plan.rows(*domain) {
|
||||
import_mysql_row(&mut tx, table_name, *domain, row, &target_columns).await?;
|
||||
imported = imported.saturating_add(1);
|
||||
}
|
||||
}
|
||||
enforce_mysql_identity_import_invariants(&mut tx, &identity_scope, identity_state).await?;
|
||||
tx.commit().await.map_sql_err()?;
|
||||
Ok(imported)
|
||||
}
|
||||
|
||||
async fn capture_mysql_identity_import_state(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
scope: &IdentityImportScope,
|
||||
) -> Result<IdentityImportState, DataLayerError> {
|
||||
let mut affected_user_ids = if scope.finalizes_oauth_links {
|
||||
scope.user_ids.iter().cloned().collect::<BTreeSet<_>>()
|
||||
} else {
|
||||
BTreeSet::new()
|
||||
};
|
||||
for link_id in &scope.oauth_link_ids {
|
||||
if let Some(user_id) =
|
||||
sqlx::query_scalar::<_, String>("SELECT user_id FROM user_oauth_links WHERE id = ?")
|
||||
.bind(link_id)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
{
|
||||
affected_user_ids.insert(user_id);
|
||||
}
|
||||
}
|
||||
for provider_type in &scope.oauth_provider_types {
|
||||
let user_ids = sqlx::query_scalar::<_, String>(
|
||||
"SELECT user_id FROM user_oauth_links WHERE provider_type = ?",
|
||||
)
|
||||
.bind(provider_type)
|
||||
.fetch_all(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
affected_user_ids.extend(user_ids);
|
||||
}
|
||||
Ok(IdentityImportState { affected_user_ids })
|
||||
}
|
||||
|
||||
async fn enforce_mysql_identity_import_invariants(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
scope: &IdentityImportScope,
|
||||
mut state: IdentityImportState,
|
||||
) -> Result<(), DataLayerError> {
|
||||
for user_id in &scope.user_ids {
|
||||
let auth_source =
|
||||
sqlx::query_scalar::<_, String>("SELECT auth_source FROM users WHERE id = ? LIMIT 1")
|
||||
.bind(user_id)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::InvalidInput(format!(
|
||||
"imported users row '{user_id}' did not produce a user record"
|
||||
))
|
||||
})?;
|
||||
if !matches!(auth_source.as_str(), "local" | "ldap" | "oauth") {
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"imported user '{user_id}' has unsupported auth_source '{auth_source}'"
|
||||
)));
|
||||
}
|
||||
if auth_source == "oauth" {
|
||||
sqlx::query("UPDATE users SET email_verified = 0 WHERE id = ?")
|
||||
.bind(user_id)
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
}
|
||||
}
|
||||
|
||||
for link_id in &scope.oauth_link_ids {
|
||||
let user_id = sqlx::query_scalar::<_, String>(
|
||||
"SELECT user_id FROM user_oauth_links WHERE id = ? LIMIT 1",
|
||||
)
|
||||
.bind(link_id)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::InvalidInput(format!(
|
||||
"imported OAuth link row '{link_id}' did not produce a link record"
|
||||
))
|
||||
})?;
|
||||
state.affected_user_ids.insert(user_id);
|
||||
}
|
||||
|
||||
for link_id in &scope.oauth_link_ids {
|
||||
if let Some((provider_type, provider_user_id)) = sqlx::query_as::<_, (String, String)>(
|
||||
r#"
|
||||
SELECT imported.provider_type, imported.provider_user_id
|
||||
FROM user_oauth_links imported
|
||||
JOIN user_oauth_links duplicate
|
||||
ON duplicate.provider_type = imported.provider_type
|
||||
AND duplicate.provider_user_id = imported.provider_user_id
|
||||
AND duplicate.id <> imported.id
|
||||
WHERE imported.id = ?
|
||||
LIMIT 1
|
||||
"#,
|
||||
)
|
||||
.bind(link_id)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
{
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"OAuth import assigns provider identity '{provider_type}:{provider_user_id}' more than once"
|
||||
)));
|
||||
}
|
||||
}
|
||||
|
||||
for link_id in &scope.oauth_link_ids {
|
||||
if let Some((user_id, provider_type)) = sqlx::query_as::<_, (String, String)>(
|
||||
r#"
|
||||
SELECT imported.user_id, imported.provider_type
|
||||
FROM user_oauth_links imported
|
||||
JOIN user_oauth_links duplicate
|
||||
ON duplicate.user_id = imported.user_id
|
||||
AND duplicate.provider_type = imported.provider_type
|
||||
AND duplicate.id <> imported.id
|
||||
WHERE imported.id = ?
|
||||
LIMIT 1
|
||||
"#,
|
||||
)
|
||||
.bind(link_id)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
{
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"OAuth import links user '{user_id}' to provider '{provider_type}' more than once"
|
||||
)));
|
||||
}
|
||||
}
|
||||
|
||||
for link_id in &scope.oauth_link_ids {
|
||||
if let Some(invalid_id) = sqlx::query_scalar::<_, String>(
|
||||
r#"
|
||||
SELECT links.id
|
||||
FROM user_oauth_links links
|
||||
LEFT JOIN users ON users.id = links.user_id
|
||||
LEFT JOIN oauth_providers providers ON providers.provider_type = links.provider_type
|
||||
WHERE links.id = ?
|
||||
AND (
|
||||
users.id IS NULL
|
||||
OR providers.provider_type IS NULL
|
||||
OR BINARY links.provider_type <> BINARY LOWER(TRIM(links.provider_type))
|
||||
OR links.provider_type = ''
|
||||
OR BINARY links.provider_user_id <> BINARY TRIM(links.provider_user_id)
|
||||
OR links.provider_user_id = ''
|
||||
OR BINARY providers.provider_type <> BINARY LOWER(TRIM(providers.provider_type))
|
||||
)
|
||||
LIMIT 1
|
||||
"#,
|
||||
)
|
||||
.bind(link_id)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
{
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"OAuth import produced invalid or orphaned link '{invalid_id}'"
|
||||
)));
|
||||
}
|
||||
}
|
||||
|
||||
if !scope.validates_oauth_login_methods {
|
||||
return Ok(());
|
||||
}
|
||||
for user_id in state.affected_user_ids {
|
||||
if sqlx::query_scalar::<_, String>(
|
||||
r#"
|
||||
SELECT users.id
|
||||
FROM users
|
||||
WHERE users.id = ?
|
||||
AND users.auth_source = 'oauth'
|
||||
AND users.is_active = 1
|
||||
AND users.is_deleted = 0
|
||||
AND NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM user_oauth_links links
|
||||
JOIN oauth_providers providers ON providers.provider_type = links.provider_type
|
||||
WHERE links.user_id = users.id
|
||||
AND providers.is_enabled = 1
|
||||
AND BINARY links.provider_type = BINARY LOWER(TRIM(links.provider_type))
|
||||
AND BINARY links.provider_user_id = BINARY TRIM(links.provider_user_id)
|
||||
AND links.provider_user_id <> ''
|
||||
)
|
||||
LIMIT 1
|
||||
"#,
|
||||
)
|
||||
.bind(&user_id)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.is_some()
|
||||
{
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"OAuth import would leave active user '{user_id}' without an enabled identity binding"
|
||||
)));
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn mysql_domain_table(
|
||||
domain: ExportDomain,
|
||||
) -> Result<(&'static str, &'static str), DataLayerError> {
|
||||
match domain {
|
||||
ExportDomain::Users => Ok(("users", "id")),
|
||||
ExportDomain::ApiKeys => Ok(("api_keys", "id")),
|
||||
ExportDomain::Providers => Ok(("providers", "id")),
|
||||
ExportDomain::ProviderKeys => Ok(("provider_api_keys", "id")),
|
||||
ExportDomain::Endpoints => Ok(("provider_endpoints", "id")),
|
||||
ExportDomain::Models => Ok(("models", "id")),
|
||||
ExportDomain::GlobalModels => Ok(("global_models", "id")),
|
||||
ExportDomain::AuthModules => Ok(("auth_modules", "id")),
|
||||
ExportDomain::OAuthProviders => Ok(("oauth_providers", "provider_type")),
|
||||
ExportDomain::UserOAuthLinks => Ok(("user_oauth_links", "id")),
|
||||
ExportDomain::UserGroups => Ok(("user_groups", "id")),
|
||||
ExportDomain::UserGroupMembers => Ok(("user_group_members", "group_id")),
|
||||
ExportDomain::ProxyNodes => Ok(("proxy_nodes", "id")),
|
||||
ExportDomain::SystemConfigs => Ok(("system_configs", "id")),
|
||||
ExportDomain::Wallets => Err(DataLayerError::InvalidInput(
|
||||
"mysql wallet export uses multiple tables and must be handled as a domain".to_string(),
|
||||
)),
|
||||
ExportDomain::Usage => Ok(("`usage`", "request_id")),
|
||||
ExportDomain::Billing => Err(DataLayerError::InvalidInput(
|
||||
"mysql billing export uses multiple tables and must be handled as a domain".to_string(),
|
||||
)),
|
||||
ExportDomain::Auxiliary => Err(DataLayerError::InvalidInput(
|
||||
"mysql auxiliary export uses multiple tables and must be handled as a domain"
|
||||
.to_string(),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
async fn export_mysql_auxiliary_records(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
records: &mut Vec<DataExportRecord>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
for table in AUXILIARY_TABLES {
|
||||
let table_sql = mysql_quote_identifier(table.name)?;
|
||||
let order_sql = table
|
||||
.primary_key
|
||||
.iter()
|
||||
.map(|column| mysql_quote_identifier(column).map(|column| format!("{column} ASC")))
|
||||
.collect::<Result<Vec<_>, _>>()?
|
||||
.join(", ");
|
||||
let rows = sqlx::query(&format!("SELECT * FROM {table_sql} ORDER BY {order_sql}"))
|
||||
.fetch_all(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
for row in rows {
|
||||
let payload = mysql_row_payload(&row)?;
|
||||
let id = auxiliary_row_id(*table, &payload)?;
|
||||
records.push(DataExportRecord::row(
|
||||
ExportDomain::Auxiliary,
|
||||
id,
|
||||
payload_with_table(payload, table.name)?,
|
||||
));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn mysql_export_row_id(
|
||||
domain: ExportDomain,
|
||||
row: &sqlx::mysql::MySqlRow,
|
||||
id_column: &str,
|
||||
) -> Result<String, DataLayerError> {
|
||||
if domain == ExportDomain::UserGroupMembers {
|
||||
let group_id = mysql_required_export_text(row, "group_id", domain)?;
|
||||
let user_id = mysql_required_export_text(row, "user_id", domain)?;
|
||||
return Ok(format!("{group_id}:{user_id}"));
|
||||
}
|
||||
mysql_required_export_text(row, id_column, domain)
|
||||
}
|
||||
|
||||
fn mysql_required_export_text(
|
||||
row: &sqlx::mysql::MySqlRow,
|
||||
column: &str,
|
||||
domain: ExportDomain,
|
||||
) -> Result<String, DataLayerError> {
|
||||
row.try_get::<Option<String>, _>(column)
|
||||
.map_sql_err()?
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"{} export row has null id column '{}'",
|
||||
domain.as_str(),
|
||||
column
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
async fn export_mysql_billing_records(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
records: &mut Vec<DataExportRecord>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
for (table_name, id_column) in [
|
||||
("billing_rules", "id"),
|
||||
("dimension_collectors", "id"),
|
||||
("usage_settlement_snapshots", "request_id"),
|
||||
] {
|
||||
let sql = format!("SELECT * FROM {table_name} ORDER BY {id_column} ASC");
|
||||
let rows = sqlx::query(&sql).fetch_all(&mut **tx).await.map_sql_err()?;
|
||||
for row in rows {
|
||||
let id = row
|
||||
.try_get::<Option<String>, _>(id_column)
|
||||
.map_sql_err()?
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"billing export row in table '{table_name}' has null id"
|
||||
))
|
||||
})?;
|
||||
records.push(DataExportRecord::row(
|
||||
ExportDomain::Billing,
|
||||
format!("{table_name}:{id}"),
|
||||
payload_with_table(mysql_row_payload(&row)?, table_name)?,
|
||||
));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn export_mysql_wallet_records(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
records: &mut Vec<DataExportRecord>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
for (table_name, id_column) in mysql_wallet_tables() {
|
||||
let sql = format!("SELECT * FROM {table_name} ORDER BY {id_column} ASC");
|
||||
let rows = sqlx::query(&sql).fetch_all(&mut **tx).await.map_sql_err()?;
|
||||
for row in rows {
|
||||
let id = row
|
||||
.try_get::<Option<String>, _>(id_column)
|
||||
.map_sql_err()?
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"wallet export row in table '{table_name}' has null id"
|
||||
))
|
||||
})?;
|
||||
records.push(DataExportRecord::row(
|
||||
ExportDomain::Wallets,
|
||||
format!("{table_name}:{id}"),
|
||||
payload_with_table(mysql_row_payload(&row)?, table_name)?,
|
||||
));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn import_mysql_row(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
table_name: &str,
|
||||
domain: ExportDomain,
|
||||
row: &ExportRow,
|
||||
target_columns: &MysqlImportColumns,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let mut object =
|
||||
filter_import_payload("mysql", table_name, domain, row, &target_columns.names)?;
|
||||
deactivate_imported_credentials(table_name, &mut object, |column_name| {
|
||||
target_columns.names.contains(column_name)
|
||||
});
|
||||
|
||||
let columns = object.keys().map(String::as_str).collect::<Vec<_>>();
|
||||
for primary_key in &target_columns.primary_key {
|
||||
if object.get(primary_key).is_none_or(Value::is_null) {
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"{} export row '{}' is missing non-null primary key column '{}' for mysql table '{}'",
|
||||
domain.as_str(),
|
||||
row.id,
|
||||
primary_key,
|
||||
table_name
|
||||
)));
|
||||
}
|
||||
}
|
||||
|
||||
let primary_key_predicate = target_columns
|
||||
.primary_key
|
||||
.iter()
|
||||
.map(|column| mysql_quote_identifier(column).map(|column| format!("{column} = ?")))
|
||||
.collect::<Result<Vec<_>, _>>()?
|
||||
.join(" AND ");
|
||||
let lock_sql =
|
||||
format!("SELECT 1 FROM {table_name} WHERE {primary_key_predicate} LIMIT 1 FOR UPDATE");
|
||||
let mut lock_query = sqlx::query(&lock_sql);
|
||||
for column in &target_columns.primary_key {
|
||||
lock_query =
|
||||
bind_mysql_import_column(lock_query, &object, target_columns, table_name, column)?;
|
||||
}
|
||||
let exists = lock_query
|
||||
.fetch_optional(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.is_some();
|
||||
|
||||
if exists {
|
||||
let update_columns = columns
|
||||
.iter()
|
||||
.copied()
|
||||
.filter(|column| !target_columns.primary_key.iter().any(|key| key == column))
|
||||
.collect::<Vec<_>>();
|
||||
if update_columns.is_empty() {
|
||||
return Ok(());
|
||||
}
|
||||
let update_sql = update_columns
|
||||
.iter()
|
||||
.map(|column| mysql_quote_identifier(column).map(|column| format!("{column} = ?")))
|
||||
.collect::<Result<Vec<_>, _>>()?
|
||||
.join(", ");
|
||||
let sql = format!("UPDATE {table_name} SET {update_sql} WHERE {primary_key_predicate}");
|
||||
let mut query = sqlx::query(&sql);
|
||||
for column in update_columns {
|
||||
query = bind_mysql_import_column(query, &object, target_columns, table_name, column)?;
|
||||
}
|
||||
for column in &target_columns.primary_key {
|
||||
query = bind_mysql_import_column(query, &object, target_columns, table_name, column)?;
|
||||
}
|
||||
query.execute(&mut **tx).await.map_sql_err()?;
|
||||
return Ok(());
|
||||
}
|
||||
|
||||
let column_sql = columns
|
||||
.iter()
|
||||
.map(|column| mysql_quote_identifier(column))
|
||||
.collect::<Result<Vec<_>, _>>()?
|
||||
.join(", ");
|
||||
let placeholder_sql = vec!["?"; columns.len()].join(", ");
|
||||
let sql = format!("INSERT INTO {table_name} ({column_sql}) VALUES ({placeholder_sql})");
|
||||
let mut query = sqlx::query(&sql);
|
||||
for column in columns {
|
||||
query = bind_mysql_import_column(query, &object, target_columns, table_name, column)?;
|
||||
}
|
||||
query.execute(&mut **tx).await.map_sql_err()?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn bind_mysql_import_column<'q>(
|
||||
query: sqlx::query::Query<'q, sqlx::MySql, sqlx::mysql::MySqlArguments>,
|
||||
object: &'q serde_json::Map<String, Value>,
|
||||
target_columns: &MysqlImportColumns,
|
||||
table_name: &str,
|
||||
column: &str,
|
||||
) -> Result<sqlx::query::Query<'q, sqlx::MySql, sqlx::mysql::MySqlArguments>, DataLayerError> {
|
||||
let value = object
|
||||
.get(column)
|
||||
.expect("column name came from payload object keys");
|
||||
let data_type = target_columns
|
||||
.data_types
|
||||
.get(column)
|
||||
.map(String::as_str)
|
||||
.unwrap_or_default();
|
||||
bind_mysql_import_value(query, value, table_name, column, data_type)
|
||||
}
|
||||
|
||||
async fn import_mysql_billing_row(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
row: &ExportRow,
|
||||
column_cache: &mut BTreeMap<String, MysqlImportColumns>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let (table_name, payload) = billing_payload_table(row)?;
|
||||
let table_name = mysql_billing_table_name(&table_name)?;
|
||||
let target_columns = mysql_import_columns_cached(tx, column_cache, table_name).await?;
|
||||
import_mysql_row(
|
||||
tx,
|
||||
table_name,
|
||||
ExportDomain::Billing,
|
||||
&ExportRow {
|
||||
id: row.id.clone(),
|
||||
payload,
|
||||
},
|
||||
&target_columns,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn import_mysql_auxiliary_row(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
row: &ExportRow,
|
||||
column_cache: &mut BTreeMap<String, MysqlImportColumns>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let (table_name, payload) = domain_payload_table(row, "auxiliary", None)?;
|
||||
let table = auxiliary_table(&table_name)?;
|
||||
let target_columns = mysql_import_columns_cached(tx, column_cache, table.name).await?;
|
||||
import_mysql_row(
|
||||
tx,
|
||||
table.name,
|
||||
ExportDomain::Auxiliary,
|
||||
&ExportRow {
|
||||
id: row.id.clone(),
|
||||
payload,
|
||||
},
|
||||
&target_columns,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn mysql_billing_table_name(table_name: &str) -> Result<&'static str, DataLayerError> {
|
||||
match table_name {
|
||||
"billing_rules" => Ok("billing_rules"),
|
||||
"dimension_collectors" => Ok("dimension_collectors"),
|
||||
"usage_settlement_snapshots" => Ok("usage_settlement_snapshots"),
|
||||
other => Err(DataLayerError::InvalidInput(format!(
|
||||
"unsupported mysql billing export table '{other}'"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
async fn import_mysql_wallet_row(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
row: &ExportRow,
|
||||
column_cache: &mut BTreeMap<String, MysqlImportColumns>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let (table_name, payload) = domain_payload_table(row, "wallet", Some("wallets"))?;
|
||||
let table_name = mysql_wallet_table_name(&table_name)?;
|
||||
let target_columns = mysql_import_columns_cached(tx, column_cache, table_name).await?;
|
||||
import_mysql_row(
|
||||
tx,
|
||||
table_name,
|
||||
ExportDomain::Wallets,
|
||||
&ExportRow {
|
||||
id: row.id.clone(),
|
||||
payload,
|
||||
},
|
||||
&target_columns,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn mysql_wallet_tables() -> &'static [(&'static str, &'static str)] {
|
||||
&[
|
||||
("wallets", "id"),
|
||||
("wallet_transactions", "id"),
|
||||
("wallet_daily_usage_ledgers", "id"),
|
||||
("payment_orders", "id"),
|
||||
("payment_callbacks", "id"),
|
||||
("refund_requests", "id"),
|
||||
("redeem_code_batches", "id"),
|
||||
("redeem_codes", "id"),
|
||||
]
|
||||
}
|
||||
|
||||
fn mysql_wallet_table_name(table_name: &str) -> Result<&'static str, DataLayerError> {
|
||||
mysql_wallet_tables()
|
||||
.iter()
|
||||
.find(|(candidate, _)| *candidate == table_name)
|
||||
.map(|(table, _)| *table)
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::InvalidInput(format!(
|
||||
"unsupported mysql wallet export table '{table_name}'"
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
async fn mysql_import_columns_cached(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
cache: &mut BTreeMap<String, MysqlImportColumns>,
|
||||
table_name: &str,
|
||||
) -> Result<MysqlImportColumns, DataLayerError> {
|
||||
if let Some(columns) = cache.get(table_name) {
|
||||
return Ok(columns.clone());
|
||||
}
|
||||
|
||||
let columns = load_mysql_import_columns(tx, table_name).await?;
|
||||
cache.insert(table_name.to_string(), columns.clone());
|
||||
Ok(columns)
|
||||
}
|
||||
|
||||
async fn load_mysql_import_columns(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
table_name: &str,
|
||||
) -> Result<MysqlImportColumns, DataLayerError> {
|
||||
let relation_name = table_name.trim_matches('`');
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT
|
||||
CAST(COLUMN_NAME AS CHAR) AS column_name,
|
||||
CAST(DATA_TYPE AS CHAR) AS data_type,
|
||||
CAST(COLUMN_KEY AS CHAR) AS column_key,
|
||||
ORDINAL_POSITION AS ordinal_position
|
||||
FROM information_schema.columns
|
||||
WHERE table_schema = DATABASE()
|
||||
AND table_name = ?
|
||||
"#,
|
||||
)
|
||||
.bind(relation_name)
|
||||
.fetch_all(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
let mut columns = MysqlImportColumns::default();
|
||||
let mut primary_key = BTreeMap::new();
|
||||
for row in rows {
|
||||
let name = row.try_get::<String, _>("column_name").map_sql_err()?;
|
||||
let data_type = row
|
||||
.try_get::<String, _>("data_type")
|
||||
.map_sql_err()?
|
||||
.to_ascii_lowercase();
|
||||
columns.names.insert(name.clone());
|
||||
columns.data_types.insert(name.clone(), data_type);
|
||||
if row
|
||||
.try_get::<String, _>("column_key")
|
||||
.map_sql_err()?
|
||||
.eq_ignore_ascii_case("PRI")
|
||||
{
|
||||
primary_key.insert(
|
||||
row.try_get::<u32, _>("ordinal_position").map_sql_err()?,
|
||||
name,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
if columns.names.is_empty() {
|
||||
return Err(DataLayerError::UnexpectedValue(format!(
|
||||
"mysql import target table '{table_name}' has no visible columns"
|
||||
)));
|
||||
}
|
||||
if primary_key.is_empty() {
|
||||
return Err(DataLayerError::UnexpectedValue(format!(
|
||||
"mysql import target table '{table_name}' has no primary key"
|
||||
)));
|
||||
}
|
||||
columns.primary_key = primary_key.into_values().collect();
|
||||
|
||||
Ok(columns)
|
||||
}
|
||||
|
||||
fn bind_mysql_import_value<'q>(
|
||||
query: sqlx::query::Query<'q, sqlx::MySql, sqlx::mysql::MySqlArguments>,
|
||||
json_value: &'q Value,
|
||||
table_name: &str,
|
||||
column_name: &str,
|
||||
data_type: &str,
|
||||
) -> Result<sqlx::query::Query<'q, sqlx::MySql, sqlx::mysql::MySqlArguments>, DataLayerError> {
|
||||
if matches!(
|
||||
data_type,
|
||||
"binary" | "varbinary" | "blob" | "tinyblob" | "mediumblob" | "longblob"
|
||||
) {
|
||||
return match normalize_imported_binary("mysql", column_name, json_value)? {
|
||||
Some(bytes) => Ok(query.bind(bytes)),
|
||||
None => Ok(query.bind(Option::<Vec<u8>>::None)),
|
||||
};
|
||||
}
|
||||
if matches!(data_type, "decimal" | "numeric") {
|
||||
return match normalize_mysql_decimal_value(column_name, json_value)? {
|
||||
Some(value) => Ok(query.bind(value)),
|
||||
None => Ok(query.bind(Option::<String>::None)),
|
||||
};
|
||||
}
|
||||
let has_integer_type = matches!(
|
||||
data_type,
|
||||
"tinyint" | "smallint" | "mediumint" | "int" | "integer" | "bigint"
|
||||
);
|
||||
if !has_integer_type || !import_column_stores_timestamp(column_name) {
|
||||
return bind_mysql_json_value(query, json_value);
|
||||
}
|
||||
|
||||
match normalize_imported_integer_timestamp("mysql", table_name, column_name, json_value)? {
|
||||
Some(timestamp) => Ok(query.bind(timestamp)),
|
||||
None => Ok(query.bind(Option::<i64>::None)),
|
||||
}
|
||||
}
|
||||
|
||||
fn normalize_mysql_decimal_value(
|
||||
column_name: &str,
|
||||
value: &Value,
|
||||
) -> Result<Option<String>, DataLayerError> {
|
||||
match value {
|
||||
Value::Null => Ok(None),
|
||||
Value::Number(value) => Ok(Some(value.to_string())),
|
||||
Value::String(value) => Ok(Some(value.clone())),
|
||||
Value::Bool(_) | Value::Array(_) | Value::Object(_) => {
|
||||
Err(DataLayerError::InvalidInput(format!(
|
||||
"mysql decimal import column '{column_name}' must contain a number or numeric string"
|
||||
)))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn mysql_quote_identifier(identifier: &str) -> Result<String, DataLayerError> {
|
||||
if identifier.trim().is_empty() {
|
||||
return Err(DataLayerError::InvalidInput(
|
||||
"mysql import column name cannot be empty".to_string(),
|
||||
));
|
||||
}
|
||||
if !identifier
|
||||
.chars()
|
||||
.all(|ch| ch.is_ascii_alphanumeric() || ch == '_')
|
||||
{
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"mysql import column name '{identifier}' contains unsupported characters"
|
||||
)));
|
||||
}
|
||||
Ok(format!("`{identifier}`"))
|
||||
}
|
||||
|
||||
fn bind_mysql_json_value<'q>(
|
||||
query: sqlx::query::Query<'q, sqlx::MySql, sqlx::mysql::MySqlArguments>,
|
||||
value: &'q Value,
|
||||
) -> Result<sqlx::query::Query<'q, sqlx::MySql, sqlx::mysql::MySqlArguments>, DataLayerError> {
|
||||
Ok(match value {
|
||||
Value::Null => query.bind(Option::<String>::None),
|
||||
Value::Bool(value) => query.bind(i64::from(*value)),
|
||||
Value::Number(value) => {
|
||||
if let Some(value) = value.as_i64() {
|
||||
query.bind(value)
|
||||
} else if let Some(value) = value.as_u64() {
|
||||
let value = i64::try_from(value).map_err(|_| {
|
||||
DataLayerError::InvalidInput(format!(
|
||||
"mysql import integer value {value} exceeds i64"
|
||||
))
|
||||
})?;
|
||||
query.bind(value)
|
||||
} else if let Some(value) = value.as_f64() {
|
||||
query.bind(value)
|
||||
} else {
|
||||
return Err(DataLayerError::InvalidInput(
|
||||
"mysql import number is not representable".to_string(),
|
||||
));
|
||||
}
|
||||
}
|
||||
Value::String(value) => query.bind(value),
|
||||
Value::Array(_) | Value::Object(_) => {
|
||||
let value = serde_json::to_string(value)
|
||||
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?;
|
||||
query.bind(value)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
fn mysql_row_payload(row: &sqlx::mysql::MySqlRow) -> Result<Value, DataLayerError> {
|
||||
let mut object = serde_json::Map::new();
|
||||
for (index, column) in row.columns().iter().enumerate() {
|
||||
object.insert(column.name().to_string(), mysql_value_to_json(row, index)?);
|
||||
}
|
||||
Ok(Value::Object(object))
|
||||
}
|
||||
|
||||
fn mysql_value_to_json(row: &sqlx::mysql::MySqlRow, index: usize) -> Result<Value, DataLayerError> {
|
||||
let raw = row.try_get_raw(index).map_sql_err()?;
|
||||
if raw.is_null() {
|
||||
return Ok(Value::Null);
|
||||
}
|
||||
|
||||
match raw.type_info().name().to_ascii_uppercase().as_str() {
|
||||
"BOOL" | "BOOLEAN" => Ok(Value::Bool(row.try_get::<bool, _>(index).map_sql_err()?)),
|
||||
"TINYINT" | "TINY" | "SMALLINT" | "SHORT" | "MEDIUMINT" | "INT24" | "INT" | "INTEGER"
|
||||
| "LONG" | "BIGINT" | "LONGLONG" | "YEAR" => {
|
||||
Ok(Value::from(row.try_get::<i64, _>(index).map_sql_err()?))
|
||||
}
|
||||
"FLOAT" | "DOUBLE" => {
|
||||
let value = row.try_get::<f64, _>(index).map_sql_err()?;
|
||||
serde_json::Number::from_f64(value)
|
||||
.map(Value::Number)
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"mysql export column {} contains non-finite float",
|
||||
index
|
||||
))
|
||||
})
|
||||
}
|
||||
"DECIMAL" | "NEWDECIMAL" => Ok(Value::String(
|
||||
row.try_get::<sqlx::types::BigDecimal, _>(index)
|
||||
.map_sql_err()?
|
||||
.to_string(),
|
||||
)),
|
||||
"VARCHAR" | "VAR_STRING" | "STRING" | "TEXT" | "TINYTEXT" | "MEDIUMTEXT" | "LONGTEXT"
|
||||
| "JSON" | "ENUM" | "SET" | "DATE" | "DATETIME" | "TIMESTAMP" | "TIME" => Ok(
|
||||
Value::String(row.try_get::<String, _>(index).map_sql_err()?),
|
||||
),
|
||||
"BLOB" | "TINYBLOB" | "MEDIUMBLOB" | "LONGBLOB" | "BIT" | "GEOMETRY" => {
|
||||
let bytes = row.try_get::<Vec<u8>, _>(index).map_sql_err()?;
|
||||
Ok(Value::Array(bytes.into_iter().map(Value::from).collect()))
|
||||
}
|
||||
other => Err(DataLayerError::UnexpectedValue(format!(
|
||||
"unsupported mysql export column type '{other}' at index {index}"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::normalize_mysql_decimal_value;
|
||||
use serde_json::json;
|
||||
|
||||
#[test]
|
||||
fn decimal_import_binds_numbers_and_strings_as_decimal_text() {
|
||||
let value = json!(12345.12345678);
|
||||
assert_eq!(
|
||||
normalize_mysql_decimal_value("billing_total_cost_usd", &value)
|
||||
.expect("decimal value should normalize")
|
||||
.as_deref(),
|
||||
Some("12345.12345678")
|
||||
);
|
||||
assert_eq!(
|
||||
normalize_mysql_decimal_value(
|
||||
"billing_total_cost_usd",
|
||||
&json!("123456789012.12345678")
|
||||
)
|
||||
.expect("decimal string should normalize")
|
||||
.as_deref(),
|
||||
Some("123456789012.12345678")
|
||||
);
|
||||
assert!(normalize_mysql_decimal_value("billing_total_cost_usd", &json!(true)).is_err());
|
||||
}
|
||||
}
|
||||
@@ -1,772 +0,0 @@
|
||||
use super::*;
|
||||
|
||||
#[derive(Debug, Clone, Default, PartialEq, Eq)]
|
||||
struct SqliteImportColumns {
|
||||
names: ImportColumnNames,
|
||||
declared_types: BTreeMap<String, String>,
|
||||
primary_key: Vec<String>,
|
||||
}
|
||||
|
||||
pub async fn export_sqlite_core_jsonl(
|
||||
pool: &crate::driver::sqlite::SqlitePool,
|
||||
created_at_unix_secs: u64,
|
||||
) -> Result<String, DataLayerError> {
|
||||
export_sqlite_jsonl(pool, sqlite_core_export_domains(), created_at_unix_secs).await
|
||||
}
|
||||
|
||||
pub async fn export_sqlite_jsonl(
|
||||
pool: &crate::driver::sqlite::SqlitePool,
|
||||
domains: Vec<ExportDomain>,
|
||||
created_at_unix_secs: u64,
|
||||
) -> Result<String, DataLayerError> {
|
||||
let mut tx = pool.begin().await.map_sql_err()?;
|
||||
let manifest = DataExportManifest::new(
|
||||
created_at_unix_secs,
|
||||
Some(DatabaseDriver::Sqlite),
|
||||
domains.clone(),
|
||||
);
|
||||
let mut records = vec![DataExportRecord::manifest(manifest)];
|
||||
|
||||
for domain in domains {
|
||||
if domain == ExportDomain::Auxiliary {
|
||||
export_sqlite_auxiliary_records(&mut tx, &mut records).await?;
|
||||
continue;
|
||||
}
|
||||
if domain == ExportDomain::Billing {
|
||||
export_sqlite_billing_records(&mut tx, &mut records).await?;
|
||||
continue;
|
||||
}
|
||||
if domain == ExportDomain::Wallets {
|
||||
export_sqlite_wallet_records(&mut tx, &mut records).await?;
|
||||
continue;
|
||||
}
|
||||
let (table_name, id_column) = sqlite_domain_table(domain)?;
|
||||
let order_by = export_order_by(domain, id_column);
|
||||
let sql = format!("SELECT * FROM {table_name} ORDER BY {order_by}");
|
||||
let rows = sqlx::query(&sql).fetch_all(&mut *tx).await.map_sql_err()?;
|
||||
for row in rows {
|
||||
let id = sqlite_export_row_id(domain, &row, id_column)?;
|
||||
records.push(DataExportRecord::row(domain, id, sqlite_row_payload(&row)?));
|
||||
}
|
||||
}
|
||||
|
||||
tx.commit().await.map_sql_err()?;
|
||||
encode_jsonl(&records)
|
||||
}
|
||||
|
||||
pub async fn import_sqlite_jsonl(
|
||||
pool: &crate::driver::sqlite::SqlitePool,
|
||||
input: &str,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
let plan = build_import_plan(input)?;
|
||||
import_sqlite_plan(pool, &plan).await
|
||||
}
|
||||
|
||||
pub async fn import_sqlite_plan(
|
||||
pool: &crate::driver::sqlite::SqlitePool,
|
||||
plan: &DataImportPlan,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
let identity_scope = IdentityImportScope::from_plan(plan)?;
|
||||
let mut tx = pool.begin().await.map_sql_err()?;
|
||||
let identity_state = capture_sqlite_identity_import_state(&mut tx, &identity_scope).await?;
|
||||
let mut imported = 0usize;
|
||||
let mut column_cache = BTreeMap::<String, SqliteImportColumns>::new();
|
||||
for domain in &plan.manifest.domains {
|
||||
if *domain == ExportDomain::Auxiliary {
|
||||
for row in plan.rows(*domain) {
|
||||
import_sqlite_auxiliary_row(&mut tx, row, &mut column_cache).await?;
|
||||
imported = imported.saturating_add(1);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
if *domain == ExportDomain::Billing {
|
||||
for row in plan.rows(*domain) {
|
||||
import_sqlite_billing_row(&mut tx, row, &mut column_cache).await?;
|
||||
imported = imported.saturating_add(1);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
if *domain == ExportDomain::Wallets {
|
||||
for row in plan.rows(*domain) {
|
||||
import_sqlite_wallet_row(&mut tx, row, &mut column_cache).await?;
|
||||
imported = imported.saturating_add(1);
|
||||
}
|
||||
continue;
|
||||
}
|
||||
let (table_name, _id_column) = sqlite_domain_table(*domain)?;
|
||||
let target_columns =
|
||||
sqlite_import_columns_cached(&mut tx, &mut column_cache, table_name).await?;
|
||||
for row in plan.rows(*domain) {
|
||||
import_sqlite_row(&mut tx, table_name, *domain, row, &target_columns).await?;
|
||||
imported = imported.saturating_add(1);
|
||||
}
|
||||
}
|
||||
enforce_sqlite_identity_import_invariants(&mut tx, &identity_scope, identity_state).await?;
|
||||
tx.commit().await.map_sql_err()?;
|
||||
Ok(imported)
|
||||
}
|
||||
|
||||
async fn capture_sqlite_identity_import_state(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
|
||||
scope: &IdentityImportScope,
|
||||
) -> Result<IdentityImportState, DataLayerError> {
|
||||
let mut affected_user_ids = if scope.finalizes_oauth_links {
|
||||
scope.user_ids.iter().cloned().collect::<BTreeSet<_>>()
|
||||
} else {
|
||||
BTreeSet::new()
|
||||
};
|
||||
for link_id in &scope.oauth_link_ids {
|
||||
if let Some(user_id) =
|
||||
sqlx::query_scalar::<_, String>("SELECT user_id FROM user_oauth_links WHERE id = ?")
|
||||
.bind(link_id)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
{
|
||||
affected_user_ids.insert(user_id);
|
||||
}
|
||||
}
|
||||
for provider_type in &scope.oauth_provider_types {
|
||||
let user_ids = sqlx::query_scalar::<_, String>(
|
||||
"SELECT user_id FROM user_oauth_links WHERE provider_type = ?",
|
||||
)
|
||||
.bind(provider_type)
|
||||
.fetch_all(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
affected_user_ids.extend(user_ids);
|
||||
}
|
||||
Ok(IdentityImportState { affected_user_ids })
|
||||
}
|
||||
|
||||
async fn enforce_sqlite_identity_import_invariants(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
|
||||
scope: &IdentityImportScope,
|
||||
mut state: IdentityImportState,
|
||||
) -> Result<(), DataLayerError> {
|
||||
for user_id in &scope.user_ids {
|
||||
let auth_source =
|
||||
sqlx::query_scalar::<_, String>("SELECT auth_source FROM users WHERE id = ? LIMIT 1")
|
||||
.bind(user_id)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::InvalidInput(format!(
|
||||
"imported users row '{user_id}' did not produce a user record"
|
||||
))
|
||||
})?;
|
||||
if !matches!(auth_source.as_str(), "local" | "ldap" | "oauth") {
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"imported user '{user_id}' has unsupported auth_source '{auth_source}'"
|
||||
)));
|
||||
}
|
||||
if auth_source == "oauth" {
|
||||
sqlx::query("UPDATE users SET email_verified = 0 WHERE id = ?")
|
||||
.bind(user_id)
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
}
|
||||
}
|
||||
|
||||
for link_id in &scope.oauth_link_ids {
|
||||
let user_id = sqlx::query_scalar::<_, String>(
|
||||
"SELECT user_id FROM user_oauth_links WHERE id = ? LIMIT 1",
|
||||
)
|
||||
.bind(link_id)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::InvalidInput(format!(
|
||||
"imported OAuth link row '{link_id}' did not produce a link record"
|
||||
))
|
||||
})?;
|
||||
state.affected_user_ids.insert(user_id);
|
||||
}
|
||||
|
||||
for link_id in &scope.oauth_link_ids {
|
||||
if let Some((provider_type, provider_user_id)) = sqlx::query_as::<_, (String, String)>(
|
||||
r#"
|
||||
SELECT imported.provider_type, imported.provider_user_id
|
||||
FROM user_oauth_links imported
|
||||
JOIN user_oauth_links duplicate
|
||||
ON duplicate.provider_type = imported.provider_type
|
||||
AND duplicate.provider_user_id = imported.provider_user_id
|
||||
AND duplicate.id <> imported.id
|
||||
WHERE imported.id = ?
|
||||
LIMIT 1
|
||||
"#,
|
||||
)
|
||||
.bind(link_id)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
{
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"OAuth import assigns provider identity '{provider_type}:{provider_user_id}' more than once"
|
||||
)));
|
||||
}
|
||||
}
|
||||
|
||||
for link_id in &scope.oauth_link_ids {
|
||||
if let Some((user_id, provider_type)) = sqlx::query_as::<_, (String, String)>(
|
||||
r#"
|
||||
SELECT imported.user_id, imported.provider_type
|
||||
FROM user_oauth_links imported
|
||||
JOIN user_oauth_links duplicate
|
||||
ON duplicate.user_id = imported.user_id
|
||||
AND duplicate.provider_type = imported.provider_type
|
||||
AND duplicate.id <> imported.id
|
||||
WHERE imported.id = ?
|
||||
LIMIT 1
|
||||
"#,
|
||||
)
|
||||
.bind(link_id)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
{
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"OAuth import links user '{user_id}' to provider '{provider_type}' more than once"
|
||||
)));
|
||||
}
|
||||
}
|
||||
|
||||
for link_id in &scope.oauth_link_ids {
|
||||
if let Some(invalid_id) = sqlx::query_scalar::<_, String>(
|
||||
r#"
|
||||
SELECT links.id
|
||||
FROM user_oauth_links links
|
||||
LEFT JOIN users ON users.id = links.user_id
|
||||
LEFT JOIN oauth_providers providers ON providers.provider_type = links.provider_type
|
||||
WHERE links.id = ?
|
||||
AND (
|
||||
users.id IS NULL
|
||||
OR providers.provider_type IS NULL
|
||||
OR links.provider_type <> LOWER(TRIM(links.provider_type))
|
||||
OR links.provider_type = ''
|
||||
OR links.provider_user_id <> TRIM(links.provider_user_id)
|
||||
OR links.provider_user_id = ''
|
||||
OR providers.provider_type <> LOWER(TRIM(providers.provider_type))
|
||||
)
|
||||
LIMIT 1
|
||||
"#,
|
||||
)
|
||||
.bind(link_id)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
{
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"OAuth import produced invalid or orphaned link '{invalid_id}'"
|
||||
)));
|
||||
}
|
||||
}
|
||||
|
||||
if !scope.validates_oauth_login_methods {
|
||||
return Ok(());
|
||||
}
|
||||
for user_id in state.affected_user_ids {
|
||||
if sqlx::query_scalar::<_, String>(
|
||||
r#"
|
||||
SELECT users.id
|
||||
FROM users
|
||||
WHERE users.id = ?
|
||||
AND users.auth_source = 'oauth'
|
||||
AND users.is_active = 1
|
||||
AND users.is_deleted = 0
|
||||
AND NOT EXISTS (
|
||||
SELECT 1
|
||||
FROM user_oauth_links links
|
||||
JOIN oauth_providers providers ON providers.provider_type = links.provider_type
|
||||
WHERE links.user_id = users.id
|
||||
AND providers.is_enabled = 1
|
||||
AND links.provider_type = LOWER(TRIM(links.provider_type))
|
||||
AND links.provider_user_id = TRIM(links.provider_user_id)
|
||||
AND links.provider_user_id <> ''
|
||||
)
|
||||
LIMIT 1
|
||||
"#,
|
||||
)
|
||||
.bind(&user_id)
|
||||
.fetch_optional(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?
|
||||
.is_some()
|
||||
{
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"OAuth import would leave active user '{user_id}' without an enabled identity binding"
|
||||
)));
|
||||
}
|
||||
}
|
||||
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn sqlite_domain_table(
|
||||
domain: ExportDomain,
|
||||
) -> Result<(&'static str, &'static str), DataLayerError> {
|
||||
match domain {
|
||||
ExportDomain::Users => Ok(("users", "id")),
|
||||
ExportDomain::ApiKeys => Ok(("api_keys", "id")),
|
||||
ExportDomain::Providers => Ok(("providers", "id")),
|
||||
ExportDomain::ProviderKeys => Ok(("provider_api_keys", "id")),
|
||||
ExportDomain::Endpoints => Ok(("provider_endpoints", "id")),
|
||||
ExportDomain::Models => Ok(("models", "id")),
|
||||
ExportDomain::GlobalModels => Ok(("global_models", "id")),
|
||||
ExportDomain::AuthModules => Ok(("auth_modules", "id")),
|
||||
ExportDomain::OAuthProviders => Ok(("oauth_providers", "provider_type")),
|
||||
ExportDomain::UserOAuthLinks => Ok(("user_oauth_links", "id")),
|
||||
ExportDomain::UserGroups => Ok(("user_groups", "id")),
|
||||
ExportDomain::UserGroupMembers => Ok(("user_group_members", "group_id")),
|
||||
ExportDomain::ProxyNodes => Ok(("proxy_nodes", "id")),
|
||||
ExportDomain::SystemConfigs => Ok(("system_configs", "id")),
|
||||
ExportDomain::Wallets => Err(DataLayerError::InvalidInput(
|
||||
"sqlite wallet export uses multiple tables and must be handled as a domain".to_string(),
|
||||
)),
|
||||
ExportDomain::Usage => Ok((r#""usage""#, "request_id")),
|
||||
ExportDomain::Billing => Err(DataLayerError::InvalidInput(
|
||||
"sqlite billing export uses multiple tables and must be handled as a domain"
|
||||
.to_string(),
|
||||
)),
|
||||
ExportDomain::Auxiliary => Err(DataLayerError::InvalidInput(
|
||||
"sqlite auxiliary export uses multiple tables and must be handled as a domain"
|
||||
.to_string(),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
async fn export_sqlite_auxiliary_records(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
|
||||
records: &mut Vec<DataExportRecord>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
for table in AUXILIARY_TABLES {
|
||||
let table_sql = sqlite_quote_identifier(table.name)?;
|
||||
let order_sql = table
|
||||
.primary_key
|
||||
.iter()
|
||||
.map(|column| sqlite_quote_identifier(column).map(|column| format!("{column} ASC")))
|
||||
.collect::<Result<Vec<_>, _>>()?
|
||||
.join(", ");
|
||||
let rows = sqlx::query(&format!("SELECT * FROM {table_sql} ORDER BY {order_sql}"))
|
||||
.fetch_all(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
for row in rows {
|
||||
let payload = sqlite_row_payload(&row)?;
|
||||
let id = auxiliary_row_id(*table, &payload)?;
|
||||
records.push(DataExportRecord::row(
|
||||
ExportDomain::Auxiliary,
|
||||
id,
|
||||
payload_with_table(payload, table.name)?,
|
||||
));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn sqlite_export_row_id(
|
||||
domain: ExportDomain,
|
||||
row: &sqlx::sqlite::SqliteRow,
|
||||
id_column: &str,
|
||||
) -> Result<String, DataLayerError> {
|
||||
if domain == ExportDomain::UserGroupMembers {
|
||||
let group_id = sqlite_required_export_text(row, "group_id", domain)?;
|
||||
let user_id = sqlite_required_export_text(row, "user_id", domain)?;
|
||||
return Ok(format!("{group_id}:{user_id}"));
|
||||
}
|
||||
sqlite_required_export_text(row, id_column, domain)
|
||||
}
|
||||
|
||||
fn sqlite_required_export_text(
|
||||
row: &sqlx::sqlite::SqliteRow,
|
||||
column: &str,
|
||||
domain: ExportDomain,
|
||||
) -> Result<String, DataLayerError> {
|
||||
row.try_get::<Option<String>, _>(column)
|
||||
.map_sql_err()?
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"{} export row has null id column '{}'",
|
||||
domain.as_str(),
|
||||
column
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
async fn export_sqlite_billing_records(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
|
||||
records: &mut Vec<DataExportRecord>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
for table_name in [
|
||||
"billing_rules",
|
||||
"dimension_collectors",
|
||||
"usage_settlement_snapshots",
|
||||
] {
|
||||
let id_column = if table_name == "usage_settlement_snapshots" {
|
||||
"request_id"
|
||||
} else {
|
||||
"id"
|
||||
};
|
||||
let sql = format!("SELECT * FROM {table_name} ORDER BY {id_column} ASC");
|
||||
let rows = sqlx::query(&sql).fetch_all(&mut **tx).await.map_sql_err()?;
|
||||
for row in rows {
|
||||
let id = row
|
||||
.try_get::<Option<String>, _>(id_column)
|
||||
.map_sql_err()?
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"billing export row in table '{table_name}' has null id"
|
||||
))
|
||||
})?;
|
||||
records.push(DataExportRecord::row(
|
||||
ExportDomain::Billing,
|
||||
format!("{table_name}:{id}"),
|
||||
payload_with_table(sqlite_row_payload(&row)?, table_name)?,
|
||||
));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn export_sqlite_wallet_records(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
|
||||
records: &mut Vec<DataExportRecord>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
for (table_name, id_column) in sqlite_wallet_tables() {
|
||||
let sql = format!("SELECT * FROM {table_name} ORDER BY {id_column} ASC");
|
||||
let rows = sqlx::query(&sql).fetch_all(&mut **tx).await.map_sql_err()?;
|
||||
for row in rows {
|
||||
let id = row
|
||||
.try_get::<Option<String>, _>(id_column)
|
||||
.map_sql_err()?
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"wallet export row in table '{table_name}' has null id"
|
||||
))
|
||||
})?;
|
||||
records.push(DataExportRecord::row(
|
||||
ExportDomain::Wallets,
|
||||
format!("{table_name}:{id}"),
|
||||
payload_with_table(sqlite_row_payload(&row)?, table_name)?,
|
||||
));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn import_sqlite_row(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
|
||||
table_name: &str,
|
||||
domain: ExportDomain,
|
||||
row: &ExportRow,
|
||||
target_columns: &SqliteImportColumns,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let mut object =
|
||||
filter_import_payload("sqlite", table_name, domain, row, &target_columns.names)?;
|
||||
deactivate_imported_credentials(table_name, &mut object, |column_name| {
|
||||
target_columns.names.contains(column_name)
|
||||
});
|
||||
|
||||
let columns = object.keys().map(String::as_str).collect::<Vec<_>>();
|
||||
let column_sql = columns
|
||||
.iter()
|
||||
.map(|column| sqlite_quote_identifier(column))
|
||||
.collect::<Result<Vec<_>, _>>()?
|
||||
.join(", ");
|
||||
let placeholder_sql = vec!["?"; columns.len()].join(", ");
|
||||
let conflict_columns = target_columns
|
||||
.primary_key
|
||||
.iter()
|
||||
.map(|column| sqlite_quote_identifier(column))
|
||||
.collect::<Result<Vec<_>, _>>()?
|
||||
.join(", ");
|
||||
let update_sql = columns
|
||||
.iter()
|
||||
.filter(|column| !target_columns.primary_key.iter().any(|key| key == *column))
|
||||
.map(|column| {
|
||||
let quoted = sqlite_quote_identifier(column)?;
|
||||
Ok(format!("{quoted} = excluded.{quoted}"))
|
||||
})
|
||||
.collect::<Result<Vec<_>, DataLayerError>>()?
|
||||
.join(", ");
|
||||
let conflict_sql = if update_sql.is_empty() {
|
||||
format!("ON CONFLICT ({conflict_columns}) DO NOTHING")
|
||||
} else {
|
||||
format!("ON CONFLICT ({conflict_columns}) DO UPDATE SET {update_sql}")
|
||||
};
|
||||
let sql = format!(
|
||||
"INSERT INTO {table_name} ({column_sql}) VALUES ({placeholder_sql}) {conflict_sql}"
|
||||
);
|
||||
let mut query = sqlx::query(&sql);
|
||||
for column in columns {
|
||||
let value = object
|
||||
.get(column)
|
||||
.expect("column name came from payload object keys");
|
||||
let declared_type = target_columns
|
||||
.declared_types
|
||||
.get(column)
|
||||
.map(String::as_str)
|
||||
.unwrap_or_default();
|
||||
query = bind_sqlite_import_value(query, value, table_name, column, declared_type)?;
|
||||
}
|
||||
query.execute(&mut **tx).await.map_sql_err()?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn import_sqlite_billing_row(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
|
||||
row: &ExportRow,
|
||||
column_cache: &mut BTreeMap<String, SqliteImportColumns>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let (table_name, payload) = billing_payload_table(row)?;
|
||||
let table_name = sqlite_billing_table_name(&table_name)?;
|
||||
let target_columns = sqlite_import_columns_cached(tx, column_cache, table_name).await?;
|
||||
import_sqlite_row(
|
||||
tx,
|
||||
table_name,
|
||||
ExportDomain::Billing,
|
||||
&ExportRow {
|
||||
id: row.id.clone(),
|
||||
payload,
|
||||
},
|
||||
&target_columns,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
async fn import_sqlite_auxiliary_row(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
|
||||
row: &ExportRow,
|
||||
column_cache: &mut BTreeMap<String, SqliteImportColumns>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let (table_name, payload) = domain_payload_table(row, "auxiliary", None)?;
|
||||
let table = auxiliary_table(&table_name)?;
|
||||
let target_columns = sqlite_import_columns_cached(tx, column_cache, table.name).await?;
|
||||
import_sqlite_row(
|
||||
tx,
|
||||
table.name,
|
||||
ExportDomain::Auxiliary,
|
||||
&ExportRow {
|
||||
id: row.id.clone(),
|
||||
payload,
|
||||
},
|
||||
&target_columns,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn sqlite_billing_table_name(table_name: &str) -> Result<&'static str, DataLayerError> {
|
||||
match table_name {
|
||||
"billing_rules" => Ok("billing_rules"),
|
||||
"dimension_collectors" => Ok("dimension_collectors"),
|
||||
"usage_settlement_snapshots" => Ok("usage_settlement_snapshots"),
|
||||
other => Err(DataLayerError::InvalidInput(format!(
|
||||
"unsupported sqlite billing export table '{other}'"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
async fn import_sqlite_wallet_row(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
|
||||
row: &ExportRow,
|
||||
column_cache: &mut BTreeMap<String, SqliteImportColumns>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let (table_name, payload) = domain_payload_table(row, "wallet", Some("wallets"))?;
|
||||
let table_name = sqlite_wallet_table_name(&table_name)?;
|
||||
let target_columns = sqlite_import_columns_cached(tx, column_cache, table_name).await?;
|
||||
import_sqlite_row(
|
||||
tx,
|
||||
table_name,
|
||||
ExportDomain::Wallets,
|
||||
&ExportRow {
|
||||
id: row.id.clone(),
|
||||
payload,
|
||||
},
|
||||
&target_columns,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn sqlite_wallet_tables() -> &'static [(&'static str, &'static str)] {
|
||||
&[
|
||||
("wallets", "id"),
|
||||
("wallet_transactions", "id"),
|
||||
("wallet_daily_usage_ledgers", "id"),
|
||||
("payment_orders", "id"),
|
||||
("payment_callbacks", "id"),
|
||||
("refund_requests", "id"),
|
||||
("redeem_code_batches", "id"),
|
||||
("redeem_codes", "id"),
|
||||
]
|
||||
}
|
||||
|
||||
fn sqlite_wallet_table_name(table_name: &str) -> Result<&'static str, DataLayerError> {
|
||||
sqlite_wallet_tables()
|
||||
.iter()
|
||||
.find(|(candidate, _)| *candidate == table_name)
|
||||
.map(|(table, _)| *table)
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::InvalidInput(format!(
|
||||
"unsupported sqlite wallet export table '{table_name}'"
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
async fn sqlite_import_columns_cached(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
|
||||
cache: &mut BTreeMap<String, SqliteImportColumns>,
|
||||
table_name: &str,
|
||||
) -> Result<SqliteImportColumns, DataLayerError> {
|
||||
if let Some(columns) = cache.get(table_name) {
|
||||
return Ok(columns.clone());
|
||||
}
|
||||
|
||||
let columns = load_sqlite_import_columns(tx, table_name).await?;
|
||||
cache.insert(table_name.to_string(), columns.clone());
|
||||
Ok(columns)
|
||||
}
|
||||
|
||||
async fn load_sqlite_import_columns(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
|
||||
table_name: &str,
|
||||
) -> Result<SqliteImportColumns, DataLayerError> {
|
||||
let sql = format!("PRAGMA table_info({table_name})");
|
||||
let rows = sqlx::query(&sql).fetch_all(&mut **tx).await.map_sql_err()?;
|
||||
let mut columns = SqliteImportColumns::default();
|
||||
let mut primary_key = BTreeMap::new();
|
||||
for row in rows {
|
||||
let name = row.try_get::<String, _>("name").map_sql_err()?;
|
||||
let declared_type = row
|
||||
.try_get::<Option<String>, _>("type")
|
||||
.map_sql_err()?
|
||||
.unwrap_or_default();
|
||||
columns.names.insert(name.clone());
|
||||
columns.declared_types.insert(name.clone(), declared_type);
|
||||
let primary_key_position = row.try_get::<i64, _>("pk").map_sql_err()?;
|
||||
if primary_key_position > 0 {
|
||||
primary_key.insert(primary_key_position, name);
|
||||
}
|
||||
}
|
||||
|
||||
if columns.names.is_empty() {
|
||||
return Err(DataLayerError::UnexpectedValue(format!(
|
||||
"sqlite import target table '{table_name}' has no visible columns"
|
||||
)));
|
||||
}
|
||||
if primary_key.is_empty() {
|
||||
return Err(DataLayerError::UnexpectedValue(format!(
|
||||
"sqlite import target table '{table_name}' has no primary key"
|
||||
)));
|
||||
}
|
||||
columns.primary_key = primary_key.into_values().collect();
|
||||
|
||||
Ok(columns)
|
||||
}
|
||||
|
||||
fn bind_sqlite_import_value<'q>(
|
||||
query: sqlx::query::Query<'q, sqlx::Sqlite, sqlx::sqlite::SqliteArguments<'q>>,
|
||||
json_value: &'q Value,
|
||||
table_name: &str,
|
||||
column_name: &str,
|
||||
declared_type: &str,
|
||||
) -> Result<sqlx::query::Query<'q, sqlx::Sqlite, sqlx::sqlite::SqliteArguments<'q>>, DataLayerError>
|
||||
{
|
||||
if declared_type.to_ascii_uppercase().contains("BLOB") {
|
||||
return match normalize_imported_binary("sqlite", column_name, json_value)? {
|
||||
Some(bytes) => Ok(query.bind(bytes)),
|
||||
None => Ok(query.bind(Option::<Vec<u8>>::None)),
|
||||
};
|
||||
}
|
||||
let has_integer_affinity = declared_type.to_ascii_uppercase().contains("INT");
|
||||
if !has_integer_affinity || !import_column_stores_timestamp(column_name) {
|
||||
return bind_sqlite_json_value(query, json_value);
|
||||
}
|
||||
|
||||
match normalize_imported_integer_timestamp("sqlite", table_name, column_name, json_value)? {
|
||||
Some(timestamp) => Ok(query.bind(timestamp)),
|
||||
None => Ok(query.bind(Option::<i64>::None)),
|
||||
}
|
||||
}
|
||||
|
||||
fn sqlite_row_payload(row: &sqlx::sqlite::SqliteRow) -> Result<Value, DataLayerError> {
|
||||
let mut object = serde_json::Map::new();
|
||||
for (index, column) in row.columns().iter().enumerate() {
|
||||
object.insert(
|
||||
column.name().to_string(),
|
||||
sqlite_value_to_json(row, index, column.name())?,
|
||||
);
|
||||
}
|
||||
Ok(Value::Object(object))
|
||||
}
|
||||
|
||||
fn sqlite_value_to_json(
|
||||
row: &sqlx::sqlite::SqliteRow,
|
||||
index: usize,
|
||||
column_name: &str,
|
||||
) -> Result<Value, DataLayerError> {
|
||||
let raw = row.try_get_raw(index).map_sql_err()?;
|
||||
if raw.is_null() {
|
||||
return Ok(Value::Null);
|
||||
}
|
||||
|
||||
match raw.type_info().name().to_ascii_uppercase().as_str() {
|
||||
"INTEGER" => {
|
||||
let value = row.try_get::<i64, _>(index).map_sql_err()?;
|
||||
if sqlite_integer_column_is_boolean(column_name) {
|
||||
match value {
|
||||
0 => return Ok(Value::Bool(false)),
|
||||
1 => return Ok(Value::Bool(true)),
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
Ok(Value::from(value))
|
||||
}
|
||||
"REAL" | "FLOAT" | "DOUBLE" => {
|
||||
let value = row.try_get::<f64, _>(index).map_sql_err()?;
|
||||
serde_json::Number::from_f64(value)
|
||||
.map(Value::Number)
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"sqlite export column {} contains non-finite float",
|
||||
index
|
||||
))
|
||||
})
|
||||
}
|
||||
"TEXT" => Ok(Value::String(
|
||||
row.try_get::<String, _>(index).map_sql_err()?,
|
||||
)),
|
||||
"BLOB" => {
|
||||
let bytes = row.try_get::<Vec<u8>, _>(index).map_sql_err()?;
|
||||
Ok(Value::Array(bytes.into_iter().map(Value::from).collect()))
|
||||
}
|
||||
other => Err(DataLayerError::UnexpectedValue(format!(
|
||||
"unsupported sqlite export column type '{other}' at index {index}"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
fn sqlite_integer_column_is_boolean(column_name: &str) -> bool {
|
||||
column_name.starts_with("is_")
|
||||
|| column_name.starts_with("has_")
|
||||
|| column_name.starts_with("supports_")
|
||||
|| column_name.starts_with("enable_")
|
||||
|| column_name.starts_with("use_")
|
||||
|| matches!(
|
||||
column_name,
|
||||
"announcement_notifications"
|
||||
| "auto_delete_on_expiry"
|
||||
| "auto_fetch_models"
|
||||
| "email_notifications"
|
||||
| "email_verified"
|
||||
| "format_converted"
|
||||
| "keep_priority_on_conversion"
|
||||
| "signature_valid"
|
||||
| "tunnel_connected"
|
||||
| "tunnel_mode"
|
||||
| "usage_alerts"
|
||||
| "webhook_sent"
|
||||
)
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -3,58 +3,13 @@
|
||||
//! Each driver owns its migrator and startup preparation. The facade keeps
|
||||
//! the established public entry points used by gateway bootstrap code.
|
||||
|
||||
#[cfg(feature = "mysql")]
|
||||
mod mysql;
|
||||
#[cfg(feature = "postgres")]
|
||||
mod postgres;
|
||||
#[cfg(feature = "sqlite")]
|
||||
mod sqlite;
|
||||
mod types;
|
||||
|
||||
#[cfg(all(test, feature = "postgres", feature = "mysql", feature = "sqlite"))]
|
||||
#[cfg(all(test, feature = "postgres"))]
|
||||
mod tests;
|
||||
|
||||
#[cfg(feature = "postgres")]
|
||||
pub use postgres::{pending_migrations, prepare_database_for_startup, run_migrations};
|
||||
pub use types::PendingMigrationInfo;
|
||||
|
||||
#[cfg(any(feature = "mysql", feature = "sqlite"))]
|
||||
use sqlx::migrate::MigrateError;
|
||||
|
||||
#[cfg(feature = "mysql")]
|
||||
pub async fn run_mysql_migrations(pool: &sqlx::MySqlPool) -> Result<(), MigrateError> {
|
||||
mysql::run_migrations(pool).await
|
||||
}
|
||||
|
||||
#[cfg(feature = "mysql")]
|
||||
pub async fn pending_mysql_migrations(
|
||||
pool: &sqlx::MySqlPool,
|
||||
) -> Result<Vec<PendingMigrationInfo>, MigrateError> {
|
||||
mysql::pending_migrations(pool).await
|
||||
}
|
||||
|
||||
#[cfg(feature = "mysql")]
|
||||
pub async fn prepare_mysql_database_for_startup(
|
||||
pool: &sqlx::MySqlPool,
|
||||
) -> Result<Vec<PendingMigrationInfo>, MigrateError> {
|
||||
mysql::prepare_database_for_startup(pool).await
|
||||
}
|
||||
|
||||
#[cfg(feature = "sqlite")]
|
||||
pub async fn run_sqlite_migrations(pool: &sqlx::SqlitePool) -> Result<(), MigrateError> {
|
||||
sqlite::run_migrations(pool).await
|
||||
}
|
||||
|
||||
#[cfg(feature = "sqlite")]
|
||||
pub async fn pending_sqlite_migrations(
|
||||
pool: &sqlx::SqlitePool,
|
||||
) -> Result<Vec<PendingMigrationInfo>, MigrateError> {
|
||||
sqlite::pending_migrations(pool).await
|
||||
}
|
||||
|
||||
#[cfg(feature = "sqlite")]
|
||||
pub async fn prepare_sqlite_database_for_startup(
|
||||
pool: &sqlx::SqlitePool,
|
||||
) -> Result<Vec<PendingMigrationInfo>, MigrateError> {
|
||||
sqlite::prepare_database_for_startup(pool).await
|
||||
}
|
||||
|
||||
@@ -1,5 +0,0 @@
|
||||
#[cfg(all(test, feature = "postgres", feature = "mysql", feature = "sqlite"))]
|
||||
pub(super) use aether_data_mysql::MIGRATOR;
|
||||
pub(super) use aether_data_mysql::{
|
||||
pending_migrations, prepare_database_for_startup, run_migrations,
|
||||
};
|
||||
@@ -6,7 +6,7 @@ use sqlx::{
|
||||
use super::types::PendingMigrationInfo;
|
||||
|
||||
pub use aether_data_postgres::pending_migrations;
|
||||
#[cfg(all(test, feature = "postgres", feature = "mysql", feature = "sqlite"))]
|
||||
#[cfg(all(test, feature = "postgres"))]
|
||||
pub(super) use aether_data_postgres::{
|
||||
all_up_migrations, pending_migrations_from_applied, POSTGRES_MIGRATOR,
|
||||
};
|
||||
|
||||
@@ -1,5 +0,0 @@
|
||||
#[cfg(all(test, feature = "postgres", feature = "mysql", feature = "sqlite"))]
|
||||
pub(super) use aether_data_sqlite::MIGRATOR;
|
||||
pub(super) use aether_data_sqlite::{
|
||||
pending_migrations, prepare_database_for_startup, run_migrations,
|
||||
};
|
||||
File diff suppressed because it is too large
Load Diff
Reference in New Issue
Block a user