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 fn extract_concurrent_modification_fn() -> TokenStream {
let mut tokens = TokenStream::new();
tokens.append_all(quote! {
fn extract_concurrent_modification<T, __EsErr: From<sqlx::Error>>(
res: Result<T, sqlx::Error>,
concurrent_modification: __EsErr,
) -> Result<T, __EsErr> {
match res {
Ok(v) => Ok(v),
Err(sqlx::Error::Database(ref db_err)) if db_err.is_unique_violation() => {
Err(concurrent_modification)
}
Err(e) => Err(__EsErr::from(e)),
}
}
});
tokens
}
pub fn concurrent_modification_classifier(
error: &syn::Ident,
events_table_name: &str,
) -> TokenStream {
let events_table = bare_table_name(events_table_name);
quote! {
|e| match &e {
sqlx::Error::Database(db_err)
if db_err.is_unique_violation()
&& db_err.table() == Some(#events_table) =>
{
#error::ConcurrentModification
}
_ => #error::Sqlx(e),
}
}
}
fn index_violation_arm(error: &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()) =>
{
#error::ConstraintViolation {
column: Self::map_constraint_column(db_err.constraint()),
value: es_entity::extract_constraint_value(db_err.as_ref()),
inner: e,
}
}
}
}
pub struct ErrorClassifier<'a> {
create_error: syn::Ident,
modify_error: 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 {
create_error: opts.create_error(),
modify_error: opts.modify_error(),
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_error_classifier_fn(
&self.create_error,
self.events_table_name,
self.table_name,
));
if self.needs_write_classifier {
tokens.append_all(write_error_classifier_fn(
&self.modify_error,
self.events_table_name,
));
}
}
}
fn create_error_classifier_fn(
error: &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(error, events_table);
quote! {
#[inline(always)]
fn classify_create_error(e: sqlx::Error) -> #error {
match &e {
sqlx::Error::Database(db_err)
if db_err.is_unique_violation()
&& db_err.table() == Some(#events_table) =>
{
#error::ConstraintViolation {
column: Self::map_constraint_column(Some(#index_pkey)),
value: es_entity::extract_events_pkey_id_value(db_err.as_ref()),
inner: e,
}
}
#index_arm
_ => #error::Sqlx(e),
}
}
}
}
fn write_error_classifier_fn(error: &syn::Ident, events_table_name: &str) -> TokenStream {
let events_table = bare_table_name(events_table_name);
let index_arm = index_violation_arm(error, events_table);
quote! {
#[inline(always)]
fn classify_write_error(e: sqlx::Error) -> #error {
match &e {
sqlx::Error::Database(db_err)
if db_err.is_unique_violation()
&& db_err.table() == Some(#events_table) =>
{
#error::ConcurrentModification
}
#index_arm
_ => #error::Sqlx(e),
}
}
}
}