mirror of
https://github.com/fawney19/Aether.git
synced 2026-10-08 02:17:46 +08:00
feat(data): complete portable SQL backend parity
Align MySQL and SQLite schemas, migrations, usage, stats, export, and backfill behavior with the shared data contracts. Extend gateway startup and maintenance support across all SQL drivers.
This commit is contained in:
@@ -1,16 +1,253 @@
|
||||
use sqlx::migrate::MigrateError;
|
||||
use tracing::info;
|
||||
use std::collections::{HashMap, HashSet};
|
||||
|
||||
use sqlx::{
|
||||
migrate::{Migrate, MigrateError, Migrator},
|
||||
query, Connection, MySqlConnection, Row,
|
||||
};
|
||||
use tracing::{error, info, warn};
|
||||
|
||||
use super::types::PendingBackfillInfo;
|
||||
use crate::driver::mysql::MysqlPool;
|
||||
|
||||
pub async fn run_backfills(_pool: &MysqlPool) -> Result<(), MigrateError> {
|
||||
info!("mysql database backfills are up to date");
|
||||
static BACKFILL_MIGRATOR: Migrator = sqlx::migrate!("./backfills/mysql");
|
||||
|
||||
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(())
|
||||
}
|
||||
|
||||
pub async fn pending_backfills(
|
||||
_pool: &MysqlPool,
|
||||
async fn pending_backfills_locked(
|
||||
conn: &mut MySqlConnection,
|
||||
) -> Result<Vec<PendingBackfillInfo>, MigrateError> {
|
||||
Ok(Vec::new())
|
||||
ensure_schema_backfills_table(conn).await?;
|
||||
let applied_backfills = list_applied_backfills(conn).await?;
|
||||
validate_applied_backfills(&applied_backfills)?;
|
||||
Ok(pending_backfills_from_applied(&applied_backfills))
|
||||
}
|
||||
|
||||
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,16 +1,254 @@
|
||||
use sqlx::migrate::MigrateError;
|
||||
use tracing::info;
|
||||
use std::collections::{HashMap, HashSet};
|
||||
|
||||
use sqlx::{
|
||||
migrate::{Migrate, MigrateError, Migrator},
|
||||
query, Connection, Row, SqliteConnection,
|
||||
};
|
||||
use tracing::{error, info, warn};
|
||||
|
||||
use super::types::PendingBackfillInfo;
|
||||
use crate::driver::sqlite::SqlitePool;
|
||||
|
||||
pub async fn run_backfills(_pool: &SqlitePool) -> Result<(), MigrateError> {
|
||||
info!("sqlite database backfills are up to date");
|
||||
Ok(())
|
||||
static BACKFILL_MIGRATOR: Migrator = sqlx::migrate!("./backfills/sqlite");
|
||||
|
||||
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,
|
||||
pool: &SqlitePool,
|
||||
) -> Result<Vec<PendingBackfillInfo>, MigrateError> {
|
||||
Ok(Vec::new())
|
||||
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> {
|
||||
ensure_schema_backfills_table(conn).await?;
|
||||
let applied_backfills = list_applied_backfills(conn).await?;
|
||||
validate_applied_backfills(&applied_backfills)?;
|
||||
Ok(pending_backfills_from_applied(&applied_backfills))
|
||||
}
|
||||
|
||||
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,15 +4,14 @@ use std::{
|
||||
time::{Duration, Instant},
|
||||
};
|
||||
|
||||
use sqlx::{query, query_scalar, Connection, PgConnection, PgPool};
|
||||
use sqlx::{query, query_as, 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;
|
||||
use crate::{DatabaseDriver, SqlDatabaseConfig, SqlPoolConfig};
|
||||
use crate::lifecycle::migrate::{prepare_database_for_startup, run_sqlite_migrations};
|
||||
|
||||
const LEGACY_SYNC_ENABLED_ACTIVE_FLAGS_VERSION: i64 = 20260517012000;
|
||||
const LEGACY_SYNC_ENABLED_ACTIVE_FLAGS_SQL: &str =
|
||||
@@ -88,44 +87,558 @@ fn corrected_legacy_backfill_is_not_requeued_after_application() {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn mysql_backfills_are_empty_until_driver_specific_backfills_exist() {
|
||||
let pool = sqlx::mysql::MySqlPoolOptions::new().connect_lazy_with(
|
||||
"mysql://user:pass@localhost:3306/aether"
|
||||
.parse()
|
||||
.expect("mysql options should parse"),
|
||||
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_mysql_backfills(&pool)
|
||||
.await
|
||||
.expect("mysql pending backfills should load"),
|
||||
Vec::new()
|
||||
pending_versions,
|
||||
vec![
|
||||
20260422120000,
|
||||
20260505120000,
|
||||
20260517012000,
|
||||
20260716010000
|
||||
]
|
||||
);
|
||||
|
||||
run_mysql_backfills(&pool)
|
||||
.await
|
||||
.expect("mysql backfills should no-op");
|
||||
.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 sqlite_backfills_are_empty_until_driver_specific_backfills_exist() {
|
||||
let config = SqlDatabaseConfig::new(
|
||||
DatabaseDriver::Sqlite,
|
||||
"sqlite::memory:",
|
||||
SqlPoolConfig::default(),
|
||||
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
|
||||
)
|
||||
"#,
|
||||
)
|
||||
.expect("sqlite config should build");
|
||||
let pool = crate::driver::sqlite::SqlitePoolFactory::new(config)
|
||||
.expect("sqlite factory should build")
|
||||
.connect_lazy()
|
||||
.expect("sqlite pool should build");
|
||||
.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_sqlite_backfills(&pool)
|
||||
.await
|
||||
.expect("sqlite pending backfills should load"),
|
||||
Vec::new()
|
||||
pending_versions,
|
||||
vec![
|
||||
20260422120000,
|
||||
20260505120000,
|
||||
20260517012000,
|
||||
20260716010000
|
||||
]
|
||||
);
|
||||
|
||||
run_sqlite_backfills(&pool)
|
||||
.await
|
||||
.expect("sqlite backfills should no-op");
|
||||
.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)]
|
||||
|
||||
@@ -3,6 +3,8 @@ use std::collections::{BTreeMap, BTreeSet};
|
||||
#[cfg(all(feature = "postgres", feature = "sqlite"))]
|
||||
use futures_util::TryStreamExt;
|
||||
use serde_json::Value;
|
||||
#[cfg(all(feature = "postgres", feature = "sqlite"))]
|
||||
use sqlx::Acquire;
|
||||
use sqlx::Row;
|
||||
#[cfg(any(feature = "mysql", feature = "sqlite"))]
|
||||
use sqlx::{Column, TypeInfo, ValueRef};
|
||||
@@ -41,7 +43,8 @@ use postgres::{
|
||||
#[cfg(all(test, feature = "postgres", feature = "mysql", feature = "sqlite"))]
|
||||
use postgres::normalize_postgres_import_payload;
|
||||
|
||||
pub const EXPORT_FORMAT_VERSION: u32 = 1;
|
||||
pub const EXPORT_FORMAT_VERSION: u32 = 2;
|
||||
const MIN_SUPPORTED_EXPORT_FORMAT_VERSION: u32 = 1;
|
||||
|
||||
#[derive(
|
||||
Debug, Clone, Copy, PartialEq, Eq, PartialOrd, Ord, serde::Serialize, serde::Deserialize,
|
||||
@@ -65,6 +68,7 @@ pub enum ExportDomain {
|
||||
Wallets,
|
||||
Usage,
|
||||
Billing,
|
||||
Auxiliary,
|
||||
}
|
||||
|
||||
impl ExportDomain {
|
||||
@@ -87,10 +91,296 @@ impl ExportDomain {
|
||||
Self::Wallets => "wallets",
|
||||
Self::Usage => "usage",
|
||||
Self::Billing => "billing",
|
||||
Self::Auxiliary => "auxiliary",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
|
||||
struct AuxiliaryTable {
|
||||
name: &'static str,
|
||||
primary_key: &'static [&'static str],
|
||||
}
|
||||
|
||||
const AUXILIARY_TABLES: &[AuxiliaryTable] = &[
|
||||
AuxiliaryTable {
|
||||
name: "audit_logs",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "announcements",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "announcement_reads",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "management_tokens",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "user_preferences",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "user_sessions",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "ldap_configs",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "pool_member_scores",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "api_key_provider_mappings",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "provider_usage_tracking",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "gemini_file_mappings",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "routing_groups",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "routing_group_versions",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "routing_group_bindings",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "proxy_node_events",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "proxy_node_metrics_1m",
|
||||
primary_key: &["node_id", "bucket_start_unix_secs"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "proxy_node_metrics_1h",
|
||||
primary_key: &["node_id", "bucket_start_unix_secs"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "user_invite_codes",
|
||||
primary_key: &["user_id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "user_referrals",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "referral_rewards",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "payment_gateway_configs",
|
||||
primary_key: &["provider"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "billing_plans",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "user_plan_entitlements",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "entitlement_usage_ledgers",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "request_candidates",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "video_tasks",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "usage_body_blobs",
|
||||
primary_key: &["body_ref"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "usage_http_audits",
|
||||
primary_key: &["request_id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "usage_routing_snapshots",
|
||||
primary_key: &["request_id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "usage_counter_deltas",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "background_task_runs",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "background_task_events",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_hourly",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_summary",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_hourly_user",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_hourly_user_model",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "user_model_usage_counts",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_hourly_model",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_hourly_provider",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_daily",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_daily_model",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_daily_provider",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_daily_api_key",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_daily_error",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_user_daily",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_user_summary",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_user_daily_model",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_user_daily_provider",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_user_daily_api_format",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_daily_model_provider",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_user_daily_model_provider",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_daily_cost_savings",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_daily_cost_savings_provider",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_daily_cost_savings_model",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_daily_cost_savings_model_provider",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_user_daily_cost_savings",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_user_daily_cost_savings_provider",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_user_daily_cost_savings_model",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
AuxiliaryTable {
|
||||
name: "stats_user_daily_cost_savings_model_provider",
|
||||
primary_key: &["id"],
|
||||
},
|
||||
];
|
||||
|
||||
fn auxiliary_table(table_name: &str) -> Result<AuxiliaryTable, DataLayerError> {
|
||||
AUXILIARY_TABLES
|
||||
.iter()
|
||||
.copied()
|
||||
.find(|table| table.name == table_name)
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::InvalidInput(format!(
|
||||
"unsupported auxiliary export table '{table_name}'"
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
fn auxiliary_row_id(table: AuxiliaryTable, payload: &Value) -> Result<String, DataLayerError> {
|
||||
let object = payload.as_object().ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"auxiliary export row in table '{}' is not a JSON object",
|
||||
table.name
|
||||
))
|
||||
})?;
|
||||
let key = table
|
||||
.primary_key
|
||||
.iter()
|
||||
.map(|column| {
|
||||
object
|
||||
.get(*column)
|
||||
.filter(|value| !value.is_null())
|
||||
.cloned()
|
||||
.ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue(format!(
|
||||
"auxiliary export row in table '{}' has null or missing primary key column '{}'",
|
||||
table.name, column
|
||||
))
|
||||
})
|
||||
})
|
||||
.collect::<Result<Vec<_>, _>>()?;
|
||||
let encoded = serde_json::to_string(&key)
|
||||
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?;
|
||||
Ok(format!("{}:{encoded}", table.name))
|
||||
}
|
||||
|
||||
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
|
||||
pub struct DataExportManifest {
|
||||
pub format_version: u32,
|
||||
@@ -229,11 +519,167 @@ const USAGE_REQUEST_BODY_DETAIL_COLUMNS: &[&str] = &[
|
||||
"client_response_body_compressed",
|
||||
];
|
||||
|
||||
const USAGE_HTTP_BODY_DETAIL_COLUMNS: &[&str] = &[
|
||||
"request_body_ref",
|
||||
"provider_request_body_ref",
|
||||
"response_body_ref",
|
||||
"client_response_body_ref",
|
||||
"request_body_state",
|
||||
"provider_request_body_state",
|
||||
"response_body_state",
|
||||
"client_response_body_state",
|
||||
"body_capture_mode",
|
||||
];
|
||||
|
||||
#[cfg(all(feature = "postgres", feature = "sqlite"))]
|
||||
const REQUEST_BODY_DETAIL_TABLES: &[&str] = &["usage_body_blobs", "usage_http_audits"];
|
||||
const REQUEST_BODY_DETAIL_TABLES: &[&str] = &["usage_body_blobs"];
|
||||
#[cfg(all(feature = "postgres", feature = "sqlite"))]
|
||||
const LIFECYCLE_TABLES: &[&str] = &["_sqlx_migrations", "schema_backfills"];
|
||||
|
||||
fn import_column_stores_timestamp(column_name: &str) -> bool {
|
||||
column_name.ends_with("_at")
|
||||
|| column_name.ends_with("_unix_secs")
|
||||
|| column_name.ends_with("_unix_ms")
|
||||
|| column_name.ends_with("_date")
|
||||
|| matches!(
|
||||
column_name,
|
||||
"start_time" | "end_time" | "window_start" | "window_end" | "hour_utc" | "date"
|
||||
)
|
||||
}
|
||||
|
||||
fn import_timestamp_uses_millis(table_name: &str, column_name: &str) -> bool {
|
||||
if !column_name.ends_with("_unix_ms") {
|
||||
return false;
|
||||
}
|
||||
|
||||
// This legacy field is named `_unix_ms`, but every repository and API path
|
||||
// has always stored and consumed it as Unix seconds.
|
||||
let relation_name = table_name
|
||||
.rsplit('.')
|
||||
.next()
|
||||
.unwrap_or(table_name)
|
||||
.trim_matches(['"', '`']);
|
||||
!(relation_name == "usage" && column_name == "created_at_unix_ms")
|
||||
}
|
||||
|
||||
fn normalize_imported_integer_timestamp(
|
||||
driver_name: &str,
|
||||
table_name: &str,
|
||||
column_name: &str,
|
||||
value: &Value,
|
||||
) -> Result<Option<i64>, DataLayerError> {
|
||||
let invalid = || {
|
||||
DataLayerError::InvalidInput(format!(
|
||||
"{driver_name} import timestamp column '{column_name}' must contain an integer or supported datetime"
|
||||
))
|
||||
};
|
||||
|
||||
let timestamp = match value {
|
||||
Value::Null => return Ok(None),
|
||||
Value::Number(value) => value
|
||||
.as_i64()
|
||||
.or_else(|| value.as_u64().and_then(|value| i64::try_from(value).ok()))
|
||||
.ok_or_else(invalid)?,
|
||||
Value::String(value) => {
|
||||
if let Ok(timestamp) = value.trim().parse::<i64>() {
|
||||
timestamp
|
||||
} else {
|
||||
let datetime = parse_imported_datetime(value).ok_or_else(invalid)?;
|
||||
if import_timestamp_uses_millis(table_name, column_name) {
|
||||
datetime.timestamp_millis()
|
||||
} else {
|
||||
datetime.timestamp()
|
||||
}
|
||||
}
|
||||
}
|
||||
Value::Bool(_) | Value::Array(_) | Value::Object(_) => return Err(invalid()),
|
||||
};
|
||||
Ok(Some(timestamp))
|
||||
}
|
||||
|
||||
fn parse_imported_datetime(value: &str) -> Option<chrono::DateTime<chrono::Utc>> {
|
||||
let value = value.trim();
|
||||
if let Ok(datetime) = chrono::DateTime::parse_from_rfc3339(value) {
|
||||
return Some(datetime.with_timezone(&chrono::Utc));
|
||||
}
|
||||
if let Ok(datetime) = chrono::DateTime::parse_from_str(value, "%Y-%m-%d %H:%M:%S%.f%:z") {
|
||||
return Some(datetime.with_timezone(&chrono::Utc));
|
||||
}
|
||||
for format in ["%Y-%m-%d %H:%M:%S%.f", "%Y-%m-%dT%H:%M:%S%.f"] {
|
||||
if let Ok(datetime) = chrono::NaiveDateTime::parse_from_str(value, format) {
|
||||
return Some(datetime.and_utc());
|
||||
}
|
||||
}
|
||||
chrono::NaiveDate::parse_from_str(value, "%Y-%m-%d")
|
||||
.ok()
|
||||
.and_then(|date| date.and_hms_opt(0, 0, 0))
|
||||
.map(|datetime| datetime.and_utc())
|
||||
}
|
||||
|
||||
#[cfg(any(feature = "mysql", feature = "postgres", feature = "sqlite"))]
|
||||
fn normalize_imported_binary(
|
||||
driver_name: &str,
|
||||
column_name: &str,
|
||||
value: &Value,
|
||||
) -> Result<Option<Vec<u8>>, DataLayerError> {
|
||||
let invalid = |detail: &str| {
|
||||
DataLayerError::InvalidInput(format!(
|
||||
"{driver_name} import binary column '{column_name}' {detail}"
|
||||
))
|
||||
};
|
||||
match value {
|
||||
Value::Null => Ok(None),
|
||||
Value::Array(values) => values
|
||||
.iter()
|
||||
.map(|value| {
|
||||
value
|
||||
.as_u64()
|
||||
.and_then(|value| u8::try_from(value).ok())
|
||||
.ok_or_else(|| invalid("contains a non-byte array value"))
|
||||
})
|
||||
.collect::<Result<Vec<_>, _>>()
|
||||
.map(Some),
|
||||
Value::String(value) => {
|
||||
let encoded = value
|
||||
.trim()
|
||||
.strip_prefix("\\x")
|
||||
.ok_or_else(|| invalid("must use PostgreSQL \\x hex encoding"))?;
|
||||
if !encoded.len().is_multiple_of(2) {
|
||||
return Err(invalid("contains odd-length hex data"));
|
||||
}
|
||||
let mut bytes = Vec::with_capacity(encoded.len() / 2);
|
||||
for index in (0..encoded.len()).step_by(2) {
|
||||
let byte = u8::from_str_radix(&encoded[index..index + 2], 16).map_err(|err| {
|
||||
invalid(&format!(
|
||||
"contains invalid hex data at byte {}: {err}",
|
||||
index / 2
|
||||
))
|
||||
})?;
|
||||
bytes.push(byte);
|
||||
}
|
||||
Ok(Some(bytes))
|
||||
}
|
||||
Value::Bool(_) | Value::Number(_) | Value::Object(_) => {
|
||||
Err(invalid("must contain a byte array or PostgreSQL hex value"))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "postgres")]
|
||||
fn postgres_bytea_json_value(column_name: &str, value: &Value) -> Result<Value, DataLayerError> {
|
||||
let Some(bytes) = normalize_imported_binary("postgres", column_name, value)? else {
|
||||
return Ok(Value::Null);
|
||||
};
|
||||
let mut encoded = String::with_capacity(2 + bytes.len() * 2);
|
||||
encoded.push_str("\\x");
|
||||
for byte in bytes {
|
||||
use std::fmt::Write as _;
|
||||
write!(&mut encoded, "{byte:02x}")
|
||||
.map_err(|err| DataLayerError::UnexpectedValue(err.to_string()))?;
|
||||
}
|
||||
Ok(Value::String(encoded))
|
||||
}
|
||||
|
||||
pub fn encode_jsonl(records: &[DataExportRecord]) -> Result<String, DataLayerError> {
|
||||
validate_export_records(records)?;
|
||||
|
||||
@@ -300,10 +746,12 @@ pub fn validate_export_records(records: &[DataExportRecord]) -> Result<(), DataL
|
||||
"export JSONL must start with a manifest record".to_string(),
|
||||
));
|
||||
};
|
||||
if manifest.format_version != EXPORT_FORMAT_VERSION {
|
||||
if !(MIN_SUPPORTED_EXPORT_FORMAT_VERSION..=EXPORT_FORMAT_VERSION)
|
||||
.contains(&manifest.format_version)
|
||||
{
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"unsupported export format version {}; expected {}",
|
||||
manifest.format_version, EXPORT_FORMAT_VERSION
|
||||
"unsupported export format version {}; supported versions are {} through {}",
|
||||
manifest.format_version, MIN_SUPPORTED_EXPORT_FORMAT_VERSION, EXPORT_FORMAT_VERSION
|
||||
)));
|
||||
}
|
||||
|
||||
@@ -369,6 +817,7 @@ pub fn sqlite_core_export_domains() -> Vec<ExportDomain> {
|
||||
ExportDomain::Wallets,
|
||||
ExportDomain::Usage,
|
||||
ExportDomain::Billing,
|
||||
ExportDomain::Auxiliary,
|
||||
]
|
||||
}
|
||||
|
||||
@@ -490,23 +939,39 @@ pub async fn copy_database_records(
|
||||
import_database_jsonl(target, &encode_jsonl(&records)?).await
|
||||
}
|
||||
|
||||
fn omit_request_body_details_from_records(records: &mut [DataExportRecord]) {
|
||||
for record in records {
|
||||
fn omit_request_body_details_from_records(records: &mut Vec<DataExportRecord>) {
|
||||
records.retain_mut(|record| {
|
||||
let DataExportRecord::Row {
|
||||
domain: ExportDomain::Usage,
|
||||
payload,
|
||||
..
|
||||
domain, payload, ..
|
||||
} = record
|
||||
else {
|
||||
continue;
|
||||
return true;
|
||||
};
|
||||
|
||||
if let Some(object) = payload.as_object_mut() {
|
||||
for column_name in USAGE_REQUEST_BODY_DETAIL_COLUMNS {
|
||||
object.remove(*column_name);
|
||||
let Some(object) = payload.as_object_mut() else {
|
||||
return true;
|
||||
};
|
||||
match *domain {
|
||||
ExportDomain::Usage => {
|
||||
for column_name in USAGE_REQUEST_BODY_DETAIL_COLUMNS {
|
||||
object.remove(*column_name);
|
||||
}
|
||||
}
|
||||
ExportDomain::Auxiliary
|
||||
if object.get("__table").and_then(Value::as_str) == Some("usage_body_blobs") =>
|
||||
{
|
||||
return false;
|
||||
}
|
||||
ExportDomain::Auxiliary
|
||||
if object.get("__table").and_then(Value::as_str) == Some("usage_http_audits") =>
|
||||
{
|
||||
for column_name in USAGE_HTTP_BODY_DETAIL_COLUMNS {
|
||||
object.remove(*column_name);
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
true
|
||||
});
|
||||
}
|
||||
|
||||
#[cfg(all(feature = "postgres", feature = "sqlite"))]
|
||||
@@ -522,24 +987,24 @@ async fn copy_postgres_to_sqlite_from_target_schema(
|
||||
crate::driver::postgres::PostgresPoolFactory::new(source.to_postgres_config()?)?
|
||||
.connect_lazy()?;
|
||||
let sqlite_pool = crate::driver::sqlite::SqlitePoolFactory::new(target)?.connect_lazy()?;
|
||||
let mut postgres_tx = postgres_pool.begin().await.map_sql_err()?;
|
||||
sqlx::query("SET TRANSACTION ISOLATION LEVEL REPEATABLE READ READ ONLY")
|
||||
.execute(&mut *postgres_tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
let source_tables = load_postgres_public_table_names(&postgres_pool).await?;
|
||||
let source_tables = load_postgres_public_table_names(&mut postgres_tx).await?;
|
||||
let target_tables = load_sqlite_copy_table_names(&sqlite_pool).await?;
|
||||
|
||||
ensure_no_nonempty_source_tables_outside_target_schema(
|
||||
&postgres_pool,
|
||||
&mut postgres_tx,
|
||||
&source_tables,
|
||||
&target_tables,
|
||||
options,
|
||||
)
|
||||
.await?;
|
||||
|
||||
let mut imported = 0usize;
|
||||
sqlx::raw_sql("PRAGMA foreign_keys = OFF")
|
||||
.execute(&sqlite_pool)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
let mut table_plans = Vec::new();
|
||||
for table_name in target_tables {
|
||||
if copy_table_is_lifecycle(&table_name)
|
||||
|| copy_table_is_sqlite_internal(&table_name)
|
||||
@@ -550,7 +1015,7 @@ async fn copy_postgres_to_sqlite_from_target_schema(
|
||||
}
|
||||
|
||||
let table_plan = build_postgres_sqlite_copy_table_plan(
|
||||
&postgres_pool,
|
||||
&mut postgres_tx,
|
||||
&sqlite_pool,
|
||||
&table_name,
|
||||
options,
|
||||
@@ -559,22 +1024,39 @@ async fn copy_postgres_to_sqlite_from_target_schema(
|
||||
if table_plan.columns.is_empty() {
|
||||
continue;
|
||||
}
|
||||
imported = imported.saturating_add(
|
||||
copy_postgres_sqlite_table(&postgres_pool, &sqlite_pool, &table_plan).await?,
|
||||
);
|
||||
table_plans.push(table_plan);
|
||||
}
|
||||
|
||||
sqlx::raw_sql("PRAGMA foreign_keys = ON")
|
||||
.execute(&sqlite_pool)
|
||||
let mut connection = sqlite_pool.acquire().await.map_sql_err()?;
|
||||
sqlx::raw_sql("PRAGMA foreign_keys = OFF")
|
||||
.execute(&mut *connection)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
ensure_sqlite_foreign_key_check_passes(&sqlite_pool).await?;
|
||||
let copy_result = async {
|
||||
let mut tx = connection.begin().await.map_sql_err()?;
|
||||
let mut imported = 0usize;
|
||||
for table_plan in &table_plans {
|
||||
imported = imported.saturating_add(
|
||||
copy_postgres_sqlite_table(&mut postgres_tx, &mut tx, table_plan).await?,
|
||||
);
|
||||
}
|
||||
ensure_sqlite_foreign_key_check_passes(&mut tx).await?;
|
||||
tx.commit().await.map_sql_err()?;
|
||||
Ok::<_, DataLayerError>(imported)
|
||||
}
|
||||
.await;
|
||||
sqlx::raw_sql("PRAGMA foreign_keys = ON")
|
||||
.execute(&mut *connection)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let imported = copy_result?;
|
||||
postgres_tx.commit().await.map_sql_err()?;
|
||||
Ok(imported)
|
||||
}
|
||||
|
||||
#[cfg(all(feature = "postgres", feature = "sqlite"))]
|
||||
async fn ensure_no_nonempty_source_tables_outside_target_schema(
|
||||
postgres_pool: &crate::driver::postgres::PostgresPool,
|
||||
postgres_tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
|
||||
source_tables: &BTreeSet<String>,
|
||||
target_tables: &BTreeSet<String>,
|
||||
options: DataCopyOptions,
|
||||
@@ -587,7 +1069,7 @@ async fn ensure_no_nonempty_source_tables_outside_target_schema(
|
||||
{
|
||||
continue;
|
||||
}
|
||||
if postgres_public_table_has_rows(postgres_pool, table_name).await? {
|
||||
if postgres_public_table_has_rows(postgres_tx, table_name).await? {
|
||||
missing.push(table_name.clone());
|
||||
}
|
||||
}
|
||||
@@ -603,15 +1085,15 @@ async fn ensure_no_nonempty_source_tables_outside_target_schema(
|
||||
|
||||
#[cfg(all(feature = "postgres", feature = "sqlite"))]
|
||||
async fn build_postgres_sqlite_copy_table_plan(
|
||||
postgres_pool: &crate::driver::postgres::PostgresPool,
|
||||
postgres_tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
|
||||
sqlite_pool: &crate::driver::sqlite::SqlitePool,
|
||||
table_name: &str,
|
||||
options: DataCopyOptions,
|
||||
) -> Result<SchemaCopyTable, DataLayerError> {
|
||||
let sqlite_columns = load_sqlite_copy_columns(sqlite_pool, table_name).await?;
|
||||
let postgres_columns =
|
||||
load_postgres_import_columns(postgres_pool, &format!("public.{table_name}")).await?;
|
||||
let source_has_rows = postgres_public_table_has_rows(postgres_pool, table_name).await?;
|
||||
load_postgres_import_columns(&mut **postgres_tx, &format!("public.{table_name}")).await?;
|
||||
let source_has_rows = postgres_public_table_has_rows(postgres_tx, table_name).await?;
|
||||
let mut columns = Vec::new();
|
||||
|
||||
for sqlite_column in sqlite_columns {
|
||||
@@ -621,6 +1103,12 @@ async fn build_postgres_sqlite_copy_table_plan(
|
||||
{
|
||||
continue;
|
||||
}
|
||||
if options.omit_request_body_details
|
||||
&& table_name == "usage_http_audits"
|
||||
&& USAGE_HTTP_BODY_DETAIL_COLUMNS.contains(&sqlite_column.name.as_str())
|
||||
{
|
||||
continue;
|
||||
}
|
||||
|
||||
if let Some(postgres_column) = postgres_columns.get(&sqlite_column.name) {
|
||||
columns.push(SchemaCopyColumn {
|
||||
@@ -652,13 +1140,13 @@ async fn build_postgres_sqlite_copy_table_plan(
|
||||
|
||||
#[cfg(all(feature = "postgres", feature = "sqlite"))]
|
||||
async fn copy_postgres_sqlite_table(
|
||||
postgres_pool: &crate::driver::postgres::PostgresPool,
|
||||
sqlite_pool: &crate::driver::sqlite::SqlitePool,
|
||||
postgres_tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
|
||||
sqlite_tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
|
||||
table: &SchemaCopyTable,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
let source_sql = postgres_schema_copy_select_sql(table)?;
|
||||
let target_sql = sqlite_schema_copy_insert_sql(table)?;
|
||||
let mut rows = sqlx::query(&source_sql).fetch(postgres_pool);
|
||||
let mut rows = sqlx::query(&source_sql).fetch(&mut **postgres_tx);
|
||||
let mut imported = 0usize;
|
||||
|
||||
while let Some(row) = rows.try_next().await.map_sql_err()? {
|
||||
@@ -679,7 +1167,7 @@ async fn copy_postgres_sqlite_table(
|
||||
})?;
|
||||
query = bind_sqlite_copy_value(query, value, &column.sqlite)?;
|
||||
}
|
||||
query.execute(sqlite_pool).await.map_sql_err()?;
|
||||
query.execute(&mut **sqlite_tx).await.map_sql_err()?;
|
||||
imported = imported.saturating_add(1);
|
||||
}
|
||||
|
||||
@@ -694,7 +1182,7 @@ fn postgres_schema_copy_select_sql(table: &SchemaCopyTable) -> Result<String, Da
|
||||
);
|
||||
let mut payload_parts = Vec::new();
|
||||
for column in &table.columns {
|
||||
if let Some(expr) = postgres_schema_copy_override_expr(column)? {
|
||||
if let Some(expr) = postgres_schema_copy_override_expr(&table.table_name, column)? {
|
||||
payload_parts.push(sql_string_literal(&column.sqlite.name));
|
||||
payload_parts.push(expr);
|
||||
}
|
||||
@@ -729,6 +1217,7 @@ fn postgres_schema_copy_select_sql(table: &SchemaCopyTable) -> Result<String, Da
|
||||
|
||||
#[cfg(all(feature = "postgres", feature = "sqlite"))]
|
||||
fn postgres_schema_copy_override_expr(
|
||||
table_name: &str,
|
||||
column: &SchemaCopyColumn,
|
||||
) -> Result<Option<String>, DataLayerError> {
|
||||
let column_sql = format!("t.{}", postgres_quote_identifier(&column.sqlite.name)?);
|
||||
@@ -755,7 +1244,7 @@ fn postgres_schema_copy_override_expr(
|
||||
} else {
|
||||
column_sql.clone()
|
||||
};
|
||||
let multiplier = if sqlite_copy_column_stores_unix_millis(&column.sqlite.name) {
|
||||
let multiplier = if import_timestamp_uses_millis(table_name, &column.sqlite.name) {
|
||||
" * 1000"
|
||||
} else {
|
||||
""
|
||||
@@ -778,14 +1267,46 @@ fn sqlite_schema_copy_insert_sql(table: &SchemaCopyTable) -> Result<String, Data
|
||||
.collect::<Result<Vec<_>, _>>()?
|
||||
.join(", ");
|
||||
let placeholder_sql = vec!["?"; table.columns.len()].join(", ");
|
||||
let mut primary_key = table
|
||||
.columns
|
||||
.iter()
|
||||
.filter(|column| column.sqlite.primary_key_position > 0)
|
||||
.collect::<Vec<_>>();
|
||||
primary_key.sort_by_key(|column| column.sqlite.primary_key_position);
|
||||
if primary_key.is_empty() {
|
||||
return Ok(format!(
|
||||
"INSERT INTO {table_sql} ({column_sql}) VALUES ({placeholder_sql})"
|
||||
));
|
||||
}
|
||||
|
||||
let conflict_columns = primary_key
|
||||
.iter()
|
||||
.map(|column| sqlite_quote_identifier(&column.sqlite.name))
|
||||
.collect::<Result<Vec<_>, _>>()?
|
||||
.join(", ");
|
||||
let update_sql = table
|
||||
.columns
|
||||
.iter()
|
||||
.filter(|column| column.sqlite.primary_key_position == 0)
|
||||
.map(|column| {
|
||||
let quoted = sqlite_quote_identifier(&column.sqlite.name)?;
|
||||
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}")
|
||||
};
|
||||
Ok(format!(
|
||||
"INSERT OR REPLACE INTO {table_sql} ({column_sql}) VALUES ({placeholder_sql})"
|
||||
"INSERT INTO {table_sql} ({column_sql}) VALUES ({placeholder_sql}) {conflict_sql}"
|
||||
))
|
||||
}
|
||||
|
||||
#[cfg(all(feature = "postgres", feature = "sqlite"))]
|
||||
async fn load_postgres_public_table_names(
|
||||
pool: &crate::driver::postgres::PostgresPool,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
|
||||
) -> Result<BTreeSet<String>, DataLayerError> {
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
@@ -796,7 +1317,7 @@ WHERE table_schema = 'public'
|
||||
ORDER BY table_name
|
||||
"#,
|
||||
)
|
||||
.fetch_all(pool)
|
||||
.fetch_all(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
@@ -872,24 +1393,24 @@ async fn load_sqlite_copy_columns(
|
||||
|
||||
#[cfg(all(feature = "postgres", feature = "sqlite"))]
|
||||
async fn postgres_public_table_has_rows(
|
||||
pool: &crate::driver::postgres::PostgresPool,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
|
||||
table_name: &str,
|
||||
) -> Result<bool, DataLayerError> {
|
||||
let table_sql = format!("public.{}", postgres_quote_identifier(table_name)?);
|
||||
sqlx::query_scalar::<_, bool>(&format!(
|
||||
"SELECT EXISTS (SELECT 1 FROM {table_sql} LIMIT 1)"
|
||||
))
|
||||
.fetch_one(pool)
|
||||
.fetch_one(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()
|
||||
}
|
||||
|
||||
#[cfg(all(feature = "postgres", feature = "sqlite"))]
|
||||
async fn ensure_sqlite_foreign_key_check_passes(
|
||||
pool: &crate::driver::sqlite::SqlitePool,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let rows = sqlx::query("PRAGMA foreign_key_check")
|
||||
.fetch_all(pool)
|
||||
.fetch_all(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
if rows.is_empty() {
|
||||
@@ -935,11 +1456,6 @@ fn sqlite_copy_column_is_required(column: &SqliteCopyColumn) -> bool {
|
||||
(column.not_null || column.primary_key_position > 0) && !column.has_default
|
||||
}
|
||||
|
||||
#[cfg(all(feature = "postgres", feature = "sqlite"))]
|
||||
fn sqlite_copy_column_stores_unix_millis(column_name: &str) -> bool {
|
||||
column_name.ends_with("_unix_ms")
|
||||
}
|
||||
|
||||
#[cfg(all(feature = "postgres", feature = "sqlite"))]
|
||||
fn sqlite_copy_affinity(column: &SqliteCopyColumn) -> SqliteCopyAffinity {
|
||||
let declared_type = column.declared_type.to_ascii_uppercase();
|
||||
@@ -962,7 +1478,7 @@ fn sqlite_copy_affinity(column: &SqliteCopyColumn) -> SqliteCopyAffinity {
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(all(feature = "postgres", feature = "sqlite"))]
|
||||
#[cfg(feature = "postgres")]
|
||||
fn is_postgres_bytea_column(column: &PostgresImportColumn) -> bool {
|
||||
column.data_type == "bytea" || column.udt_name == "bytea"
|
||||
}
|
||||
@@ -1194,7 +1710,17 @@ fn filter_import_payload(
|
||||
for (column_name, value) in object {
|
||||
if target_columns.contains(column_name) {
|
||||
filtered.insert(column_name.clone(), value.clone());
|
||||
continue;
|
||||
}
|
||||
if value.is_null() {
|
||||
continue;
|
||||
}
|
||||
return Err(DataLayerError::InvalidInput(format!(
|
||||
"{} export row '{}' contains column '{}' that does not exist in {driver_name} table '{table_name}'",
|
||||
domain.as_str(),
|
||||
row.id,
|
||||
column_name
|
||||
)));
|
||||
}
|
||||
|
||||
if filtered.is_empty() {
|
||||
@@ -1208,7 +1734,6 @@ fn filter_import_payload(
|
||||
Ok(filtered)
|
||||
}
|
||||
|
||||
#[cfg(any(feature = "mysql", feature = "sqlite"))]
|
||||
fn payload_with_table(payload: Value, table_name: &str) -> Result<Value, DataLayerError> {
|
||||
let mut object = payload.as_object().cloned().ok_or_else(|| {
|
||||
DataLayerError::UnexpectedValue("export row payload must be a JSON object".to_string())
|
||||
@@ -1218,7 +1743,6 @@ fn payload_with_table(payload: Value, table_name: &str) -> Result<Value, DataLay
|
||||
Ok(Value::Object(object))
|
||||
}
|
||||
|
||||
#[cfg(any(feature = "mysql", feature = "sqlite"))]
|
||||
fn normalize_billing_payload(
|
||||
table_name: &str,
|
||||
object: &mut serde_json::Map<String, Value>,
|
||||
|
||||
@@ -1,4 +1,12 @@
|
||||
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,
|
||||
@@ -12,6 +20,12 @@ pub async fn export_mysql_jsonl(
|
||||
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),
|
||||
@@ -20,24 +34,29 @@ pub async fn export_mysql_jsonl(
|
||||
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(pool, &mut records).await?;
|
||||
export_mysql_billing_records(&mut tx, &mut records).await?;
|
||||
continue;
|
||||
}
|
||||
if domain == ExportDomain::Wallets {
|
||||
export_mysql_wallet_records(pool, &mut records).await?;
|
||||
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(pool).await.map_sql_err()?;
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -53,31 +72,40 @@ pub async fn import_mysql_plan(
|
||||
pool: &crate::driver::mysql::MysqlPool,
|
||||
plan: &DataImportPlan,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
let mut tx = pool.begin().await.map_sql_err()?;
|
||||
let mut imported = 0usize;
|
||||
let mut column_cache = BTreeMap::<String, ImportColumnNames>::new();
|
||||
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(pool, row, &mut column_cache).await?;
|
||||
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(pool, row, &mut column_cache).await?;
|
||||
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(pool, &mut column_cache, table_name).await?;
|
||||
mysql_import_columns_cached(&mut tx, &mut column_cache, table_name).await?;
|
||||
for row in plan.rows(*domain) {
|
||||
import_mysql_row(pool, table_name, *domain, row, &target_columns).await?;
|
||||
import_mysql_row(&mut tx, table_name, *domain, row, &target_columns).await?;
|
||||
imported = imported.saturating_add(1);
|
||||
}
|
||||
}
|
||||
tx.commit().await.map_sql_err()?;
|
||||
Ok(imported)
|
||||
}
|
||||
|
||||
@@ -106,9 +134,42 @@ fn mysql_domain_table(
|
||||
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,
|
||||
@@ -139,7 +200,7 @@ fn mysql_required_export_text(
|
||||
}
|
||||
|
||||
async fn export_mysql_billing_records(
|
||||
pool: &crate::driver::mysql::MysqlPool,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
records: &mut Vec<DataExportRecord>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
for (table_name, id_column) in [
|
||||
@@ -148,7 +209,7 @@ async fn export_mysql_billing_records(
|
||||
("usage_settlement_snapshots", "request_id"),
|
||||
] {
|
||||
let sql = format!("SELECT * FROM {table_name} ORDER BY {id_column} ASC");
|
||||
let rows = sqlx::query(&sql).fetch_all(pool).await.map_sql_err()?;
|
||||
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)
|
||||
@@ -169,12 +230,12 @@ async fn export_mysql_billing_records(
|
||||
}
|
||||
|
||||
async fn export_mysql_wallet_records(
|
||||
pool: &crate::driver::mysql::MysqlPool,
|
||||
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(pool).await.map_sql_err()?;
|
||||
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)
|
||||
@@ -195,53 +256,115 @@ async fn export_mysql_wallet_records(
|
||||
}
|
||||
|
||||
async fn import_mysql_row(
|
||||
pool: &crate::driver::mysql::MysqlPool,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
table_name: &str,
|
||||
domain: ExportDomain,
|
||||
row: &ExportRow,
|
||||
target_columns: &ImportColumnNames,
|
||||
target_columns: &MysqlImportColumns,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let object = filter_import_payload("mysql", table_name, domain, row, target_columns)?;
|
||||
let object = filter_import_payload("mysql", table_name, domain, row, &target_columns.names)?;
|
||||
|
||||
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 update_sql = columns
|
||||
.iter()
|
||||
.map(|column| {
|
||||
let quoted = mysql_quote_identifier(column)?;
|
||||
Ok(format!("{quoted} = VALUES({quoted})"))
|
||||
})
|
||||
.collect::<Result<Vec<_>, DataLayerError>>()?
|
||||
.join(", ");
|
||||
let sql = format!(
|
||||
"INSERT INTO {table_name} ({column_sql}) VALUES ({placeholder_sql}) ON DUPLICATE KEY UPDATE {update_sql}"
|
||||
);
|
||||
let sql = format!("INSERT INTO {table_name} ({column_sql}) VALUES ({placeholder_sql})");
|
||||
let mut query = sqlx::query(&sql);
|
||||
for column in columns {
|
||||
let value = object
|
||||
.get(column)
|
||||
.expect("column name came from payload object keys");
|
||||
query = bind_mysql_json_value(query, value)?;
|
||||
query = bind_mysql_import_column(query, &object, target_columns, table_name, column)?;
|
||||
}
|
||||
query.execute(pool).await.map_sql_err()?;
|
||||
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(
|
||||
pool: &crate::driver::mysql::MysqlPool,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
row: &ExportRow,
|
||||
column_cache: &mut BTreeMap<String, ImportColumnNames>,
|
||||
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(pool, column_cache, table_name).await?;
|
||||
let target_columns = mysql_import_columns_cached(tx, column_cache, table_name).await?;
|
||||
import_mysql_row(
|
||||
pool,
|
||||
tx,
|
||||
table_name,
|
||||
ExportDomain::Billing,
|
||||
&ExportRow {
|
||||
@@ -253,6 +376,27 @@ async fn import_mysql_billing_row(
|
||||
.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"),
|
||||
@@ -265,15 +409,15 @@ fn mysql_billing_table_name(table_name: &str) -> Result<&'static str, DataLayerE
|
||||
}
|
||||
|
||||
async fn import_mysql_wallet_row(
|
||||
pool: &crate::driver::mysql::MysqlPool,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
row: &ExportRow,
|
||||
column_cache: &mut BTreeMap<String, ImportColumnNames>,
|
||||
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(pool, column_cache, table_name).await?;
|
||||
let target_columns = mysql_import_columns_cached(tx, column_cache, table_name).await?;
|
||||
import_mysql_row(
|
||||
pool,
|
||||
tx,
|
||||
table_name,
|
||||
ExportDomain::Wallets,
|
||||
&ExportRow {
|
||||
@@ -311,51 +455,130 @@ fn mysql_wallet_table_name(table_name: &str) -> Result<&'static str, DataLayerEr
|
||||
}
|
||||
|
||||
async fn mysql_import_columns_cached(
|
||||
pool: &crate::driver::mysql::MysqlPool,
|
||||
cache: &mut BTreeMap<String, ImportColumnNames>,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
cache: &mut BTreeMap<String, MysqlImportColumns>,
|
||||
table_name: &str,
|
||||
) -> Result<ImportColumnNames, DataLayerError> {
|
||||
) -> Result<MysqlImportColumns, DataLayerError> {
|
||||
if let Some(columns) = cache.get(table_name) {
|
||||
return Ok(columns.clone());
|
||||
}
|
||||
|
||||
let columns = load_mysql_import_columns(pool, table_name).await?;
|
||||
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(
|
||||
pool: &crate::driver::mysql::MysqlPool,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::MySql>,
|
||||
table_name: &str,
|
||||
) -> Result<ImportColumnNames, DataLayerError> {
|
||||
) -> Result<MysqlImportColumns, DataLayerError> {
|
||||
let relation_name = table_name.trim_matches('`');
|
||||
let rows = sqlx::query(
|
||||
r#"
|
||||
SELECT COLUMN_NAME AS column_name
|
||||
SELECT
|
||||
COLUMN_NAME AS column_name,
|
||||
DATA_TYPE AS data_type,
|
||||
COLUMN_KEY AS column_key,
|
||||
ORDINAL_POSITION AS ordinal_position
|
||||
FROM information_schema.columns
|
||||
WHERE table_schema = DATABASE()
|
||||
AND table_name = ?
|
||||
"#,
|
||||
)
|
||||
.bind(relation_name)
|
||||
.fetch_all(pool)
|
||||
.fetch_all(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
let mut columns = ImportColumnNames::new();
|
||||
let mut columns = MysqlImportColumns::default();
|
||||
let mut primary_key = BTreeMap::new();
|
||||
for row in rows {
|
||||
columns.insert(row.try_get::<String, _>("column_name").map_sql_err()?);
|
||||
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::<i64, _>("ordinal_position").map_sql_err()?,
|
||||
name,
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
if columns.is_empty() {
|
||||
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(
|
||||
@@ -454,3 +677,30 @@ fn mysql_value_to_json(row: &sqlx::mysql::MySqlRow, index: usize) -> Result<Valu
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
#[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());
|
||||
}
|
||||
}
|
||||
|
||||
@@ -12,6 +12,11 @@ pub async fn export_postgres_jsonl(
|
||||
domains: Vec<ExportDomain>,
|
||||
created_at_unix_secs: u64,
|
||||
) -> Result<String, DataLayerError> {
|
||||
let mut tx = pool.begin().await.map_sql_err()?;
|
||||
sqlx::query("SET TRANSACTION ISOLATION LEVEL REPEATABLE READ READ ONLY")
|
||||
.execute(&mut *tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let manifest = DataExportManifest::new(
|
||||
created_at_unix_secs,
|
||||
Some(DatabaseDriver::Postgres),
|
||||
@@ -20,12 +25,16 @@ pub async fn export_postgres_jsonl(
|
||||
let mut records = vec![DataExportRecord::manifest(manifest)];
|
||||
|
||||
for domain in domains {
|
||||
if domain == ExportDomain::Auxiliary {
|
||||
export_postgres_auxiliary_records(&mut tx, &mut records).await?;
|
||||
continue;
|
||||
}
|
||||
if domain == ExportDomain::Billing {
|
||||
export_postgres_billing_records(pool, &mut records).await?;
|
||||
export_postgres_billing_records(&mut tx, &mut records).await?;
|
||||
continue;
|
||||
}
|
||||
if domain == ExportDomain::Wallets {
|
||||
export_postgres_wallet_records(pool, &mut records).await?;
|
||||
export_postgres_wallet_records(&mut tx, &mut records).await?;
|
||||
continue;
|
||||
}
|
||||
let (table_name, id_column) = postgres_domain_table(domain)?;
|
||||
@@ -34,7 +43,7 @@ pub async fn export_postgres_jsonl(
|
||||
let sql = format!(
|
||||
"SELECT {export_id_sql} AS export_id, to_jsonb(t) AS payload FROM {table_name} AS t ORDER BY {order_by}"
|
||||
);
|
||||
let rows = sqlx::query(&sql).fetch_all(pool).await.map_sql_err()?;
|
||||
let rows = sqlx::query(&sql).fetch_all(&mut *tx).await.map_sql_err()?;
|
||||
for row in rows {
|
||||
let id = row.try_get::<String, _>("export_id").map_sql_err()?;
|
||||
let payload = row.try_get::<Value, _>("payload").map_sql_err()?;
|
||||
@@ -42,6 +51,7 @@ pub async fn export_postgres_jsonl(
|
||||
}
|
||||
}
|
||||
|
||||
tx.commit().await.map_sql_err()?;
|
||||
encode_jsonl(&records)
|
||||
}
|
||||
|
||||
@@ -57,19 +67,27 @@ pub async fn import_postgres_plan(
|
||||
pool: &crate::driver::postgres::PostgresPool,
|
||||
plan: &DataImportPlan,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
let mut tx = pool.begin().await.map_sql_err()?;
|
||||
let mut imported = 0usize;
|
||||
let mut column_cache = BTreeMap::<String, PostgresImportColumns>::new();
|
||||
for domain in &plan.manifest.domains {
|
||||
if *domain == ExportDomain::Auxiliary {
|
||||
for row in plan.rows(*domain) {
|
||||
import_postgres_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_postgres_billing_row(pool, row, &mut column_cache).await?;
|
||||
import_postgres_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_postgres_wallet_row(pool, row, &mut column_cache).await?;
|
||||
import_postgres_wallet_row(&mut tx, row, &mut column_cache).await?;
|
||||
imported = imported.saturating_add(1);
|
||||
}
|
||||
continue;
|
||||
@@ -81,10 +99,10 @@ pub async fn import_postgres_plan(
|
||||
continue;
|
||||
}
|
||||
let target_columns =
|
||||
postgres_import_columns_cached(pool, &mut column_cache, table_name).await?;
|
||||
postgres_import_columns_cached(&mut tx, &mut column_cache, table_name).await?;
|
||||
for row in rows {
|
||||
import_postgres_row(
|
||||
pool,
|
||||
&mut tx,
|
||||
table_name,
|
||||
&conflict_columns,
|
||||
*domain,
|
||||
@@ -95,9 +113,52 @@ pub async fn import_postgres_plan(
|
||||
imported = imported.saturating_add(1);
|
||||
}
|
||||
}
|
||||
if !plan.rows(ExportDomain::Auxiliary).is_empty() {
|
||||
reset_postgres_auxiliary_sequences(&mut tx).await?;
|
||||
}
|
||||
tx.commit().await.map_sql_err()?;
|
||||
Ok(imported)
|
||||
}
|
||||
|
||||
async fn reset_postgres_auxiliary_sequences(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
for table in AUXILIARY_TABLES {
|
||||
let [primary_key] = table.primary_key else {
|
||||
continue;
|
||||
};
|
||||
let relation_name = format!("public.{}", table.name);
|
||||
let sequence =
|
||||
sqlx::query_scalar::<_, Option<String>>("SELECT pg_get_serial_sequence($1, $2)")
|
||||
.bind(&relation_name)
|
||||
.bind(*primary_key)
|
||||
.fetch_one(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let Some(sequence) = sequence else {
|
||||
continue;
|
||||
};
|
||||
|
||||
let table_sql = postgres_quote_identifier(table.name)?;
|
||||
let primary_key_sql = postgres_quote_identifier(primary_key)?;
|
||||
let maximum = sqlx::query_scalar::<_, Option<i64>>(&format!(
|
||||
"SELECT MAX({primary_key_sql})::bigint FROM public.{table_sql}"
|
||||
))
|
||||
.fetch_one(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
let (value, is_called) = maximum.map_or((1_i64, false), |value| (value, true));
|
||||
sqlx::query("SELECT setval($1::regclass, $2, $3)")
|
||||
.bind(sequence)
|
||||
.bind(value)
|
||||
.bind(is_called)
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn postgres_domain_table(
|
||||
domain: ExportDomain,
|
||||
) -> Result<(&'static str, &'static str), DataLayerError> {
|
||||
@@ -125,9 +186,44 @@ fn postgres_domain_table(
|
||||
"postgres billing export uses multiple tables and must be handled as a domain"
|
||||
.to_string(),
|
||||
)),
|
||||
ExportDomain::Auxiliary => Err(DataLayerError::InvalidInput(
|
||||
"postgres auxiliary export uses multiple tables and must be handled as a domain"
|
||||
.to_string(),
|
||||
)),
|
||||
}
|
||||
}
|
||||
|
||||
async fn export_postgres_auxiliary_records(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
|
||||
records: &mut Vec<DataExportRecord>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
for table in AUXILIARY_TABLES {
|
||||
let table_sql = postgres_quote_identifier(table.name)?;
|
||||
let order_sql = table
|
||||
.primary_key
|
||||
.iter()
|
||||
.map(|column| postgres_quote_identifier(column).map(|column| format!("{column} ASC")))
|
||||
.collect::<Result<Vec<_>, _>>()?
|
||||
.join(", ");
|
||||
let rows = sqlx::query(&format!(
|
||||
"SELECT to_jsonb(t) AS payload FROM public.{table_sql} AS t ORDER BY {order_sql}"
|
||||
))
|
||||
.fetch_all(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
for row in rows {
|
||||
let payload = row.try_get::<Value, _>("payload").map_sql_err()?;
|
||||
let id = auxiliary_row_id(*table, &payload)?;
|
||||
records.push(DataExportRecord::row(
|
||||
ExportDomain::Auxiliary,
|
||||
id,
|
||||
payload_with_table(payload, table.name)?,
|
||||
));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
fn postgres_export_id_sql(domain: ExportDomain, id_column: &str) -> String {
|
||||
if domain == ExportDomain::UserGroupMembers {
|
||||
"group_id::text || ':' || user_id::text".to_string()
|
||||
@@ -145,7 +241,7 @@ fn postgres_conflict_columns(domain: ExportDomain, id_column: &str) -> Vec<&str>
|
||||
}
|
||||
|
||||
async fn postgres_import_columns_cached(
|
||||
pool: &crate::driver::postgres::PostgresPool,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
|
||||
cache: &mut BTreeMap<String, PostgresImportColumns>,
|
||||
table_name: &str,
|
||||
) -> Result<PostgresImportColumns, DataLayerError> {
|
||||
@@ -153,13 +249,13 @@ async fn postgres_import_columns_cached(
|
||||
return Ok(columns.clone());
|
||||
}
|
||||
|
||||
let columns = load_postgres_import_columns(pool, table_name).await?;
|
||||
let columns = load_postgres_import_columns(&mut **tx, table_name).await?;
|
||||
cache.insert(table_name.to_string(), columns.clone());
|
||||
Ok(columns)
|
||||
}
|
||||
|
||||
pub(super) async fn load_postgres_import_columns(
|
||||
pool: &crate::driver::postgres::PostgresPool,
|
||||
pub(super) async fn load_postgres_import_columns<'e>(
|
||||
executor: impl sqlx::Executor<'e, Database = sqlx::Postgres>,
|
||||
table_name: &str,
|
||||
) -> Result<PostgresImportColumns, DataLayerError> {
|
||||
let (schema_name, relation_name) = postgres_table_parts(table_name)?;
|
||||
@@ -173,7 +269,7 @@ WHERE table_schema = $1
|
||||
)
|
||||
.bind(schema_name)
|
||||
.bind(relation_name)
|
||||
.fetch_all(pool)
|
||||
.fetch_all(executor)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
|
||||
@@ -223,7 +319,7 @@ fn postgres_table_parts(table_name: &str) -> Result<(&str, &str), DataLayerError
|
||||
}
|
||||
|
||||
async fn export_postgres_billing_records(
|
||||
pool: &crate::driver::postgres::PostgresPool,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
|
||||
records: &mut Vec<DataExportRecord>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
for (table_name, export_table, id_column) in [
|
||||
@@ -238,7 +334,7 @@ async fn export_postgres_billing_records(
|
||||
let sql = format!(
|
||||
"SELECT {id_column}::text AS export_id, to_jsonb(t) || jsonb_build_object('__table', '{export_table}') AS payload FROM {table_name} AS t ORDER BY {id_column} ASC"
|
||||
);
|
||||
let rows = sqlx::query(&sql).fetch_all(pool).await.map_sql_err()?;
|
||||
let rows = sqlx::query(&sql).fetch_all(&mut **tx).await.map_sql_err()?;
|
||||
for row in rows {
|
||||
let id = row.try_get::<String, _>("export_id").map_sql_err()?;
|
||||
let payload = row.try_get::<Value, _>("payload").map_sql_err()?;
|
||||
@@ -253,14 +349,14 @@ async fn export_postgres_billing_records(
|
||||
}
|
||||
|
||||
async fn export_postgres_wallet_records(
|
||||
pool: &crate::driver::postgres::PostgresPool,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
|
||||
records: &mut Vec<DataExportRecord>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
for (table_name, export_table, id_column) in postgres_wallet_tables() {
|
||||
let sql = format!(
|
||||
"SELECT {id_column}::text AS export_id, to_jsonb(t) || jsonb_build_object('__table', '{export_table}') AS payload FROM {table_name} AS t ORDER BY {id_column} ASC"
|
||||
);
|
||||
let rows = sqlx::query(&sql).fetch_all(pool).await.map_sql_err()?;
|
||||
let rows = sqlx::query(&sql).fetch_all(&mut **tx).await.map_sql_err()?;
|
||||
for row in rows {
|
||||
let id = row.try_get::<String, _>("export_id").map_sql_err()?;
|
||||
let payload = row.try_get::<Value, _>("payload").map_sql_err()?;
|
||||
@@ -275,7 +371,7 @@ async fn export_postgres_wallet_records(
|
||||
}
|
||||
|
||||
async fn import_postgres_row(
|
||||
pool: &crate::driver::postgres::PostgresPool,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
|
||||
table_name: &str,
|
||||
conflict_columns: &[&str],
|
||||
domain: ExportDomain,
|
||||
@@ -316,7 +412,7 @@ async fn import_postgres_row(
|
||||
|
||||
sqlx::query(&sql)
|
||||
.bind(&payload)
|
||||
.execute(pool)
|
||||
.execute(&mut **tx)
|
||||
.await
|
||||
.map_sql_err()?;
|
||||
Ok(())
|
||||
@@ -351,7 +447,7 @@ pub(super) fn normalize_postgres_import_payload(
|
||||
}
|
||||
normalized.insert(
|
||||
column_name.clone(),
|
||||
normalize_postgres_import_value(column_name, target_column, value)?,
|
||||
normalize_postgres_import_value(table_name, column_name, target_column, value)?,
|
||||
);
|
||||
continue;
|
||||
}
|
||||
@@ -380,6 +476,7 @@ pub(super) fn normalize_postgres_import_payload(
|
||||
}
|
||||
|
||||
fn normalize_postgres_import_value(
|
||||
table_name: &str,
|
||||
column_name: &str,
|
||||
target_column: &PostgresImportColumn,
|
||||
value: &Value,
|
||||
@@ -392,7 +489,10 @@ fn normalize_postgres_import_value(
|
||||
return normalize_postgres_boolean_value(column_name, value);
|
||||
}
|
||||
if is_postgres_timestamp_column(target_column) {
|
||||
return normalize_postgres_timestamp_value(column_name, value);
|
||||
return normalize_postgres_timestamp_value(table_name, column_name, value);
|
||||
}
|
||||
if is_postgres_bytea_column(target_column) {
|
||||
return postgres_bytea_json_value(column_name, value);
|
||||
}
|
||||
if is_postgres_json_column(target_column) {
|
||||
return normalize_postgres_json_value(value);
|
||||
@@ -450,6 +550,7 @@ fn normalize_postgres_boolean_value(
|
||||
}
|
||||
|
||||
fn normalize_postgres_timestamp_value(
|
||||
table_name: &str,
|
||||
column_name: &str,
|
||||
value: &Value,
|
||||
) -> Result<Value, DataLayerError> {
|
||||
@@ -465,7 +566,7 @@ fn normalize_postgres_timestamp_value(
|
||||
)));
|
||||
};
|
||||
|
||||
let datetime = if column_name.ends_with("_unix_ms")
|
||||
let datetime = if import_timestamp_uses_millis(table_name, column_name)
|
||||
|| timestamp >= 100_000_000_000
|
||||
|| timestamp <= -100_000_000_000
|
||||
{
|
||||
@@ -497,17 +598,17 @@ fn normalize_postgres_json_value(value: &Value) -> Result<Value, DataLayerError>
|
||||
}
|
||||
|
||||
async fn import_postgres_billing_row(
|
||||
pool: &crate::driver::postgres::PostgresPool,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
|
||||
row: &ExportRow,
|
||||
column_cache: &mut BTreeMap<String, PostgresImportColumns>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let (export_table_name, payload) = billing_payload_table(row)?;
|
||||
let table_name = postgres_billing_table_name(&export_table_name)?;
|
||||
let target_columns = postgres_import_columns_cached(pool, column_cache, table_name).await?;
|
||||
let (table_name, conflict_column) = postgres_billing_table_name(&export_table_name)?;
|
||||
let target_columns = postgres_import_columns_cached(tx, column_cache, table_name).await?;
|
||||
import_postgres_row(
|
||||
pool,
|
||||
tx,
|
||||
table_name,
|
||||
&["id"],
|
||||
&[conflict_column],
|
||||
ExportDomain::Billing,
|
||||
&ExportRow {
|
||||
id: row.id.clone(),
|
||||
@@ -518,27 +619,66 @@ async fn import_postgres_billing_row(
|
||||
.await
|
||||
}
|
||||
|
||||
fn postgres_billing_table_name(table_name: &str) -> Result<&'static str, DataLayerError> {
|
||||
async fn import_postgres_auxiliary_row(
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
|
||||
row: &ExportRow,
|
||||
column_cache: &mut BTreeMap<String, PostgresImportColumns>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let (table_name, payload) = domain_payload_table(row, "auxiliary", None)?;
|
||||
let table = auxiliary_table(&table_name)?;
|
||||
let target_table = format!("public.{}", postgres_quote_identifier(table.name)?);
|
||||
let target_columns = postgres_import_columns_cached(tx, column_cache, &target_table).await?;
|
||||
import_postgres_row(
|
||||
tx,
|
||||
&target_table,
|
||||
table.primary_key,
|
||||
ExportDomain::Auxiliary,
|
||||
&ExportRow {
|
||||
id: row.id.clone(),
|
||||
payload,
|
||||
},
|
||||
&target_columns,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
fn postgres_billing_table_name(
|
||||
table_name: &str,
|
||||
) -> Result<(&'static str, &'static str), DataLayerError> {
|
||||
match table_name {
|
||||
"billing_rules" => Ok("public.billing_rules"),
|
||||
"dimension_collectors" => Ok("public.dimension_collectors"),
|
||||
"usage_settlement_snapshots" => Ok("public.usage_settlement_snapshots"),
|
||||
"billing_rules" => Ok(("public.billing_rules", "id")),
|
||||
"dimension_collectors" => Ok(("public.dimension_collectors", "id")),
|
||||
"usage_settlement_snapshots" => Ok(("public.usage_settlement_snapshots", "request_id")),
|
||||
other => Err(DataLayerError::InvalidInput(format!(
|
||||
"unsupported postgres billing export table '{other}'"
|
||||
))),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod billing_table_tests {
|
||||
use super::postgres_billing_table_name;
|
||||
|
||||
#[test]
|
||||
fn settlement_snapshot_import_uses_request_id_conflict_key() {
|
||||
assert_eq!(
|
||||
postgres_billing_table_name("usage_settlement_snapshots")
|
||||
.expect("settlement snapshot table should be supported"),
|
||||
("public.usage_settlement_snapshots", "request_id")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
async fn import_postgres_wallet_row(
|
||||
pool: &crate::driver::postgres::PostgresPool,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Postgres>,
|
||||
row: &ExportRow,
|
||||
column_cache: &mut BTreeMap<String, PostgresImportColumns>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let (export_table_name, payload) = domain_payload_table(row, "wallet", Some("wallets"))?;
|
||||
let (table_name, id_column) = postgres_wallet_table_name(&export_table_name)?;
|
||||
let target_columns = postgres_import_columns_cached(pool, column_cache, table_name).await?;
|
||||
let target_columns = postgres_import_columns_cached(tx, column_cache, table_name).await?;
|
||||
import_postgres_row(
|
||||
pool,
|
||||
tx,
|
||||
table_name,
|
||||
&[id_column],
|
||||
ExportDomain::Wallets,
|
||||
|
||||
@@ -1,5 +1,12 @@
|
||||
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,
|
||||
@@ -12,6 +19,7 @@ pub async fn export_sqlite_jsonl(
|
||||
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),
|
||||
@@ -20,24 +28,29 @@ pub async fn export_sqlite_jsonl(
|
||||
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(pool, &mut records).await?;
|
||||
export_sqlite_billing_records(&mut tx, &mut records).await?;
|
||||
continue;
|
||||
}
|
||||
if domain == ExportDomain::Wallets {
|
||||
export_sqlite_wallet_records(pool, &mut records).await?;
|
||||
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(pool).await.map_sql_err()?;
|
||||
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)
|
||||
}
|
||||
|
||||
@@ -53,31 +66,40 @@ pub async fn import_sqlite_plan(
|
||||
pool: &crate::driver::sqlite::SqlitePool,
|
||||
plan: &DataImportPlan,
|
||||
) -> Result<usize, DataLayerError> {
|
||||
let mut tx = pool.begin().await.map_sql_err()?;
|
||||
let mut imported = 0usize;
|
||||
let mut column_cache = BTreeMap::<String, ImportColumnNames>::new();
|
||||
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(pool, row, &mut column_cache).await?;
|
||||
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(pool, row, &mut column_cache).await?;
|
||||
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(pool, &mut column_cache, table_name).await?;
|
||||
sqlite_import_columns_cached(&mut tx, &mut column_cache, table_name).await?;
|
||||
for row in plan.rows(*domain) {
|
||||
import_sqlite_row(pool, table_name, *domain, row, &target_columns).await?;
|
||||
import_sqlite_row(&mut tx, table_name, *domain, row, &target_columns).await?;
|
||||
imported = imported.saturating_add(1);
|
||||
}
|
||||
}
|
||||
tx.commit().await.map_sql_err()?;
|
||||
Ok(imported)
|
||||
}
|
||||
|
||||
@@ -107,9 +129,42 @@ fn sqlite_domain_table(
|
||||
"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,
|
||||
@@ -140,7 +195,7 @@ fn sqlite_required_export_text(
|
||||
}
|
||||
|
||||
async fn export_sqlite_billing_records(
|
||||
pool: &crate::driver::sqlite::SqlitePool,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
|
||||
records: &mut Vec<DataExportRecord>,
|
||||
) -> Result<(), DataLayerError> {
|
||||
for table_name in [
|
||||
@@ -154,7 +209,7 @@ async fn export_sqlite_billing_records(
|
||||
"id"
|
||||
};
|
||||
let sql = format!("SELECT * FROM {table_name} ORDER BY {id_column} ASC");
|
||||
let rows = sqlx::query(&sql).fetch_all(pool).await.map_sql_err()?;
|
||||
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)
|
||||
@@ -175,12 +230,12 @@ async fn export_sqlite_billing_records(
|
||||
}
|
||||
|
||||
async fn export_sqlite_wallet_records(
|
||||
pool: &crate::driver::sqlite::SqlitePool,
|
||||
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(pool).await.map_sql_err()?;
|
||||
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)
|
||||
@@ -201,13 +256,13 @@ async fn export_sqlite_wallet_records(
|
||||
}
|
||||
|
||||
async fn import_sqlite_row(
|
||||
pool: &crate::driver::sqlite::SqlitePool,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
|
||||
table_name: &str,
|
||||
domain: ExportDomain,
|
||||
row: &ExportRow,
|
||||
target_columns: &ImportColumnNames,
|
||||
target_columns: &SqliteImportColumns,
|
||||
) -> Result<(), DataLayerError> {
|
||||
let object = filter_import_payload("sqlite", table_name, domain, row, target_columns)?;
|
||||
let object = filter_import_payload("sqlite", table_name, domain, row, &target_columns.names)?;
|
||||
|
||||
let columns = object.keys().map(String::as_str).collect::<Vec<_>>();
|
||||
let column_sql = columns
|
||||
@@ -216,29 +271,55 @@ async fn import_sqlite_row(
|
||||
.collect::<Result<Vec<_>, _>>()?
|
||||
.join(", ");
|
||||
let placeholder_sql = vec!["?"; columns.len()].join(", ");
|
||||
let sql =
|
||||
format!("INSERT OR REPLACE INTO {table_name} ({column_sql}) VALUES ({placeholder_sql})");
|
||||
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");
|
||||
query = bind_sqlite_json_value(query, value)?;
|
||||
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(pool).await.map_sql_err()?;
|
||||
query.execute(&mut **tx).await.map_sql_err()?;
|
||||
Ok(())
|
||||
}
|
||||
|
||||
async fn import_sqlite_billing_row(
|
||||
pool: &crate::driver::sqlite::SqlitePool,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
|
||||
row: &ExportRow,
|
||||
column_cache: &mut BTreeMap<String, ImportColumnNames>,
|
||||
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(pool, column_cache, table_name).await?;
|
||||
let target_columns = sqlite_import_columns_cached(tx, column_cache, table_name).await?;
|
||||
import_sqlite_row(
|
||||
pool,
|
||||
tx,
|
||||
table_name,
|
||||
ExportDomain::Billing,
|
||||
&ExportRow {
|
||||
@@ -250,6 +331,27 @@ async fn import_sqlite_billing_row(
|
||||
.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"),
|
||||
@@ -262,15 +364,15 @@ fn sqlite_billing_table_name(table_name: &str) -> Result<&'static str, DataLayer
|
||||
}
|
||||
|
||||
async fn import_sqlite_wallet_row(
|
||||
pool: &crate::driver::sqlite::SqlitePool,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
|
||||
row: &ExportRow,
|
||||
column_cache: &mut BTreeMap<String, ImportColumnNames>,
|
||||
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(pool, column_cache, table_name).await?;
|
||||
let target_columns = sqlite_import_columns_cached(tx, column_cache, table_name).await?;
|
||||
import_sqlite_row(
|
||||
pool,
|
||||
tx,
|
||||
table_name,
|
||||
ExportDomain::Wallets,
|
||||
&ExportRow {
|
||||
@@ -308,39 +410,81 @@ fn sqlite_wallet_table_name(table_name: &str) -> Result<&'static str, DataLayerE
|
||||
}
|
||||
|
||||
async fn sqlite_import_columns_cached(
|
||||
pool: &crate::driver::sqlite::SqlitePool,
|
||||
cache: &mut BTreeMap<String, ImportColumnNames>,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
|
||||
cache: &mut BTreeMap<String, SqliteImportColumns>,
|
||||
table_name: &str,
|
||||
) -> Result<ImportColumnNames, DataLayerError> {
|
||||
) -> Result<SqliteImportColumns, DataLayerError> {
|
||||
if let Some(columns) = cache.get(table_name) {
|
||||
return Ok(columns.clone());
|
||||
}
|
||||
|
||||
let columns = load_sqlite_import_columns(pool, table_name).await?;
|
||||
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(
|
||||
pool: &crate::driver::sqlite::SqlitePool,
|
||||
tx: &mut sqlx::Transaction<'_, sqlx::Sqlite>,
|
||||
table_name: &str,
|
||||
) -> Result<ImportColumnNames, DataLayerError> {
|
||||
) -> Result<SqliteImportColumns, DataLayerError> {
|
||||
let sql = format!("PRAGMA table_info({table_name})");
|
||||
let rows = sqlx::query(&sql).fetch_all(pool).await.map_sql_err()?;
|
||||
let mut columns = ImportColumnNames::new();
|
||||
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 {
|
||||
columns.insert(row.try_get::<String, _>("name").map_sql_err()?);
|
||||
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.is_empty() {
|
||||
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() {
|
||||
|
||||
@@ -1,14 +1,16 @@
|
||||
use std::collections::BTreeMap;
|
||||
use std::collections::{BTreeMap, BTreeSet};
|
||||
|
||||
use serde_json::json;
|
||||
|
||||
use super::{
|
||||
build_import_plan, decode_jsonl, encode_jsonl, export_mysql_core_jsonl, export_mysql_jsonl,
|
||||
export_postgres_core_jsonl, export_sqlite_core_jsonl, import_mysql_jsonl,
|
||||
import_postgres_jsonl, import_sqlite_jsonl, mysql_core_export_domains,
|
||||
normalize_postgres_import_payload, postgres_core_export_domains, sqlite_core_export_domains,
|
||||
DataExportManifest, DataExportRecord, DataImportPlan, ExportDomain, ExportRow,
|
||||
PostgresImportColumn,
|
||||
export_postgres_core_jsonl, export_sqlite_core_jsonl, filter_import_payload,
|
||||
import_mysql_jsonl, import_postgres_jsonl, import_sqlite_jsonl, mysql_core_export_domains,
|
||||
normalize_imported_binary, normalize_imported_integer_timestamp,
|
||||
normalize_postgres_import_payload, postgres_bytea_json_value, postgres_core_export_domains,
|
||||
sqlite_core_export_domains, sqlite_schema_copy_insert_sql, DataExportManifest,
|
||||
DataExportRecord, DataImportPlan, ExportDomain, ExportRow, PostgresImportColumn,
|
||||
SchemaCopyColumn, SchemaCopyTable, SqliteCopyColumn, AUXILIARY_TABLES,
|
||||
};
|
||||
use crate::driver::postgres::{PostgresPoolConfig, PostgresPoolFactory};
|
||||
use crate::lifecycle::migrate::{
|
||||
@@ -64,6 +66,81 @@ fn jsonl_round_trips_manifest_and_domain_rows() {
|
||||
fn core_export_domains_match_across_sql_drivers() {
|
||||
assert_eq!(sqlite_core_export_domains(), mysql_core_export_domains());
|
||||
assert_eq!(sqlite_core_export_domains(), postgres_core_export_domains());
|
||||
assert!(sqlite_core_export_domains().contains(&ExportDomain::Auxiliary));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn sqlite_core_export_covers_every_portable_table() {
|
||||
let pool = sqlx::sqlite::SqlitePoolOptions::new()
|
||||
.max_connections(1)
|
||||
.connect("sqlite::memory:")
|
||||
.await
|
||||
.expect("sqlite pool should connect");
|
||||
run_sqlite_migrations(&pool)
|
||||
.await
|
||||
.expect("sqlite migrations should run");
|
||||
|
||||
let schema_tables = sqlx::query_scalar::<_, String>(
|
||||
r#"
|
||||
SELECT name
|
||||
FROM sqlite_master
|
||||
WHERE type = 'table'
|
||||
AND name NOT LIKE 'sqlite_%'
|
||||
AND name NOT IN ('_sqlx_migrations', 'schema_backfills')
|
||||
ORDER BY name
|
||||
"#,
|
||||
)
|
||||
.fetch_all(&pool)
|
||||
.await
|
||||
.expect("sqlite schema tables should load")
|
||||
.into_iter()
|
||||
.collect::<BTreeSet<_>>();
|
||||
|
||||
let mut exported_tables = [
|
||||
"users",
|
||||
"api_keys",
|
||||
"providers",
|
||||
"provider_api_keys",
|
||||
"provider_endpoints",
|
||||
"global_models",
|
||||
"models",
|
||||
"auth_modules",
|
||||
"oauth_providers",
|
||||
"user_oauth_links",
|
||||
"user_groups",
|
||||
"user_group_members",
|
||||
"proxy_nodes",
|
||||
"system_configs",
|
||||
"usage",
|
||||
"wallets",
|
||||
"wallet_transactions",
|
||||
"wallet_daily_usage_ledgers",
|
||||
"payment_orders",
|
||||
"payment_callbacks",
|
||||
"refund_requests",
|
||||
"redeem_code_batches",
|
||||
"redeem_codes",
|
||||
"billing_rules",
|
||||
"dimension_collectors",
|
||||
"usage_settlement_snapshots",
|
||||
]
|
||||
.into_iter()
|
||||
.map(str::to_string)
|
||||
.collect::<BTreeSet<_>>();
|
||||
exported_tables.extend(AUXILIARY_TABLES.iter().map(|table| table.name.to_string()));
|
||||
|
||||
assert_eq!(schema_tables, exported_tables);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn version_one_exports_remain_importable_after_full_export_expansion() {
|
||||
let records = decode_jsonl(
|
||||
r#"{"record_type":"manifest","manifest":{"format_version":1,"created_at_unix_secs":1,"source_driver":null,"domains":["users"]}}
|
||||
{"record_type":"row","domain":"users","id":"user-1","payload":{"id":"user-1"}}"#,
|
||||
)
|
||||
.expect("version one exports should remain supported");
|
||||
|
||||
assert_eq!(records.len(), 2);
|
||||
}
|
||||
|
||||
#[test]
|
||||
@@ -171,6 +248,70 @@ fn postgres_import_payload_normalizes_sqlite_values_for_target_columns() {
|
||||
assert!(!normalized.contains_key("legacy_nullable"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cross_driver_timestamp_normalization_preserves_usage_second_contract() {
|
||||
assert_eq!(
|
||||
normalize_imported_integer_timestamp(
|
||||
"sqlite",
|
||||
r#""usage""#,
|
||||
"created_at_unix_ms",
|
||||
&json!("1970-01-01T00:00:01.234900Z"),
|
||||
)
|
||||
.expect("usage timestamp should normalize"),
|
||||
Some(1),
|
||||
);
|
||||
assert_eq!(
|
||||
normalize_imported_integer_timestamp(
|
||||
"mysql",
|
||||
"request_candidates",
|
||||
"created_at_unix_ms",
|
||||
&json!("1970-01-01T00:00:01.234900Z"),
|
||||
)
|
||||
.expect("millisecond timestamp should normalize"),
|
||||
Some(1_234),
|
||||
);
|
||||
|
||||
let target_columns = BTreeMap::from([(
|
||||
"created_at_unix_ms".to_string(),
|
||||
postgres_column("timestamp with time zone", "timestamptz"),
|
||||
)]);
|
||||
let row = ExportRow {
|
||||
id: "usage-1".to_string(),
|
||||
payload: json!({ "created_at_unix_ms": 1_700_000_000 }),
|
||||
};
|
||||
let normalized = normalize_postgres_import_payload(
|
||||
"public.usage",
|
||||
ExportDomain::Usage,
|
||||
&row,
|
||||
&target_columns,
|
||||
)
|
||||
.expect("postgres usage timestamp should normalize");
|
||||
assert_eq!(
|
||||
normalized["created_at_unix_ms"],
|
||||
json!("2023-11-14T22:13:20+00:00")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn cross_driver_binary_normalization_preserves_raw_bytes() {
|
||||
assert_eq!(
|
||||
normalize_imported_binary("sqlite", "payload_gzip", &json!([0, 1, 127, 255]))
|
||||
.expect("byte array should normalize"),
|
||||
Some(vec![0, 1, 127, 255]),
|
||||
);
|
||||
assert_eq!(
|
||||
normalize_imported_binary("mysql", "payload_gzip", &json!("\\x00017fff"))
|
||||
.expect("postgres hex should normalize"),
|
||||
Some(vec![0, 1, 127, 255]),
|
||||
);
|
||||
assert!(normalize_imported_binary("sqlite", "payload_gzip", &json!([256])).is_err());
|
||||
assert_eq!(
|
||||
postgres_bytea_json_value("payload_gzip", &json!([0, 1, 127, 255]))
|
||||
.expect("postgres bytea should normalize"),
|
||||
json!("\\x00017fff"),
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn postgres_import_payload_rejects_non_null_unknown_columns() {
|
||||
let target_columns = BTreeMap::from([(
|
||||
@@ -197,6 +338,95 @@ fn postgres_import_payload_rejects_non_null_unknown_columns() {
|
||||
assert!(err.to_string().contains("does not exist"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mysql_and_sqlite_import_payloads_reject_non_null_unknown_columns() {
|
||||
let target_columns = BTreeSet::from(["id".to_string()]);
|
||||
let row = ExportRow {
|
||||
id: "user-1".to_string(),
|
||||
payload: json!({
|
||||
"id": "user-1",
|
||||
"legacy_nullable": null,
|
||||
"unexpected_column": "value"
|
||||
}),
|
||||
};
|
||||
|
||||
for driver_name in ["mysql", "sqlite"] {
|
||||
let err = filter_import_payload(
|
||||
driver_name,
|
||||
"users",
|
||||
ExportDomain::Users,
|
||||
&row,
|
||||
&target_columns,
|
||||
)
|
||||
.expect_err("non-null unknown columns should fail");
|
||||
|
||||
assert!(err.to_string().contains("unexpected_column"));
|
||||
assert!(err.to_string().contains("does not exist"));
|
||||
assert!(err.to_string().contains(driver_name));
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn mysql_and_sqlite_import_payloads_ignore_unknown_null_columns() {
|
||||
let target_columns = BTreeSet::from(["id".to_string()]);
|
||||
let row = ExportRow {
|
||||
id: "user-1".to_string(),
|
||||
payload: json!({
|
||||
"id": "user-1",
|
||||
"legacy_nullable": null
|
||||
}),
|
||||
};
|
||||
|
||||
let filtered = filter_import_payload(
|
||||
"sqlite",
|
||||
"users",
|
||||
ExportDomain::Users,
|
||||
&row,
|
||||
&target_columns,
|
||||
)
|
||||
.expect("unknown null columns should remain backward compatible");
|
||||
|
||||
assert_eq!(
|
||||
filtered,
|
||||
serde_json::Map::from_iter([("id".to_string(), json!("user-1"))])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn postgres_to_sqlite_copy_uses_primary_key_upsert_instead_of_replace() {
|
||||
let table = SchemaCopyTable {
|
||||
table_name: "usage".to_string(),
|
||||
columns: vec![
|
||||
SchemaCopyColumn {
|
||||
sqlite: SqliteCopyColumn {
|
||||
name: "request_id".to_string(),
|
||||
declared_type: "TEXT".to_string(),
|
||||
not_null: true,
|
||||
has_default: false,
|
||||
primary_key_position: 1,
|
||||
},
|
||||
postgres: postgres_column("character varying", "varchar"),
|
||||
},
|
||||
SchemaCopyColumn {
|
||||
sqlite: SqliteCopyColumn {
|
||||
name: "status".to_string(),
|
||||
declared_type: "TEXT".to_string(),
|
||||
not_null: true,
|
||||
has_default: false,
|
||||
primary_key_position: 0,
|
||||
},
|
||||
postgres: postgres_column("character varying", "varchar"),
|
||||
},
|
||||
],
|
||||
};
|
||||
|
||||
let sql = sqlite_schema_copy_insert_sql(&table).expect("copy SQL should build");
|
||||
|
||||
assert!(!sql.contains("OR REPLACE"));
|
||||
assert!(sql.contains("ON CONFLICT (\"request_id\") DO UPDATE SET"));
|
||||
assert!(sql.contains("\"status\" = excluded.\"status\""));
|
||||
}
|
||||
|
||||
fn postgres_column(data_type: &str, udt_name: &str) -> PostgresImportColumn {
|
||||
PostgresImportColumn {
|
||||
data_type: data_type.to_ascii_lowercase(),
|
||||
@@ -215,6 +445,183 @@ fn postgres_not_null_default_column(data_type: &str, udt_name: &str) -> Postgres
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn sqlite_import_rejects_non_integer_timestamp_values() {
|
||||
let pool = sqlx::sqlite::SqlitePoolOptions::new()
|
||||
.max_connections(1)
|
||||
.connect("sqlite::memory:")
|
||||
.await
|
||||
.expect("sqlite pool should connect");
|
||||
run_sqlite_migrations(&pool)
|
||||
.await
|
||||
.expect("sqlite migrations should run");
|
||||
|
||||
for invalid_value in [
|
||||
json!("not-a-timestamp"),
|
||||
json!(1.5),
|
||||
json!(true),
|
||||
json!({"unexpected": "object"}),
|
||||
] {
|
||||
let encoded = encode_jsonl(&[
|
||||
DataExportRecord::manifest(DataExportManifest::new(
|
||||
1_700_000_000,
|
||||
Some(DatabaseDriver::Postgres),
|
||||
vec![ExportDomain::GlobalModels],
|
||||
)),
|
||||
DataExportRecord::row(
|
||||
ExportDomain::GlobalModels,
|
||||
"invalid-timestamp",
|
||||
json!({
|
||||
"id": "invalid-timestamp",
|
||||
"name": "invalid-timestamp",
|
||||
"created_at": invalid_value,
|
||||
"updated_at": 1
|
||||
}),
|
||||
),
|
||||
])
|
||||
.expect("invalid timestamp fixture should encode");
|
||||
|
||||
let err = import_sqlite_jsonl(&pool, &encoded)
|
||||
.await
|
||||
.expect_err("non-integer timestamp should be rejected");
|
||||
assert!(err.to_string().contains(
|
||||
"timestamp column 'created_at' must contain an integer or supported datetime"
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn sqlite_import_updates_parent_without_cascading_child_rows() {
|
||||
let pool = sqlx::sqlite::SqlitePoolOptions::new()
|
||||
.max_connections(1)
|
||||
.connect("sqlite::memory:")
|
||||
.await
|
||||
.expect("sqlite pool should connect");
|
||||
run_sqlite_migrations(&pool)
|
||||
.await
|
||||
.expect("sqlite migrations should run");
|
||||
sqlx::query("PRAGMA foreign_keys = ON")
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("foreign keys should be enabled");
|
||||
sqlx::raw_sql(
|
||||
r#"
|
||||
INSERT INTO users (id, email, username, created_at, updated_at)
|
||||
VALUES ('import-user', '[email protected]', 'import-user', 1, 1);
|
||||
INSERT INTO user_groups (
|
||||
id, name, normalized_name, description, priority,
|
||||
allowed_providers_mode, allowed_api_formats_mode, allowed_models_mode, rate_limit_mode,
|
||||
created_at, updated_at
|
||||
)
|
||||
VALUES (
|
||||
'import-group', 'Before', 'import-group', 'preserve-me', 0,
|
||||
'inherit', 'inherit', 'inherit', 'inherit', 1, 1
|
||||
);
|
||||
INSERT INTO user_group_members (group_id, user_id, created_at)
|
||||
VALUES ('import-group', 'import-user', 1);
|
||||
"#,
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("parent and child fixtures should insert");
|
||||
|
||||
let encoded = encode_jsonl(&[
|
||||
DataExportRecord::manifest(DataExportManifest::new(
|
||||
1_700_000_000,
|
||||
Some(DatabaseDriver::Postgres),
|
||||
vec![ExportDomain::UserGroups],
|
||||
)),
|
||||
DataExportRecord::row(
|
||||
ExportDomain::UserGroups,
|
||||
"import-group",
|
||||
json!({
|
||||
"id": "import-group",
|
||||
"name": "After",
|
||||
"normalized_name": "import-group",
|
||||
"priority": 10,
|
||||
"allowed_providers_mode": "inherit",
|
||||
"allowed_api_formats_mode": "inherit",
|
||||
"allowed_models_mode": "inherit",
|
||||
"rate_limit_mode": "inherit",
|
||||
"created_at": 1,
|
||||
"updated_at": 2
|
||||
}),
|
||||
),
|
||||
])
|
||||
.expect("group export should encode");
|
||||
|
||||
assert_eq!(
|
||||
import_sqlite_jsonl(&pool, &encoded)
|
||||
.await
|
||||
.expect("group import should update in place"),
|
||||
1
|
||||
);
|
||||
let group = sqlx::query_as::<_, (String, String)>(
|
||||
"SELECT name, description FROM user_groups WHERE id = 'import-group'",
|
||||
)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("updated group should load");
|
||||
assert_eq!(group, ("After".to_string(), "preserve-me".to_string()));
|
||||
let member_count: i64 = sqlx::query_scalar(
|
||||
"SELECT COUNT(*) FROM user_group_members WHERE group_id = 'import-group'",
|
||||
)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("group member count should load");
|
||||
assert_eq!(member_count, 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn sqlite_import_rolls_back_rows_after_late_failure() {
|
||||
let pool = sqlx::sqlite::SqlitePoolOptions::new()
|
||||
.max_connections(1)
|
||||
.connect("sqlite::memory:")
|
||||
.await
|
||||
.expect("sqlite pool should connect");
|
||||
run_sqlite_migrations(&pool)
|
||||
.await
|
||||
.expect("sqlite migrations should run");
|
||||
let encoded = encode_jsonl(&[
|
||||
DataExportRecord::manifest(DataExportManifest::new(
|
||||
1_700_000_000,
|
||||
Some(DatabaseDriver::Postgres),
|
||||
vec![ExportDomain::GlobalModels],
|
||||
)),
|
||||
DataExportRecord::row(
|
||||
ExportDomain::GlobalModels,
|
||||
"rollback-valid",
|
||||
json!({
|
||||
"id": "rollback-valid",
|
||||
"name": "rollback-valid",
|
||||
"created_at": 1,
|
||||
"updated_at": 1
|
||||
}),
|
||||
),
|
||||
DataExportRecord::row(
|
||||
ExportDomain::GlobalModels,
|
||||
"rollback-invalid",
|
||||
json!({
|
||||
"id": "rollback-invalid",
|
||||
"name": "rollback-invalid",
|
||||
"created_at": "invalid-timestamp",
|
||||
"updated_at": 1
|
||||
}),
|
||||
),
|
||||
])
|
||||
.expect("rollback fixture should encode");
|
||||
|
||||
import_sqlite_jsonl(&pool, &encoded)
|
||||
.await
|
||||
.expect_err("late invalid row should fail the import");
|
||||
let count: i64 =
|
||||
sqlx::query_scalar("SELECT COUNT(*) FROM global_models WHERE id LIKE 'rollback-%'")
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("rolled back row count should load");
|
||||
assert_eq!(count, 0);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn sqlite_core_export_reads_migrated_database_rows() {
|
||||
let pool = sqlx::sqlite::SqlitePoolOptions::new()
|
||||
@@ -243,7 +650,7 @@ VALUES ('provider-key-1', 'provider-1', 'Provider Key', 'ciphertext-provider', '
|
||||
INSERT INTO provider_endpoints (id, provider_id, name, base_url, created_at, updated_at)
|
||||
VALUES ('endpoint-1', 'provider-1', 'Primary', 'https://example.test', '1970-01-01T00:00:01Z', '1970-01-01T00:00:02Z');
|
||||
INSERT INTO global_models (id, name, created_at, updated_at)
|
||||
VALUES ('global-model-1', 'gpt-test', '1970-01-01T00:00:01Z', '1970-01-01T00:00:02Z');
|
||||
VALUES ('global-model-1', 'gpt-test', '1970-01-01T00:00:01Z', '1970-01-01 00:00:02.123456');
|
||||
INSERT INTO models (id, provider_id, global_model_id, provider_model_name, created_at, updated_at)
|
||||
VALUES ('model-1', 'provider-1', 'global-model-1', 'gpt-test', '1970-01-01T00:00:01Z', '1970-01-01T00:00:02Z');
|
||||
INSERT INTO billing_rules (id, global_model_id, name, task_type, expression, variables, dimension_mappings, is_enabled, created_at, updated_at)
|
||||
@@ -255,7 +662,21 @@ VALUES ('config-1', 'billing.enabled', 'true', '1970-01-01T00:00:01Z', '1970-01-
|
||||
INSERT INTO wallets (id, user_id, created_at, updated_at)
|
||||
VALUES ('wallet-1', 'user-1', '1970-01-01T00:00:01Z', '1970-01-01T00:00:02Z');
|
||||
INSERT INTO "usage" (request_id, id, user_id, provider_name, model, status, billing_status, created_at_unix_ms, updated_at_unix_secs)
|
||||
VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'completed', 'settled', 1, 2);
|
||||
VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'completed', 'settled', '1970-01-01T00:00:01.234900Z', 2);
|
||||
INSERT INTO audit_logs (id, event_type, description, request_id, created_at)
|
||||
VALUES ('audit-1', 'request.completed', 'Exported audit', 'request-1', '1970-01-01T00:00:02Z');
|
||||
INSERT INTO usage_body_blobs (body_ref, request_id, body_field, payload_gzip, created_at, updated_at)
|
||||
VALUES ('body-ref-1', 'request-1', 'request', X'00117FFF', '1970-01-01T00:00:01Z', '1970-01-01T00:00:02Z');
|
||||
INSERT INTO usage_http_audits (request_id, request_body_ref, request_body_state, body_capture_mode, created_at, updated_at)
|
||||
VALUES ('request-1', 'body-ref-1', 'captured', 'full', '1970-01-01T00:00:01Z', '1970-01-01T00:00:02Z');
|
||||
INSERT INTO usage_routing_snapshots (
|
||||
request_id, candidate_id, candidate_index, selected_provider_id,
|
||||
selected_endpoint_id, selected_provider_api_key_id, created_at, updated_at
|
||||
)
|
||||
VALUES (
|
||||
'request-1', 'candidate-1', 2, 'provider-1',
|
||||
'endpoint-1', 'provider-key-1', '1970-01-01T00:00:01Z', '1970-01-01T00:00:02Z'
|
||||
);
|
||||
"#,
|
||||
)
|
||||
.execute(&pool)
|
||||
@@ -304,6 +725,21 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
|
||||
import_plan.rows(ExportDomain::Billing)[0].payload["dimension_mappings"]["input"],
|
||||
"input_tokens"
|
||||
);
|
||||
assert!(import_plan
|
||||
.rows(ExportDomain::Auxiliary)
|
||||
.iter()
|
||||
.any(|row| row.payload["__table"] == "audit_logs" && row.payload["id"] == "audit-1"));
|
||||
assert!(import_plan
|
||||
.rows(ExportDomain::Auxiliary)
|
||||
.iter()
|
||||
.any(|row| row.payload["__table"] == "usage_body_blobs"
|
||||
&& row.payload["payload_gzip"] == json!([0, 17, 127, 255])));
|
||||
assert!(import_plan
|
||||
.rows(ExportDomain::Auxiliary)
|
||||
.iter()
|
||||
.any(|row| row.payload["__table"] == "usage_routing_snapshots"
|
||||
&& row.payload["candidate_id"] == "candidate-1"
|
||||
&& row.payload["selected_provider_id"] == "provider-1"));
|
||||
|
||||
let target_pool = sqlx::sqlite::SqlitePoolOptions::new()
|
||||
.max_connections(1)
|
||||
@@ -316,7 +752,7 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
|
||||
let imported = import_sqlite_jsonl(&target_pool, &encoded)
|
||||
.await
|
||||
.expect("sqlite import should load exported rows");
|
||||
assert_eq!(imported, 16);
|
||||
assert_eq!(imported, 20);
|
||||
|
||||
let imported_api_key =
|
||||
sqlx::query_as::<_, (String,)>("SELECT key_encrypted FROM api_keys WHERE id = 'api-key-1'")
|
||||
@@ -325,13 +761,31 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
|
||||
.expect("imported api key should load");
|
||||
assert_eq!(imported_api_key.0, "ciphertext-1");
|
||||
|
||||
let imported_usage = sqlx::query_as::<_, (String,)>(
|
||||
"SELECT request_id FROM \"usage\" WHERE request_id = 'request-1'",
|
||||
let imported_usage = sqlx::query_as::<_, (String, i64, String)>(
|
||||
"SELECT request_id, created_at_unix_ms, typeof(created_at_unix_ms) FROM \"usage\" WHERE request_id = 'request-1'",
|
||||
)
|
||||
.fetch_one(&target_pool)
|
||||
.await
|
||||
.expect("imported usage should load");
|
||||
assert_eq!(imported_usage.0, "request-1");
|
||||
assert_eq!(
|
||||
imported_usage,
|
||||
("request-1".to_string(), 1, "integer".to_string())
|
||||
);
|
||||
|
||||
let imported_global_model_timestamps = sqlx::query_as::<_, (i64, i64, String, String)>(
|
||||
r#"
|
||||
SELECT created_at, updated_at, typeof(created_at), typeof(updated_at)
|
||||
FROM global_models
|
||||
WHERE id = 'global-model-1'
|
||||
"#,
|
||||
)
|
||||
.fetch_one(&target_pool)
|
||||
.await
|
||||
.expect("imported global model timestamps should decode as integers");
|
||||
assert_eq!(
|
||||
imported_global_model_timestamps,
|
||||
(1, 2, "integer".to_string(), "integer".to_string())
|
||||
);
|
||||
|
||||
let imported_group_member = sqlx::query_as::<_, (String, String)>(
|
||||
"SELECT group_id, user_id FROM user_group_members WHERE group_id = 'group-1' AND user_id = 'user-1'",
|
||||
@@ -350,6 +804,29 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
|
||||
.expect("imported billing rule should load");
|
||||
assert_eq!(imported_billing_rule.0, "input_tokens * 0.01");
|
||||
|
||||
let imported_body: Vec<u8> = sqlx::query_scalar(
|
||||
"SELECT payload_gzip FROM usage_body_blobs WHERE body_ref = 'body-ref-1'",
|
||||
)
|
||||
.fetch_one(&target_pool)
|
||||
.await
|
||||
.expect("imported body blob should load");
|
||||
assert_eq!(imported_body, vec![0, 17, 127, 255]);
|
||||
|
||||
let imported_routing = sqlx::query_as::<_, (String, i64, String)>(
|
||||
r#"
|
||||
SELECT candidate_id, candidate_index, selected_provider_id
|
||||
FROM usage_routing_snapshots
|
||||
WHERE request_id = 'request-1'
|
||||
"#,
|
||||
)
|
||||
.fetch_one(&target_pool)
|
||||
.await
|
||||
.expect("imported routing snapshot should load");
|
||||
assert_eq!(
|
||||
imported_routing,
|
||||
("candidate-1".to_string(), 2, "provider-1".to_string())
|
||||
);
|
||||
|
||||
if let Some(database_url) = std::env::var("AETHER_TEST_POSTGRES_URL")
|
||||
.ok()
|
||||
.filter(|value| !value.trim().is_empty())
|
||||
@@ -375,7 +852,7 @@ VALUES ('request-1', 'request-1', 'user-1', 'Provider One', 'gpt-test', 'complet
|
||||
let imported = import_postgres_jsonl(&postgres_pool, &encoded)
|
||||
.await
|
||||
.expect("postgres import should load exported rows");
|
||||
assert_eq!(imported, 16);
|
||||
assert_eq!(imported, 20);
|
||||
|
||||
let imported_api_key = sqlx::query_as::<_, (String,)>(
|
||||
"SELECT key_encrypted FROM api_keys WHERE id = 'api-key-1'",
|
||||
@@ -616,6 +1093,17 @@ async fn postgres_core_export_reads_migrated_database_rows_when_url_is_set() {
|
||||
.await
|
||||
.expect("imported sqlite api key should load");
|
||||
assert_eq!(imported_api_key.0, "ciphertext-1");
|
||||
let imported_global_model_timestamps = sqlx::query_as::<_, (i64, i64, String, String)>(
|
||||
"SELECT created_at, updated_at, typeof(created_at), typeof(updated_at) FROM global_models WHERE id = ?",
|
||||
)
|
||||
.bind(&global_model_id)
|
||||
.fetch_one(&target_pool)
|
||||
.await
|
||||
.expect("imported sqlite global model timestamps should decode as integers");
|
||||
assert_eq!(
|
||||
imported_global_model_timestamps,
|
||||
(1, 2, "integer".to_string(), "integer".to_string())
|
||||
);
|
||||
let imported_group_member = sqlx::query_as::<_, (String, String)>(
|
||||
"SELECT group_id, user_id FROM user_group_members WHERE group_id = ? AND user_id = ?",
|
||||
)
|
||||
|
||||
@@ -453,15 +453,173 @@ fn create_table_names(sql: &str) -> BTreeSet<String> {
|
||||
let trimmed = line.trim_start();
|
||||
let table_part = trimmed
|
||||
.strip_prefix("CREATE TABLE IF NOT EXISTS public.")
|
||||
.or_else(|| trimmed.strip_prefix("CREATE TABLE IF NOT EXISTS "))?;
|
||||
.or_else(|| trimmed.strip_prefix("CREATE TABLE IF NOT EXISTS "))
|
||||
.or_else(|| trimmed.strip_prefix("CREATE TABLE public."))
|
||||
.or_else(|| trimmed.strip_prefix("CREATE TABLE "))?;
|
||||
let table_name = table_part
|
||||
.split(|ch: char| ch.is_ascii_whitespace() || ch == '(')
|
||||
.next()?;
|
||||
Some(table_name.trim_matches('"').to_string())
|
||||
Some(
|
||||
table_name
|
||||
.trim_matches(|ch| ch == '"' || ch == '`')
|
||||
.to_string(),
|
||||
)
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn portable_driver_migrations_create_the_postgres_table_set() {
|
||||
let mut postgres_tables = POSTGRES_MIGRATOR
|
||||
.iter()
|
||||
.filter(|migration| migration.migration_type.is_up_migration())
|
||||
.flat_map(|migration| create_table_names(migration.sql.as_ref()))
|
||||
.collect::<BTreeSet<_>>();
|
||||
postgres_tables.remove("schema_backfills");
|
||||
|
||||
let mysql_tables = super::mysql::MIGRATOR
|
||||
.iter()
|
||||
.filter(|migration| migration.migration_type.is_up_migration())
|
||||
.flat_map(|migration| create_table_names(migration.sql.as_ref()))
|
||||
.collect::<BTreeSet<_>>();
|
||||
let sqlite_tables = super::sqlite::MIGRATOR
|
||||
.iter()
|
||||
.filter(|migration| migration.migration_type.is_up_migration())
|
||||
.flat_map(|migration| create_table_names(migration.sql.as_ref()))
|
||||
.collect::<BTreeSet<_>>();
|
||||
|
||||
assert_eq!(mysql_tables, postgres_tables, "MySQL table set drifted");
|
||||
assert_eq!(sqlite_tables, postgres_tables, "SQLite table set drifted");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn migrated_sqlite_columns_match_the_generated_logical_schema() {
|
||||
const GENERATED_SQLITE_SCHEMA: &[&str] = &[
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/schema/generated/sqlite/baseline/001_identity.sql"
|
||||
)),
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/schema/generated/sqlite/baseline/002_provider_catalog.sql"
|
||||
)),
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/schema/generated/sqlite/baseline/003_auth_config.sql"
|
||||
)),
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/schema/generated/sqlite/baseline/004_proxy_nodes.sql"
|
||||
)),
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/schema/generated/sqlite/baseline/005_wallet_billing.sql"
|
||||
)),
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/schema/generated/sqlite/baseline/006_usage.sql"
|
||||
)),
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/schema/generated/sqlite/baseline/007_stats.sql"
|
||||
)),
|
||||
include_str!(concat!(
|
||||
env!("CARGO_MANIFEST_DIR"),
|
||||
"/schema/generated/sqlite/baseline/008_background_tasks.sql"
|
||||
)),
|
||||
];
|
||||
|
||||
let migrated = sqlx::sqlite::SqlitePoolOptions::new()
|
||||
.max_connections(1)
|
||||
.connect("sqlite::memory:")
|
||||
.await
|
||||
.expect("migrated sqlite pool should connect");
|
||||
super::run_sqlite_migrations(&migrated)
|
||||
.await
|
||||
.expect("sqlite migrations should run");
|
||||
|
||||
let generated = sqlx::sqlite::SqlitePoolOptions::new()
|
||||
.max_connections(1)
|
||||
.connect("sqlite::memory:")
|
||||
.await
|
||||
.expect("generated sqlite pool should connect");
|
||||
for source in GENERATED_SQLITE_SCHEMA {
|
||||
sqlx::raw_sql(source)
|
||||
.execute(&generated)
|
||||
.await
|
||||
.expect("generated sqlite schema fragment should run");
|
||||
}
|
||||
|
||||
let migrated_tables = sqlite_portable_table_names(&migrated).await;
|
||||
let generated_tables = sqlite_portable_table_names(&generated).await;
|
||||
assert_eq!(migrated_tables, generated_tables);
|
||||
|
||||
for table in generated_tables {
|
||||
let migrated_columns = sqlite_table_column_names(&migrated, &table).await;
|
||||
let generated_columns = sqlite_table_column_names(&generated, &table).await;
|
||||
assert_eq!(
|
||||
migrated_columns, generated_columns,
|
||||
"SQLite migration columns drifted for table {table}"
|
||||
);
|
||||
}
|
||||
|
||||
let migrated_indexes = sqlite_named_index_names(&migrated).await;
|
||||
let generated_indexes = sqlite_named_index_names(&generated).await;
|
||||
let missing_indexes = generated_indexes
|
||||
.difference(&migrated_indexes)
|
||||
.cloned()
|
||||
.collect::<BTreeSet<_>>();
|
||||
assert!(
|
||||
missing_indexes.is_empty(),
|
||||
"SQLite migrations are missing generated logical indexes: {missing_indexes:?}"
|
||||
);
|
||||
}
|
||||
|
||||
async fn sqlite_portable_table_names(pool: &SqlitePool) -> BTreeSet<String> {
|
||||
query_scalar::<_, String>(
|
||||
r#"
|
||||
SELECT name
|
||||
FROM sqlite_master
|
||||
WHERE type = 'table'
|
||||
AND name NOT LIKE 'sqlite_%'
|
||||
AND name NOT IN ('_sqlx_migrations', 'schema_backfills')
|
||||
ORDER BY name
|
||||
"#,
|
||||
)
|
||||
.fetch_all(pool)
|
||||
.await
|
||||
.expect("sqlite table names should load")
|
||||
.into_iter()
|
||||
.collect()
|
||||
}
|
||||
|
||||
async fn sqlite_table_column_names(pool: &SqlitePool, table: &str) -> BTreeSet<String> {
|
||||
query_scalar::<_, String>("SELECT name FROM pragma_table_info(?) ORDER BY cid")
|
||||
.bind(table)
|
||||
.fetch_all(pool)
|
||||
.await
|
||||
.expect("sqlite table columns should load")
|
||||
.into_iter()
|
||||
.collect()
|
||||
}
|
||||
|
||||
async fn sqlite_named_index_names(pool: &SqlitePool) -> BTreeSet<String> {
|
||||
query_scalar::<_, String>(
|
||||
r#"
|
||||
SELECT name
|
||||
FROM sqlite_master
|
||||
WHERE type = 'index'
|
||||
AND sql IS NOT NULL
|
||||
ORDER BY name
|
||||
"#,
|
||||
)
|
||||
.fetch_all(pool)
|
||||
.await
|
||||
.expect("sqlite named indexes should load")
|
||||
.into_iter()
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn empty_database_snapshot_sql_includes_usage_body_blobs_and_audit_admin_role() {
|
||||
assert!(EMPTY_DATABASE_SNAPSHOT_SQL.contains("'audit_admin'"));
|
||||
@@ -862,6 +1020,9 @@ fn mysql_and_sqlite_migrations_include_enabled_incrementals() {
|
||||
20260527000000,
|
||||
20260528000000,
|
||||
20260528020000,
|
||||
20260725010000,
|
||||
20260725020000,
|
||||
20260725030000,
|
||||
]
|
||||
);
|
||||
assert_eq!(
|
||||
@@ -889,10 +1050,304 @@ fn mysql_and_sqlite_migrations_include_enabled_incrementals() {
|
||||
20260527000000,
|
||||
20260528000000,
|
||||
20260528020000,
|
||||
20260725000000,
|
||||
20260725010000,
|
||||
20260725020000,
|
||||
20260725030000,
|
||||
20260725040000,
|
||||
]
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn sqlite_imported_timestamp_migration_normalizes_text_storage() {
|
||||
let pool = SqlitePool::connect("sqlite::memory:")
|
||||
.await
|
||||
.expect("sqlite pool should connect");
|
||||
super::run_sqlite_migrations(&pool)
|
||||
.await
|
||||
.expect("sqlite migrations should run");
|
||||
|
||||
query(
|
||||
r#"
|
||||
INSERT INTO global_models (id, name, created_at, updated_at)
|
||||
VALUES
|
||||
('timestamp-rfc3339', 'timestamp-rfc3339', '1970-01-01T00:00:01Z', '1970-01-01T08:00:02+08:00'),
|
||||
('timestamp-sqlalchemy', 'timestamp-sqlalchemy', '1970-01-01 00:00:03.123456', '1970-01-01 00:00:04.987654'),
|
||||
('timestamp-integer', 'timestamp-integer', 5, 6);
|
||||
"#,
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("timestamp fixtures should insert");
|
||||
query(
|
||||
r#"
|
||||
INSERT INTO "usage" (request_id, created_at_unix_ms, updated_at_unix_secs)
|
||||
VALUES ('timestamp-usage', '1970-01-01T00:00:01.234900Z', '1970-01-01T00:00:02Z');
|
||||
"#,
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("usage timestamp fixture should insert");
|
||||
|
||||
let migration = super::sqlite::MIGRATOR
|
||||
.iter()
|
||||
.find(|migration| migration.version == 20260725000000)
|
||||
.expect("timestamp normalization migration should be embedded");
|
||||
sqlx::raw_sql(migration.sql.as_ref())
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("timestamp normalization migration should apply");
|
||||
|
||||
let rows = sqlx::query_as::<_, (String, i64, i64, String, String)>(
|
||||
r#"
|
||||
SELECT id, created_at, updated_at, typeof(created_at), typeof(updated_at)
|
||||
FROM global_models
|
||||
WHERE id LIKE 'timestamp-%'
|
||||
ORDER BY id
|
||||
"#,
|
||||
)
|
||||
.fetch_all(&pool)
|
||||
.await
|
||||
.expect("normalized timestamps should decode as integers");
|
||||
|
||||
assert_eq!(
|
||||
rows,
|
||||
vec![
|
||||
(
|
||||
"timestamp-integer".to_string(),
|
||||
5,
|
||||
6,
|
||||
"integer".to_string(),
|
||||
"integer".to_string(),
|
||||
),
|
||||
(
|
||||
"timestamp-rfc3339".to_string(),
|
||||
1,
|
||||
2,
|
||||
"integer".to_string(),
|
||||
"integer".to_string(),
|
||||
),
|
||||
(
|
||||
"timestamp-sqlalchemy".to_string(),
|
||||
3,
|
||||
4,
|
||||
"integer".to_string(),
|
||||
"integer".to_string(),
|
||||
),
|
||||
]
|
||||
);
|
||||
|
||||
let usage_timestamps = sqlx::query_as::<_, (i64, i64, String, String)>(
|
||||
r#"
|
||||
SELECT created_at_unix_ms, updated_at_unix_secs,
|
||||
typeof(created_at_unix_ms), typeof(updated_at_unix_secs)
|
||||
FROM "usage"
|
||||
WHERE request_id = 'timestamp-usage'
|
||||
"#,
|
||||
)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("normalized usage timestamps should decode as integers");
|
||||
assert_eq!(
|
||||
usage_timestamps,
|
||||
(1, 2, "integer".to_string(), "integer".to_string())
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn sqlite_imported_timestamp_migration_rejects_non_integer_storage() {
|
||||
let pool = SqlitePool::connect("sqlite::memory:")
|
||||
.await
|
||||
.expect("sqlite pool should connect");
|
||||
super::run_sqlite_migrations(&pool)
|
||||
.await
|
||||
.expect("sqlite migrations should run");
|
||||
query(
|
||||
r#"
|
||||
INSERT INTO global_models (id, name, created_at, updated_at)
|
||||
VALUES ('timestamp-invalid', 'timestamp-invalid', 1.5, 1);
|
||||
"#,
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("non-integer timestamp fixture should insert");
|
||||
|
||||
let migration = super::sqlite::MIGRATOR
|
||||
.iter()
|
||||
.find(|migration| migration.version == 20260725000000)
|
||||
.expect("timestamp normalization migration should be embedded");
|
||||
let err = sqlx::raw_sql(migration.sql.as_ref())
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect_err("non-integer timestamp should fail the migration");
|
||||
assert!(err
|
||||
.to_string()
|
||||
.contains("imported_timestamp_storage_must_be_integer"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn sqlite_remaining_timestamp_migration_repairs_other_repository_domains() {
|
||||
let pool = SqlitePool::connect("sqlite::memory:")
|
||||
.await
|
||||
.expect("sqlite pool should connect");
|
||||
super::run_sqlite_migrations(&pool)
|
||||
.await
|
||||
.expect("sqlite migrations should run");
|
||||
|
||||
query(
|
||||
r#"
|
||||
INSERT INTO users (id, email, username, auth_source, created_at, updated_at)
|
||||
VALUES ('timestamp-user', 'timestamp@example.com', 'timestamp-user', 'local', 1, 1);
|
||||
|
||||
INSERT INTO audit_logs (id, event_type, description, created_at)
|
||||
VALUES ('timestamp-audit', 'test', 'test', '1970-01-01T00:00:01Z');
|
||||
|
||||
INSERT INTO request_candidates (
|
||||
id, request_id, candidate_index, status, created_at, started_at, finished_at
|
||||
) VALUES (
|
||||
'timestamp-candidate', 'timestamp-request', 0, 'success',
|
||||
'1970-01-01T00:00:02Z', '1970-01-01T00:00:03Z', '1970-01-01T00:00:04Z'
|
||||
);
|
||||
|
||||
INSERT INTO stats_daily (id, date, created_at, updated_at)
|
||||
VALUES (
|
||||
'timestamp-stats', '1970-01-02',
|
||||
'1970-01-01T00:00:05Z', '1970-01-01T00:00:06Z'
|
||||
);
|
||||
|
||||
INSERT INTO user_sessions (
|
||||
id, user_id, client_device_id, refresh_token_hash,
|
||||
last_seen_at, expires_at, created_at, updated_at
|
||||
) VALUES (
|
||||
'timestamp-session', 'timestamp-user', 'device', 'hash',
|
||||
'1970-01-01T00:00:07Z', '1970-01-01T00:00:08Z',
|
||||
'1970-01-01T00:00:09Z', '1970-01-01T00:00:10Z'
|
||||
);
|
||||
"#,
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("remaining timestamp fixtures should insert");
|
||||
|
||||
let migration = super::sqlite::MIGRATOR
|
||||
.iter()
|
||||
.find(|migration| migration.version == 20260725040000)
|
||||
.expect("remaining timestamp migration should be embedded");
|
||||
sqlx::raw_sql(migration.sql.as_ref())
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("remaining timestamp migration should apply");
|
||||
|
||||
let audit = sqlx::query_as::<_, (i64, String)>(
|
||||
"SELECT created_at, typeof(created_at) FROM audit_logs WHERE id = 'timestamp-audit'",
|
||||
)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("normalized audit timestamp should load");
|
||||
assert_eq!(audit, (1, "integer".to_string()));
|
||||
|
||||
let candidate = sqlx::query_as::<_, (i64, i64, i64, String, String, String)>(
|
||||
r#"
|
||||
SELECT created_at, started_at, finished_at,
|
||||
typeof(created_at), typeof(started_at), typeof(finished_at)
|
||||
FROM request_candidates
|
||||
WHERE id = 'timestamp-candidate'
|
||||
"#,
|
||||
)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("normalized candidate timestamps should load");
|
||||
assert_eq!(
|
||||
candidate,
|
||||
(
|
||||
2,
|
||||
3,
|
||||
4,
|
||||
"integer".to_string(),
|
||||
"integer".to_string(),
|
||||
"integer".to_string(),
|
||||
)
|
||||
);
|
||||
|
||||
let stats = sqlx::query_as::<_, (i64, i64, i64, String, String, String)>(
|
||||
r#"
|
||||
SELECT date, created_at, updated_at,
|
||||
typeof(date), typeof(created_at), typeof(updated_at)
|
||||
FROM stats_daily
|
||||
WHERE id = 'timestamp-stats'
|
||||
"#,
|
||||
)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("normalized stats timestamps should load");
|
||||
assert_eq!(
|
||||
stats,
|
||||
(
|
||||
86_400,
|
||||
5,
|
||||
6,
|
||||
"integer".to_string(),
|
||||
"integer".to_string(),
|
||||
"integer".to_string(),
|
||||
)
|
||||
);
|
||||
|
||||
let session = sqlx::query_as::<_, (i64, i64, i64, i64, String, String, String, String)>(
|
||||
r#"
|
||||
SELECT last_seen_at, expires_at, created_at, updated_at,
|
||||
typeof(last_seen_at), typeof(expires_at), typeof(created_at), typeof(updated_at)
|
||||
FROM user_sessions
|
||||
WHERE id = 'timestamp-session'
|
||||
"#,
|
||||
)
|
||||
.fetch_one(&pool)
|
||||
.await
|
||||
.expect("normalized session timestamps should load");
|
||||
assert_eq!(
|
||||
session,
|
||||
(
|
||||
7,
|
||||
8,
|
||||
9,
|
||||
10,
|
||||
"integer".to_string(),
|
||||
"integer".to_string(),
|
||||
"integer".to_string(),
|
||||
"integer".to_string(),
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn sqlite_remaining_timestamp_migration_rejects_invalid_storage() {
|
||||
let pool = SqlitePool::connect("sqlite::memory:")
|
||||
.await
|
||||
.expect("sqlite pool should connect");
|
||||
super::run_sqlite_migrations(&pool)
|
||||
.await
|
||||
.expect("sqlite migrations should run");
|
||||
query(
|
||||
r#"
|
||||
INSERT INTO audit_logs (id, event_type, description, created_at)
|
||||
VALUES ('timestamp-invalid-audit', 'test', 'test', 1.5);
|
||||
"#,
|
||||
)
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect("invalid timestamp fixture should insert");
|
||||
|
||||
let migration = super::sqlite::MIGRATOR
|
||||
.iter()
|
||||
.find(|migration| migration.version == 20260725040000)
|
||||
.expect("remaining timestamp migration should be embedded");
|
||||
let err = sqlx::raw_sql(migration.sql.as_ref())
|
||||
.execute(&pool)
|
||||
.await
|
||||
.expect_err("invalid remaining timestamp should fail the migration");
|
||||
assert!(err.to_string().contains("invalid_count = 0"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn endpoint_api_root_migration_moves_v1_from_stored_default_paths() {
|
||||
let pool = SqlitePool::connect("sqlite::memory:")
|
||||
|
||||
Reference in New Issue
Block a user