use sqlx::{Executor, Row};
use super::dialect::{SessionDb, SessionPool, sql};
use super::error::SessionStoreError;
type Conn = <SessionDb as sqlx::Database>::Connection;
const STATEMENT_SEPARATOR: &str = "--;;";
const MIGRATIONS: &[(&str, &str)] = &[("0001_sessions", sql::SCHEMA)];
pub(super) async fn apply(pool: &SessionPool) -> Result<(), SessionStoreError> {
let mut conn = pool.acquire().await?;
apply_on(&mut conn).await
}
async fn apply_on(conn: &mut Conn) -> Result<(), SessionStoreError> {
conn.execute(sql::CREATE_HISTORY).await?;
let Some(lock) = sql::LOCK else {
return apply_pending(conn).await;
};
sqlx::query(lock).execute(&mut *conn).await?;
let result = apply_pending(&mut *conn).await;
if let Some(unlock) = sql::UNLOCK {
let _ = sqlx::query(unlock).execute(&mut *conn).await;
}
result
}
async fn apply_pending(conn: &mut Conn) -> Result<(), SessionStoreError> {
for &(version, migration) in MIGRATIONS {
let applied: i64 = sqlx::query(sql::COUNT_APPLIED)
.bind(version)
.fetch_one(&mut *conn)
.await?
.try_get::<i64, _>(0)?;
if applied > 0 {
continue;
}
for statement in statements(migration) {
conn.execute(statement).await?;
}
sqlx::query(sql::RECORD_APPLIED)
.bind(version)
.execute(&mut *conn)
.await?;
}
Ok(())
}
fn statements(migration: &str) -> impl Iterator<Item = &str> {
let mut pieces = Vec::new();
let mut start = 0;
let mut offset = 0;
for line in migration.split_inclusive('\n') {
if line.trim() == STATEMENT_SEPARATOR {
pieces.push(&migration[start..offset]);
start = offset + line.len();
}
offset += line.len();
}
pieces.push(&migration[start..]);
pieces
.into_iter()
.map(str::trim)
.filter(|statement| !statement.is_empty())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn the_bundled_migration_splits_into_at_least_one_statement() {
for &(version, migration) in MIGRATIONS {
let count = statements(migration).count();
assert!(count > 0, "{version} produced no statements");
}
}
#[test]
fn a_comment_that_mentions_the_separator_does_not_split_the_file() {
let sql = "-- Statements are separated by a line reading `--;;`.\n\
CREATE TABLE a (i INT)";
let split: Vec<_> = statements(sql).collect();
assert_eq!(split.len(), 1, "split into {split:?}");
assert!(split[0].starts_with("-- Statements"));
}
#[test]
fn no_bundled_statement_begins_mid_comment() {
for &(version, migration) in MIGRATIONS {
for statement in statements(migration) {
assert!(
!statement.starts_with('`'),
"{version} produced a fragment, not a statement: {statement:.40}"
);
}
}
}
#[test]
fn the_bundled_migration_creates_the_sessions_table() {
for &(version, migration) in MIGRATIONS {
assert!(
statements(migration)
.any(|statement| statement
.contains("CREATE TABLE IF NOT EXISTS arcature_sessions")),
"{version} does not create arcature_sessions"
);
}
}
}