use crate::db::Database;
use crate::sql::{self, Dialect, Value};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ColumnKind {
Uuid,
Text,
Blob,
I64,
Bool,
}
#[derive(Debug, Clone, Copy)]
pub struct TableSpec {
pub name: &'static str,
pub key: &'static [&'static str],
pub columns: &'static [(&'static str, ColumnKind)],
}
use ColumnKind::{Blob, Bool, I64, Text, Uuid};
pub const TABLES: &[TableSpec] = &[
TableSpec {
name: "nonces",
key: &["value"],
columns: &[("value", Text), ("created_at", I64)],
},
TableSpec {
name: "eab_keys",
key: &["kid"],
columns: &[
("kid", Uuid),
("secret", Blob),
("label", Text),
("profile", Text),
("status", Text),
("created_at", I64),
],
},
TableSpec {
name: "jobs",
key: &["id"],
columns: &[
("id", Uuid),
("kind", Text),
("dedup_key", Text),
("payload", Text),
("status", Text),
("run_at", I64),
("attempts", I64),
("max_attempts", I64),
("deadline", I64),
("lease_until", I64),
("lease_owner", Text),
("last_error", Text),
("created_at", I64),
("updated_at", I64),
],
},
TableSpec {
name: "audit_log",
key: &["id"],
columns: &[
("id", I64),
("created_at", I64),
("event", Text),
("outcome", Text),
("profile", Text),
("actor_kind", Text),
("actor_id", Text),
("account_id", Text),
("order_id", Text),
("cert_serial", Text),
("identifiers", Text),
("client_ip", Text),
("client_ptr", Text),
("user_agent", Text),
("request_id", Text),
("reason", Text),
("detail", Text),
],
},
TableSpec {
name: "revocations",
key: &["issuer", "serial"],
columns: &[
("issuer", Text),
("serial", Text),
("revoked_at", I64),
("reason", I64),
("not_after", I64),
],
},
TableSpec {
name: "crls",
key: &["issuer"],
columns: &[
("issuer", Text),
("crl_number", I64),
("der", Blob),
("this_update", I64),
("next_update", I64),
],
},
TableSpec {
name: "http01_tokens",
key: &["token"],
columns: &[
("token", Text),
("key_authorization", Text),
("created_at", I64),
("expires_at", I64),
],
},
TableSpec {
name: "admin_users",
key: &["id"],
columns: &[
("id", Uuid),
("username", Text),
("password_hash", Text),
("status", Text),
("totp_secret", Blob),
("totp_pending_secret", Blob),
("totp_last_step", I64),
("created_at", I64),
("updated_at", I64),
("last_login_at", I64),
("role", Text),
("contact_email", Text),
("known_login_ips", Text),
],
},
TableSpec {
name: "accounts",
key: &["id"],
columns: &[
("id", Uuid),
("profile", Text),
("pubkey", Blob),
("contact", Text),
("status", Text),
("created_at", I64),
("created_ip", Text),
("created_ptr", Text),
("last_seen_at", I64),
("last_seen_ip", Text),
("last_seen_ptr", Text),
("eab_kid", Uuid),
("terms_of_service_agreed", Bool),
],
},
TableSpec {
name: "orders",
key: &["id"],
columns: &[
("id", Uuid),
("profile", Text),
("account_id", Uuid),
("status", Text),
("identifiers", Text),
("expires", I64),
("not_before", I64),
("not_after", I64),
("error", Text),
("certificate", Text),
("replaces", Text),
("created_at", I64),
("created_ip", Text),
("created_ptr", Text),
("cert_serial", Text),
("cert_pubkey", Blob),
("revoked_at", I64),
("revocation_reason", I64),
("cert_not_after", I64),
],
},
TableSpec {
name: "authorizations",
key: &["id"],
columns: &[
("id", Uuid),
("order_id", Uuid),
("identifier", Text),
("status", Text),
("expires", I64),
("created_at", I64),
],
},
TableSpec {
name: "challenges",
key: &["id"],
columns: &[
("id", Uuid),
("authz_id", Uuid),
("type", Text),
("token", Text),
("status", Text),
("validated", I64),
("created_at", I64),
("error", Text),
],
},
TableSpec {
name: "upstream_orders",
key: &["order_id"],
columns: &[
("order_id", Uuid),
("upstream_order_url", Text),
("upstream_finalize_url", Text),
("upstream_certificate_url", Text),
("csr_der", Blob),
("status", Text),
("error", Text),
("created_at", I64),
("updated_at", I64),
("client_ip", Text),
("client_ptr", Text),
("user_agent", Text),
("request_id", Text),
],
},
TableSpec {
name: "admin_sessions",
key: &["token_hash"],
columns: &[
("token_hash", Text),
("user_id", Uuid),
("csrf_token", Text),
("state", Text),
("mfa_attempts", I64),
("created_at", I64),
("expires_at", I64),
("last_seen_at", I64),
("created_ip", Text),
("user_agent", Text),
],
},
TableSpec {
name: "admin_recovery_codes",
key: &["id"],
columns: &[
("id", Uuid),
("user_id", Uuid),
("code_hash", Text),
("created_at", I64),
("used_at", I64),
],
},
];
const BATCH: i64 = 1000;
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TableCount {
pub table: &'static str,
pub rows: u64,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct TransferReport {
pub tables: Vec<TableCount>,
}
impl TransferReport {
#[must_use]
pub fn total(&self) -> u64 {
self.tables.iter().map(|table| table.rows).sum()
}
}
pub async fn non_empty_tables(database: &Database) -> Result<Vec<TableCount>, sqlx::Error> {
let mut found = Vec::new();
for table in TABLES {
let rows = count(table.name, database).await?;
if rows > 0 {
found.push(TableCount {
table: table.name,
rows,
});
}
}
Ok(found)
}
async fn count(table: &'static str, database: &Database) -> Result<u64, sqlx::Error> {
let sql = format!("SELECT COUNT(*) FROM {table};");
let count: i64 = sql::query(sqlx::AssertSqlSafe(sql))
.fetch_one(database)
.await?
.try_get(0usize)?;
Ok(u64::try_from(count).unwrap_or(0))
}
impl Database {
pub async fn transfer_to(&self, target: &Database) -> Result<TransferReport, sqlx::Error> {
let mut tx = target.write_transaction().await?;
let mut tables = Vec::with_capacity(TABLES.len());
for spec in TABLES {
let rows = copy_table(spec, self, &mut tx).await?;
tables.push(TableCount {
table: spec.name,
rows,
});
}
if target.dialect() == Dialect::Postgres {
sql::query(
"SELECT setval(pg_get_serial_sequence('audit_log', 'id'), \
(SELECT COALESCE(MAX(id), 1) FROM audit_log));",
)
.fetch_one(tx.conn())
.await?;
}
tx.commit().await?;
Ok(TransferReport { tables })
}
}
async fn copy_table(
spec: &TableSpec,
source: &Database,
tx: &mut crate::db::Tx,
) -> Result<u64, sqlx::Error> {
let names = spec
.columns
.iter()
.map(|(name, _)| quote(name))
.collect::<Vec<_>>()
.join(", ");
let order = spec
.key
.iter()
.map(|name| quote(name))
.collect::<Vec<_>>()
.join(", ");
let mut after: Option<Vec<Value>> = None;
let mut copied = 0u64;
loop {
let batch = read_batch(spec, source, &names, &order, after.as_deref()).await?;
if batch.is_empty() {
return Ok(copied);
}
after = Some(key_of(spec, batch.last().expect("the batch is not empty")));
copied += batch.len() as u64;
write_batch(spec, &names, &batch, tx).await?;
if batch.len() < usize::try_from(BATCH).unwrap_or(usize::MAX) {
return Ok(copied);
}
}
}
async fn read_batch(
spec: &TableSpec,
source: &Database,
names: &str,
order: &str,
after: Option<&[Value]>,
) -> Result<Vec<Vec<Value>>, sqlx::Error> {
let table = spec.name;
let where_clause = match after {
None => String::new(),
Some(_) => {
let markers = vec!["?"; spec.key.len()].join(", ");
match spec.key.len() {
1 => format!(" WHERE {order} > {markers}"),
_ => format!(" WHERE ({order}) > ({markers})"),
}
}
};
let sql = format!("SELECT {names} FROM {table}{where_clause} ORDER BY {order} LIMIT {BATCH};");
let mut query = sql::query(sqlx::AssertSqlSafe(sql));
for value in after.unwrap_or(&[]) {
query = query.bind(value.clone());
}
let rows = query.fetch_all(source).await?;
rows.iter()
.map(|row| {
spec.columns
.iter()
.map(|(name, kind)| read(row, name, *kind))
.collect::<Result<Vec<_>, _>>()
})
.collect()
}
async fn write_batch(
spec: &TableSpec,
names: &str,
batch: &[Vec<Value>],
tx: &mut crate::db::Tx,
) -> Result<(), sqlx::Error> {
let table = spec.name;
let row = format!("({})", vec!["?"; spec.columns.len()].join(", "));
let values = vec![row; batch.len()].join(", ");
let overriding = match (tx.conn().dialect(), table) {
(Dialect::Postgres, "audit_log") => " OVERRIDING SYSTEM VALUE",
_ => "",
};
let sql = format!("INSERT INTO {table} ({names}){overriding} VALUES {values};");
let mut query = sql::query(sqlx::AssertSqlSafe(sql));
for row in batch {
for value in row {
query = query.bind(value.clone());
}
}
query.execute(tx.conn()).await?;
Ok(())
}
fn key_of(spec: &TableSpec, row: &[Value]) -> Vec<Value> {
spec.key
.iter()
.map(|key| {
let index = spec
.columns
.iter()
.position(|(name, _)| name == key)
.expect("a key column is one of the table's columns");
row[index].clone()
})
.collect()
}
fn read(row: &sql::Row, name: &str, kind: ColumnKind) -> Result<Value, sqlx::Error> {
Ok(match kind {
ColumnKind::Uuid => match row.try_get::<Option<uuid::Uuid>>(name)? {
Some(value) => Value::Uuid(value),
None => Value::Null(sql::NullKind::Uuid),
},
ColumnKind::Text => match row.try_get::<Option<String>>(name)? {
Some(value) => Value::Text(value),
None => Value::Null(sql::NullKind::Text),
},
ColumnKind::Blob => match row.try_get::<Option<Vec<u8>>>(name)? {
Some(value) => Value::Blob(value),
None => Value::Null(sql::NullKind::Blob),
},
ColumnKind::I64 => match row.try_get::<Option<i64>>(name)? {
Some(value) => Value::I64(value),
None => Value::Null(sql::NullKind::I64),
},
ColumnKind::Bool => match row.try_get::<Option<bool>>(name)? {
Some(value) => Value::Bool(value),
None => Value::Null(sql::NullKind::Bool),
},
})
}
fn quote(name: &str) -> String {
format!("\"{name}\"")
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
#[tokio::test]
async fn the_manifest_names_every_column() {
let database = Database::connect_for_test().await.unwrap();
for spec in TABLES {
let declared: Vec<&str> = spec.columns.iter().map(|(name, _)| *name).collect();
let live = live_columns(&database, spec.name).await;
assert_eq!(
live, declared,
"the manifest for `{}` has drifted from the schema; a column \
missing here is a column the transfer would drop",
spec.name
);
}
}
#[tokio::test]
async fn the_manifest_names_every_table() {
let database = Database::connect_for_test().await.unwrap();
let mut live = live_tables(&database).await;
let mut declared: Vec<String> = TABLES.iter().map(|s| s.name.to_string()).collect();
live.sort();
declared.sort();
assert_eq!(
live, declared,
"every table but `_sqlx_migrations` is copied; one missing here is \
one the transfer would leave behind"
);
}
#[test]
fn the_manifest_is_in_dependency_order() {
const EDGES: &[(&str, &str)] = &[
("orders", "accounts"),
("authorizations", "orders"),
("challenges", "authorizations"),
("upstream_orders", "orders"),
("admin_sessions", "admin_users"),
("admin_recovery_codes", "admin_users"),
];
let position = |name: &str| {
TABLES
.iter()
.position(|spec| spec.name == name)
.unwrap_or_else(|| panic!("{name} should be in the manifest"))
};
for (child, parent) in EDGES {
assert!(
position(parent) < position(child),
"{parent} must be copied before {child}, or the foreign key refuses the row"
);
}
}
#[test]
fn every_key_is_a_column_of_its_table() {
for spec in TABLES {
assert!(!spec.key.is_empty(), "{} declares no key", spec.name);
for key in spec.key {
assert!(
spec.columns.iter().any(|(name, _)| name == key),
"{}.{key} is a key but not a column",
spec.name
);
}
}
}
#[tokio::test]
async fn a_copy_carries_every_table_and_leaves_the_source_alone() {
let source = Arc::new(Database::connect_for_test().await.unwrap());
let target = Database::connect_for_test().await.unwrap();
crate::testutil::seed_every_table(&source).await;
let before = crate::testutil::row_counts(&source).await;
assert!(
before.iter().all(|(_, rows)| *rows > 0),
"every table must be seeded or this proves nothing: {before:?}"
);
assert!(
non_empty_tables(&target)
.await
.expect("an empty target counts")
.is_empty(),
"the target starts empty"
);
let report = source
.transfer_to(&target)
.await
.expect("the copy should succeed");
assert_eq!(
report.total(),
before.iter().map(|(_, rows)| rows).sum::<u64>(),
"the report counts what the tables hold"
);
assert_eq!(
report.tables.len(),
TABLES.len(),
"every table is reported, including any that were empty"
);
assert_eq!(
crate::testutil::row_counts(&target).await,
before,
"the target holds what the source held"
);
assert_eq!(
crate::testutil::row_counts(&source).await,
before,
"the source is only read"
);
let agreed = crate::sql::query(
"SELECT terms_of_service_agreed FROM accounts ORDER BY created_at, id;",
)
.fetch_all(&target)
.await
.expect("the accounts are readable")
.iter()
.map(|row| {
row.try_get::<Option<bool>>(0usize)
.expect("a nullable bool")
})
.collect::<Vec<_>>();
assert!(
agreed.contains(&Some(true)) && agreed.contains(&None),
"an agreement and its absence both survive: {agreed:?}"
);
}
#[tokio::test]
async fn a_table_longer_than_one_batch_is_copied_whole() {
let source = Arc::new(Database::connect_for_test().await.unwrap());
let target = Database::connect_for_test().await.unwrap();
let rows = u64::try_from(BATCH).expect("the batch size is positive") + 1;
for index in 0..rows {
crate::nonce::Nonce::new()
.save(&source)
.await
.expect("a nonce is storable");
crate::revocation::Revocation {
issuer: "a".repeat(64),
serial: format!("{index:08x}"),
revoked_at: 1,
reason: None,
not_after: None,
}
.insert_if_absent(&source)
.await
.expect("a revocation is storable");
}
let report = source
.transfer_to(&target)
.await
.expect("the copy should succeed");
for table in ["nonces", "revocations"] {
let reported = report
.tables
.iter()
.find(|entry| entry.table == table)
.unwrap_or_else(|| panic!("{table} is in the manifest"));
assert_eq!(reported.rows, rows, "{table}: one batch and one row");
assert_eq!(
count(table, &target).await.expect("a count"),
rows,
"{table}: and the target agrees"
);
}
}
#[tokio::test]
async fn a_failed_copy_leaves_the_target_as_it_was() {
let source = Arc::new(Database::connect_for_test().await.unwrap());
let target = Arc::new(Database::connect_for_test().await.unwrap());
crate::testutil::seed_every_table(&source).await;
crate::admin_user::AdminUser::create("alice", "hash", None, &target)
.await
.expect("the target has an operator of its own");
let error = source
.transfer_to(&target)
.await
.expect_err("the duplicate username must fail the copy");
assert!(
crate::sql::is_unique_violation(&error),
"the cause is the username, not something else: {error}"
);
for (table, rows) in crate::testutil::row_counts(&target).await {
let expected = u64::from(table == "admin_users");
assert_eq!(
rows, expected,
"`{table}` should hold {expected} row(s) after the rollback"
);
}
}
async fn live_tables(database: &Database) -> Vec<String> {
let sql = match database.dialect() {
Dialect::Sqlite => {
"SELECT name FROM sqlite_master WHERE type = 'table' \
AND name NOT LIKE 'sqlite_%' AND name <> '_sqlx_migrations';"
}
Dialect::Postgres => {
"SELECT table_name FROM information_schema.tables \
WHERE table_schema = ANY (current_schemas(false)) \
AND table_type = 'BASE TABLE' AND table_name <> '_sqlx_migrations';"
}
};
sql::query(sql)
.fetch_all(database)
.await
.expect("the catalog should be readable")
.iter()
.map(|row| row.try_get::<String>(0usize).expect("a name"))
.collect()
}
async fn live_columns(database: &Database, table: &str) -> Vec<String> {
match database.dialect() {
Dialect::Sqlite => sql::query(sqlx::AssertSqlSafe(format!(
"SELECT name FROM pragma_table_info('{table}');"
)))
.fetch_all(database)
.await
.expect("the table should exist")
.iter()
.map(|row| row.try_get::<String>(0usize).expect("a name"))
.collect(),
Dialect::Postgres => sql::query(
"SELECT column_name FROM information_schema.columns \
WHERE table_schema = ANY (current_schemas(false)) AND table_name = ? \
ORDER BY ordinal_position;",
)
.bind(table)
.fetch_all(database)
.await
.expect("the table should exist")
.iter()
.map(|row| row.try_get::<String>(0usize).expect("a name"))
.collect(),
}
}
}