use sea_query::{Asterisk, Condition, DynIden, ExprTrait, Order, Query};
use serde::Serialize;
use crate::errors::OrionError;
use crate::storage::{DbPool, DbRow, DbTransaction};
#[derive(Debug, Serialize)]
pub struct PaginatedResult<T> {
pub data: Vec<T>,
pub total: i64,
pub limit: i64,
pub offset: i64,
}
#[derive(Debug, Default, serde::Deserialize, Serialize, utoipa::IntoParams)]
#[into_params(parameter_in = Query)]
pub struct VersionFilter {
pub limit: Option<i64>,
pub offset: Option<i64>,
}
pub fn optional_string_value(opt: Option<&str>) -> sea_query::Value {
opt.map(|s| s.to_string().into())
.unwrap_or(sea_query::Value::String(None))
}
pub fn map_duplicate(e: sqlx::Error, conflict_msg: impl FnOnce() -> String) -> OrionError {
match &e {
sqlx::Error::Database(db_err)
if db_err.kind() == sqlx::error::ErrorKind::UniqueViolation
|| db_err.message().contains("Only one draft version allowed") =>
{
OrionError::Conflict(conflict_msg())
}
_ => OrionError::Storage(e),
}
}
pub fn sql_now(backend: crate::storage::DbBackend) -> &'static str {
match backend {
crate::storage::DbBackend::Sqlite => "datetime('now')",
crate::storage::DbBackend::Postgres => "LOCALTIMESTAMP",
crate::storage::DbBackend::Mysql => "UTC_TIMESTAMP()",
}
}
pub fn sql_now_plus_secs(backend: crate::storage::DbBackend, secs: u64) -> String {
match backend {
crate::storage::DbBackend::Sqlite => format!("datetime('now', '+{secs} seconds')"),
crate::storage::DbBackend::Postgres => {
format!("LOCALTIMESTAMP + interval '{secs} seconds'")
}
crate::storage::DbBackend::Mysql => {
format!("DATE_ADD(UTC_TIMESTAMP(), INTERVAL {secs} SECOND)")
}
}
}
pub async fn fetch_required<T: DbRow>(
pool: &DbPool,
sql: &str,
values: sea_query_sqlx::SqlxValues,
err: impl FnOnce() -> OrionError,
) -> Result<T, OrionError> {
pool.fetch_optional_as::<T>(sql, values)
.await?
.ok_or_else(err)
}
pub async fn count_where<I>(
pool: &DbPool,
table: I,
cond: sea_query::Condition,
) -> Result<i64, OrionError>
where
I: sea_query::IntoTableRef,
{
use sea_query::{Asterisk, Expr, Func, Query};
let (sql, values) = crate::storage::build_sqlx(
Query::select()
.expr(Func::count(Expr::col(Asterisk)))
.from(table)
.cond_where(cond),
);
let (total,): (i64,) = pool.fetch_one_as::<(i64,)>(&sql, values).await?;
Ok(total)
}
pub enum Projection {
All,
Columns(Vec<DynIden>),
}
pub struct Page {
pub from: DynIden,
pub projection: Projection,
pub cond: Condition,
pub sort: DynIden,
pub order: Order,
pub limit: i64,
pub offset: i64,
}
pub async fn paginate<T: DbRow>(
pool: &DbPool,
page: Page,
) -> Result<PaginatedResult<T>, OrionError> {
let total = count_where(pool, page.from.clone(), page.cond.clone()).await?;
let (sql, values) = crate::storage::build_sqlx(&mut page_select(&page));
let data = pool.fetch_all_as::<T>(&sql, values).await?;
Ok(PaginatedResult {
data,
total,
limit: page.limit,
offset: page.offset,
})
}
pub(crate) fn page_select(page: &Page) -> sea_query::SelectStatement {
let mut select = Query::select();
match &page.projection {
Projection::All => select.column(Asterisk),
Projection::Columns(columns) => select.columns(columns.iter().cloned()),
};
select
.from(page.from.clone())
.cond_where(page.cond.clone())
.order_by(page.sort.clone(), page.order.clone())
.limit(page.limit as u64)
.offset(page.offset as u64)
.to_owned()
}
pub async fn ensure_absent<T: DbRow>(
pool: &DbPool,
sql: &str,
values: sea_query_sqlx::SqlxValues,
err: impl FnOnce() -> OrionError,
) -> Result<(), OrionError> {
if pool.fetch_optional_as::<T>(sql, values).await?.is_some() {
return Err(err());
}
Ok(())
}
pub async fn fetch_required_tx<T: DbRow>(
tx: &mut DbTransaction,
sql: &str,
values: sea_query_sqlx::SqlxValues,
err: impl FnOnce() -> OrionError,
) -> Result<T, OrionError> {
tx.fetch_optional_as::<T>(sql, values)
.await?
.ok_or_else(err)
}
pub fn clamp_pagination(limit: Option<i64>, offset: Option<i64>) -> (i64, i64) {
let limit = limit.unwrap_or(50).clamp(1, 1000);
let offset = Ord::max(offset.unwrap_or(0), 0);
(limit, offset)
}
pub fn parse_sort_order(sort_order: Option<&str>) -> sea_query::Order {
match sort_order {
Some("asc") => sea_query::Order::Asc,
_ => sea_query::Order::Desc,
}
}
pub(crate) const EXPORT_PAGE_SIZE: i64 = 500;
pub async fn snapshot_pages<T: DbRow>(
pool: &DbPool,
page_size: i64,
mut select_for: impl FnMut(i64, i64) -> sea_query::SelectStatement,
) -> Result<Vec<T>, OrionError> {
assert!(
(1..=1000).contains(&page_size),
"page_size {page_size} is outside the repository clamp (1..=1000); \
a clamped page would silently truncate the snapshot"
);
let mut tx = pool.begin_tx().await.map_err(OrionError::Storage)?;
if crate::storage::get_backend() == crate::storage::DbBackend::Postgres {
tx.execute_query(
"SET TRANSACTION ISOLATION LEVEL REPEATABLE READ",
sea_query_sqlx::SqlxValues(sea_query::Values(Vec::new())),
)
.await?;
}
let mut out = Vec::new();
let mut offset = 0i64;
loop {
let (sql, values) = crate::storage::build_sqlx(&mut select_for(page_size, offset));
let page: Vec<T> = tx.fetch_all_as(&sql, values).await?;
let page_len = page.len() as i64;
out.extend(page);
if page_len < page_size {
tx.commit().await.map_err(OrionError::Storage)?;
return Ok(out);
}
offset += page_size;
}
}
pub fn tag_like_pattern(tag: &str) -> sea_query::LikeExpr {
let escaped = tag
.replace('\\', "\\\\")
.replace('%', "\\%")
.replace('_', "\\_");
sea_query::LikeExpr::new(format!("%\"{escaped}\"%")).escape('\\')
}
pub fn cutoff_hours_ago(now: chrono::NaiveDateTime, hours: u64) -> chrono::NaiveDateTime {
chrono::Duration::try_hours(i64::try_from(hours).unwrap_or(i64::MAX))
.and_then(|d| now.checked_sub_signed(d))
.unwrap_or(chrono::NaiveDateTime::MIN)
}
pub fn cutoff_days_ago(now: chrono::NaiveDateTime, days: u64) -> chrono::NaiveDateTime {
chrono::Duration::try_days(i64::try_from(days).unwrap_or(i64::MAX))
.and_then(|d| now.checked_sub_signed(d))
.unwrap_or(chrono::NaiveDateTime::MIN)
}
pub async fn update_returning_scalar<C>(
pool: &DbPool,
update: &mut sea_query::UpdateStatement,
returning: C,
read_back: &mut sea_query::SelectStatement,
missing: impl FnOnce() -> OrionError,
) -> Result<i64, OrionError>
where
C: sea_query::IntoColumnRef,
{
use crate::storage::{DbBackend, build_sqlx, get_backend};
match get_backend() {
DbBackend::Sqlite | DbBackend::Postgres => {
update.returning(sea_query::Query::returning().column(returning));
let (sql, values) = build_sqlx(update);
pool.fetch_scalar::<i64>(&sql, values)
.await
.map_err(OrionError::Storage)
}
DbBackend::Mysql => {
let mut tx = pool.begin_tx().await.map_err(OrionError::Storage)?;
let (sql, values) = build_sqlx(update);
tx.execute_query(&sql, values).await?;
let (sql, values) = build_sqlx(read_back);
let (value,): (i64,) = fetch_required_tx(&mut tx, &sql, values, missing).await?;
tx.commit().await.map_err(OrionError::Storage)?;
Ok(value)
}
}
}
pub enum WriteStatement<'a> {
Insert(&'a mut sea_query::InsertStatement),
Update(&'a mut sea_query::UpdateStatement),
}
impl WriteStatement<'_> {
fn supports_returning(&self, backend: crate::storage::DbBackend) -> bool {
match backend {
crate::storage::DbBackend::Postgres => true,
crate::storage::DbBackend::Sqlite => matches!(self, Self::Insert(_)),
crate::storage::DbBackend::Mysql => false,
}
}
fn build_returning_all(&mut self) -> (String, sea_query_sqlx::SqlxValues) {
match self {
Self::Insert(q) => crate::storage::build_sqlx(q.returning_all()),
Self::Update(q) => crate::storage::build_sqlx(q.returning_all()),
}
}
fn build(&mut self) -> (String, sea_query_sqlx::SqlxValues) {
match self {
Self::Insert(q) => crate::storage::build_sqlx(&mut **q),
Self::Update(q) => crate::storage::build_sqlx(&mut **q),
}
}
}
pub async fn write_returning_row<T: DbRow>(
pool: &DbPool,
mut write: WriteStatement<'_>,
read_back: &mut sea_query::SelectStatement,
map_write_err: impl FnOnce(sqlx::Error) -> OrionError,
missing: impl FnOnce() -> OrionError,
) -> Result<T, OrionError> {
if write.supports_returning(crate::storage::get_backend()) {
let (sql, values) = write.build_returning_all();
pool.fetch_optional_as::<T>(&sql, values)
.await
.map_err(map_write_err)?
.ok_or_else(missing)
} else {
let mut tx = pool.begin_tx().await.map_err(OrionError::Storage)?;
let (sql, values) = write.build();
tx.execute_query(&sql, values)
.await
.map_err(map_write_err)?;
let (sql, values) = crate::storage::build_sqlx(read_back);
let row = fetch_required_tx(&mut tx, &sql, values, missing).await?;
tx.commit().await.map_err(OrionError::Storage)?;
Ok(row)
}
}
pub async fn insert_if_absent<C>(
pool: &DbPool,
mut insert: sea_query::InsertStatement,
conflict_col: C,
) -> Result<u64, OrionError>
where
C: sea_query::IntoIden,
{
use crate::storage::{DbBackend, build_sqlx, get_backend};
match get_backend() {
DbBackend::Sqlite | DbBackend::Postgres => {
insert.on_conflict(
sea_query::OnConflict::column(conflict_col)
.do_nothing()
.to_owned(),
);
let (sql, values) = build_sqlx(&mut insert);
pool.execute_query(&sql, values)
.await
.map_err(OrionError::Storage)
}
DbBackend::Mysql => {
let (sql, values) = build_sqlx(&mut insert);
let sql = sql.replacen("INSERT INTO", "INSERT IGNORE INTO", 1);
pool.execute_query(&sql, values)
.await
.map_err(OrionError::Storage)
}
}
}
const DELETE_CHUNK_ROWS: u64 = 1_000;
const DELETE_MAX_CHUNKS: usize = 5_000;
pub async fn delete_chunked(
pool: &DbPool,
table: impl sea_query::IntoIden,
id_column: impl sea_query::IntoIden,
condition: sea_query::Condition,
) -> Result<u64, OrionError> {
use sea_query::{Alias, Expr, Query};
let table = table.into_iden();
let id_column = id_column.into_iden();
let mut total = 0u64;
for chunk in 0..DELETE_MAX_CHUNKS {
let inner = Query::select()
.column(id_column.clone())
.from(table.clone())
.cond_where(condition.clone())
.limit(DELETE_CHUNK_ROWS)
.to_owned();
let materialised = Query::select()
.column(id_column.clone())
.from_subquery(inner, Alias::new("d6_chunk"))
.to_owned();
let (sql, values) = crate::storage::build_sqlx(
Query::delete()
.from_table(table.clone())
.and_where(Expr::col(id_column.clone()).in_subquery(materialised)),
);
let deleted = pool.execute_query(&sql, values).await?;
total += deleted;
if deleted < DELETE_CHUNK_ROWS {
return Ok(total);
}
if chunk + 1 == DELETE_MAX_CHUNKS {
tracing::warn!(
deleted = total,
"Retention delete hit its per-tick chunk cap; the remainder \
is left for the next tick"
);
}
tokio::task::yield_now().await;
}
Ok(total)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn optional_string_some() {
let v = optional_string_value(Some("hello"));
assert_eq!(v, sea_query::Value::String(Some("hello".into())));
}
#[test]
fn optional_string_none() {
let v = optional_string_value(None);
assert_eq!(v, sea_query::Value::String(None));
}
#[test]
fn pagination_defaults() {
assert_eq!(clamp_pagination(None, None), (50, 0));
}
#[test]
fn pagination_clamps() {
assert_eq!(clamp_pagination(Some(0), Some(-5)), (1, 0));
assert_eq!(clamp_pagination(Some(9999), Some(10)), (1000, 10));
}
#[test]
fn sort_order_asc() {
assert!(matches!(
parse_sort_order(Some("asc")),
sea_query::Order::Asc
));
}
#[test]
fn sort_order_desc() {
assert!(matches!(
parse_sort_order(Some("desc")),
sea_query::Order::Desc
));
}
#[test]
fn sort_order_none_defaults_desc() {
assert!(matches!(parse_sort_order(None), sea_query::Order::Desc));
}
fn sample_page(projection: Projection) -> Page {
use sea_query::IntoIden;
Page {
from: crate::storage::schema::Traces::Table.into_iden(),
projection,
cond: sea_query::Condition::all(),
sort: crate::storage::schema::Traces::CreatedAt.into_iden(),
order: Order::Desc,
limit: 25,
offset: 50,
}
}
fn rendered(page: &Page) -> String {
super::page_select(page).to_string(sea_query::SqliteQueryBuilder)
}
#[test]
fn projection_all_selects_every_column() {
let sql = rendered(&sample_page(Projection::All));
assert!(sql.contains("SELECT *"), "{sql}");
}
#[test]
fn projection_columns_names_only_what_it_was_given() {
use sea_query::IntoIden;
let sql = rendered(&sample_page(Projection::Columns(vec![
crate::storage::schema::Traces::Id.into_iden(),
crate::storage::schema::Traces::Status.into_iden(),
])));
assert!(sql.contains(r#"SELECT "id", "status""#), "{sql}");
assert!(!sql.contains('*'), "{sql}");
assert!(!sql.contains("access_token_hash"), "{sql}");
}
#[test]
fn returning_dispatch_matches_backend_trigger_semantics() {
use crate::storage::DbBackend::{Mysql, Postgres, Sqlite};
use crate::storage::schema::Traces;
let mut insert = Query::insert()
.into_table(Traces::Table)
.columns([Traces::Id])
.values_panic(["t1".into()])
.to_owned();
let mut update = Query::update()
.table(Traces::Table)
.value(Traces::Status, "completed")
.to_owned();
assert!(WriteStatement::Insert(&mut insert).supports_returning(Postgres));
assert!(WriteStatement::Update(&mut update).supports_returning(Postgres));
assert!(WriteStatement::Insert(&mut insert).supports_returning(Sqlite));
assert!(!WriteStatement::Update(&mut update).supports_returning(Sqlite));
assert!(!WriteStatement::Insert(&mut insert).supports_returning(Mysql));
assert!(!WriteStatement::Update(&mut update).supports_returning(Mysql));
}
#[test]
fn returning_all_is_appended_to_both_statement_kinds() {
use crate::storage::schema::Traces;
use sea_query::PostgresQueryBuilder;
let mut insert = Query::insert()
.into_table(Traces::Table)
.columns([Traces::Id])
.values_panic(["t1".into()])
.to_owned();
let (sql, _) = insert.returning_all().build(PostgresQueryBuilder);
assert!(sql.ends_with("RETURNING *"), "{sql}");
let mut update = Query::update()
.table(Traces::Table)
.value(Traces::Status, "completed")
.to_owned();
let (sql, _) = update.returning_all().build(PostgresQueryBuilder);
assert!(sql.ends_with("RETURNING *"), "{sql}");
}
#[test]
fn page_bounds_and_empty_filter_render_as_expected() {
let sql = rendered(&sample_page(Projection::All));
assert!(sql.contains("LIMIT 25"), "{sql}");
assert!(sql.contains("OFFSET 50"), "{sql}");
assert!(sql.contains(r#"ORDER BY "created_at" DESC"#), "{sql}");
assert!(
!sql.contains(r#"WHERE "#) || sql.contains("WHERE TRUE"),
"{sql}"
);
}
}