use async_trait::async_trait;
use sea_query::{Asterisk, Expr, ExprTrait, Order, Query};
use serde::Deserialize;
use crate::errors::OrionError;
use crate::storage::models::{PackageReceipt, PackageState};
use crate::storage::{DbPool, build_sqlx, schema::Packages};
use super::helpers::{fetch_required_tx, sql_now};
#[derive(Debug, Deserialize, utoipa::ToSchema)]
pub struct PutPackageReceiptRequest {
pub version: String,
pub content_hash: String,
pub state: PackageState,
}
#[async_trait]
pub trait PackageRepository: Send + Sync {
async fn put(
&self,
name: &str,
req: &PutPackageReceiptRequest,
principal: &str,
) -> Result<PackageReceipt, OrionError>;
async fn list(
&self,
limit: i64,
offset: i64,
) -> Result<super::helpers::PaginatedResult<PackageReceipt>, OrionError>;
async fn get_by_name(&self, name: &str) -> Result<Vec<PackageReceipt>, OrionError>;
}
pub struct SqlPackageRepository {
pool: DbPool,
}
impl SqlPackageRepository {
pub fn new(pool: DbPool) -> Self {
Self { pool }
}
}
fn receipt_select(name: &str, version: &str) -> sea_query::SelectStatement {
Query::select()
.column(Asterisk)
.from(Packages::Table)
.and_where(Expr::col(Packages::Name).eq(name))
.and_where(Expr::col(Packages::Version).eq(version))
.to_owned()
}
fn applied_conflict(name: &str, version: &str) -> OrionError {
OrionError::Conflict(format!(
"package '{name}' version '{version}' is already applied here with different \
content — an applied package version is immutable; bump the package version"
))
}
#[async_trait]
impl PackageRepository for SqlPackageRepository {
async fn put(
&self,
name: &str,
req: &PutPackageReceiptRequest,
principal: &str,
) -> Result<PackageReceipt, OrionError> {
crate::metrics::timed_db_op("packages.put", async {
let backend = crate::storage::get_backend();
let mut tx = self.pool.begin_tx().await?;
let (sql, values) = build_sqlx(&mut receipt_select(name, &req.version));
let existing: Option<PackageReceipt> = tx.fetch_optional_as(&sql, values).await?;
match existing {
None => {
let (sql, values) = build_sqlx(
Query::insert()
.into_table(Packages::Table)
.columns([
Packages::Name,
Packages::Version,
Packages::ContentHash,
Packages::State,
Packages::Principal,
])
.values_panic([
name.into(),
req.version.as_str().into(),
req.content_hash.as_str().into(),
req.state.as_str().into(),
principal.into(),
]),
);
tx.execute_query(&sql, values).await.map_err(|e| {
super::helpers::map_duplicate(e, || {
format!(
"package '{name}' version '{}' was recorded by a \
concurrent request — retry to read it",
req.version
)
})
})?;
}
Some(row) if row.state == PackageState::Applied.as_str() => {
if row.content_hash != req.content_hash {
return Err(applied_conflict(name, &req.version));
}
if req.state == PackageState::Staged {
return Err(OrionError::Conflict(format!(
"package '{name}' version '{}' is already applied — an \
applied package version cannot go back to staged",
req.version
)));
}
let (sql, values) = build_sqlx(
Query::update()
.table(Packages::Table)
.value(Packages::Principal, principal)
.value(Packages::UpdatedAt, Expr::cust(sql_now(backend)))
.and_where(Expr::col(Packages::Name).eq(name))
.and_where(Expr::col(Packages::Version).eq(req.version.as_str()))
.and_where(
Expr::col(Packages::State).eq(PackageState::Applied.as_str()),
)
.and_where(
Expr::col(Packages::ContentHash).eq(req.content_hash.as_str()),
),
);
if tx.execute_query(&sql, values).await? == 0 {
return Err(applied_conflict(name, &req.version));
}
}
Some(_) => {
let (sql, values) = build_sqlx(
Query::update()
.table(Packages::Table)
.value(Packages::ContentHash, req.content_hash.as_str())
.value(Packages::State, req.state.as_str())
.value(Packages::Principal, principal)
.value(Packages::UpdatedAt, Expr::cust(sql_now(backend)))
.and_where(Expr::col(Packages::Name).eq(name))
.and_where(Expr::col(Packages::Version).eq(req.version.as_str()))
.and_where(
Expr::col(Packages::State).eq(PackageState::Staged.as_str()),
),
);
if tx.execute_query(&sql, values).await? == 0 {
return Err(OrionError::Conflict(format!(
"package '{name}' version '{}' was applied by a concurrent \
request — re-check its receipt before writing again",
req.version
)));
}
}
}
let (sql, values) = build_sqlx(&mut receipt_select(name, &req.version));
let row = fetch_required_tx(&mut tx, &sql, values, || {
OrionError::internal(format!(
"package receipt '{name}' version '{}' vanished mid-write",
req.version
))
})
.await?;
tx.commit().await?;
Ok(row)
})
.await
}
async fn list(
&self,
limit: i64,
offset: i64,
) -> Result<super::helpers::PaginatedResult<PackageReceipt>, OrionError> {
crate::metrics::timed_db_op("packages.list", async {
let total = super::helpers::count_where(
&self.pool,
Packages::Table,
sea_query::Condition::all(),
)
.await?;
let (sql, values) = build_sqlx(
Query::select()
.column(Asterisk)
.from(Packages::Table)
.order_by(Packages::Name, Order::Asc)
.order_by(Packages::UpdatedAt, Order::Desc)
.order_by(Packages::Version, Order::Asc)
.limit(limit as u64)
.offset(offset as u64),
);
let data: Vec<PackageReceipt> = self.pool.fetch_all_as(&sql, values).await?;
Ok(super::helpers::PaginatedResult {
data,
total,
limit,
offset,
})
})
.await
}
async fn get_by_name(&self, name: &str) -> Result<Vec<PackageReceipt>, OrionError> {
crate::metrics::timed_db_op("packages.get_by_name", async {
let (sql, values) = build_sqlx(
Query::select()
.column(Asterisk)
.from(Packages::Table)
.and_where(Expr::col(Packages::Name).eq(name))
.order_by(Packages::UpdatedAt, Order::Desc)
.order_by(Packages::Version, Order::Desc),
);
let rows: Vec<PackageReceipt> = self.pool.fetch_all_as(&sql, values).await?;
if rows.is_empty() {
return Err(OrionError::NotFound(format!(
"Package '{name}' has no receipts"
)));
}
Ok(rows)
})
.await
}
}