use proc_macro2::TokenStream;
use quote::{TokenStreamExt, quote};
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),
}
}
}
pub fn write_error_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
}
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,
}
}
_ => #error::Sqlx(e),
}
}
}