use darling::ToTokens;
use proc_macro2::TokenStream;
use quote::{TokenStreamExt, quote};
use super::options::*;
fn bare_table_name(events_table_name: &str) -> &str {
events_table_name
.rsplit('.')
.next()
.expect("rsplit yields at least one element")
}
pub struct ErrorClassifier<'a> {
constraint_violation: syn::Ident,
events_table_name: &'a str,
table_name: &'a str,
needs_write_classifier: bool,
}
impl<'a> From<&'a RepositoryOptions> for ErrorClassifier<'a> {
fn from(opts: &'a RepositoryOptions) -> Self {
Self {
constraint_violation: opts.constraint_violation(),
events_table_name: opts.events_table_name(),
table_name: opts.table_name(),
needs_write_classifier: opts.columns.updates_needed() || opts.delete.is_soft(),
}
}
}
impl ToTokens for ErrorClassifier<'_> {
fn to_tokens(&self, tokens: &mut TokenStream) {
tokens.append_all(create_write_classifier_fn(
&self.constraint_violation,
self.events_table_name,
self.table_name,
));
if self.needs_write_classifier {
tokens.append_all(update_write_classifier_fn(
&self.constraint_violation,
self.events_table_name,
));
}
}
}
fn index_violation_arm(constraint_violation: &syn::Ident, events_table: &str) -> TokenStream {
quote! {
sqlx::Error::Database(db_err)
if db_err.table() != Some(#events_table)
&& es_entity::is_classified_constraint_violation(db_err.as_ref()) =>
{
let name = db_err.constraint().unwrap_or("unknown").to_owned();
match #constraint_violation::from_database(e, &name) {
Ok(rejection) => errlanes::Fail::Rejected(rejection),
Err(source) => errlanes::Fatal::from_error(errlanes::FatalKind::Invariant, source).with_context(name).into(),
}
}
}
}
fn create_write_classifier_fn(
constraint_violation: &syn::Ident,
events_table_name: &str,
table_name: &str,
) -> TokenStream {
let events_table = bare_table_name(events_table_name);
let index_pkey = format!("{table_name}_pkey");
let index_arm = index_violation_arm(constraint_violation, events_table);
quote! {
#[inline(always)]
fn classify_create_write(e: sqlx::Error) -> es_entity::RepoWriteError<#constraint_violation> {
match &e {
sqlx::Error::Database(db_err)
if db_err.is_unique_violation()
&& db_err.table() == Some(#events_table) =>
{
match #constraint_violation::from_database(e, #index_pkey) {
Ok(rejection) => errlanes::Fail::Rejected(rejection),
Err(source) => errlanes::Fatal::from_error(errlanes::FatalKind::Invariant, source).with_context(#index_pkey).into(),
}
}
#index_arm
_ => errlanes::Fail::from(e),
}
}
}
}
fn update_write_classifier_fn(
constraint_violation: &syn::Ident,
events_table_name: &str,
) -> TokenStream {
let events_table = bare_table_name(events_table_name);
let index_arm = index_violation_arm(constraint_violation, events_table);
quote! {
#[inline(always)]
fn classify_update_write(
e: sqlx::Error,
context: impl Into<std::borrow::Cow<'static, str>>,
) -> es_entity::RepoWriteError<#constraint_violation> {
match &e {
sqlx::Error::Database(db_err)
if db_err.is_unique_violation()
&& db_err.table() == Some(#events_table) =>
{
errlanes::Fail::from(
errlanes::Transient::from_error(errlanes::TransientKind::OptimisticConflict, e)
.with_context(context)
)
}
#index_arm
_ => errlanes::Fail::from(e),
}
}
}
}
pub fn classify_conflict_fn() -> TokenStream {
quote! {
#[inline(always)]
fn classify_conflict<T, D>(
res: Result<T, sqlx::Error>,
events_table: &'static str,
context: impl FnOnce() -> String,
) -> Result<T, es_entity::RepoWriteError<D>> {
let events_table = events_table.rsplit('.').next().unwrap_or(events_table);
match res {
Ok(v) => Ok(v),
Err(e) if e.as_database_error().is_some_and(|db_err|
db_err.is_unique_violation() && db_err.table() == Some(events_table)) =>
{
Err(errlanes::Fail::from(
errlanes::Transient::from_error(errlanes::TransientKind::OptimisticConflict, e)
.with_context(context())
))
}
Err(e) if e.as_database_error().is_some_and(|db_err| db_err.is_unique_violation()) => {
Err(errlanes::Fail::from(errlanes::Fatal::from_error(errlanes::FatalKind::Invariant, e).with_context(context())))
}
Err(e) => Err(errlanes::Fail::from(e)),
}
}
}
}
#[cfg(test)]
mod tests {
use quote::ToTokens;
use super::*;
#[test]
fn classify_conflict_strips_schema_prefix_before_comparing() {
let output = classify_conflict_fn().into_token_stream().to_string();
assert!(
output.contains("events_table . rsplit ('.') . next ()"),
"classify_conflict must strip a schema prefix off `events_table` \
before comparing it to `db_err.table()`: {output}"
);
}
}