use darling::ToTokens;
use proc_macro2::TokenStream;
use quote::{TokenStreamExt, quote};
use super::{
events_write::{EventSource, EventsInsert, SnapshotUpsert},
options::*,
};
pub struct PersistEventsFn<'a> {
entity: &'a syn::Ident,
id: &'a syn::Ident,
event: &'a syn::Ident,
events_table_name: &'a str,
event_ctx: bool,
forgettable_table_name: Option<&'a str>,
snapshot_table_name: Option<&'a str>,
}
impl<'a> From<&'a RepositoryOptions> for PersistEventsFn<'a> {
fn from(opts: &'a RepositoryOptions) -> Self {
Self {
entity: opts.entity(),
id: opts.id(),
event: opts.event(),
events_table_name: opts.events_table_name(),
event_ctx: opts.event_context_enabled(),
forgettable_table_name: opts.forgettable_table_name(),
snapshot_table_name: opts.snapshot_table_name(),
}
}
}
impl ToTokens for PersistEventsFn<'_> {
fn to_tokens(&self, tokens: &mut TokenStream) {
let events_insert = EventsInsert::new(self.events_table_name, self.event_ctx);
let source = EventSource::PerEntityStandalone {
id_param: 1,
offset_param: 3,
};
let n_base_params = 1 + events_insert.arg_exprs(&source).len();
let (ctx_var, ctx_arg) = if self.event_ctx {
(
quote! { let contexts = events.serialize_new_event_contexts(); },
quote! {
contexts.as_deref() as Option<&[es_entity::ContextData]>,
},
)
} else {
(quote! {}, quote! {})
};
let id_type = &self.id;
let event_type = &self.event;
let entity = &self.entity;
let id_tokens = quote! {
id as &#id_type
};
let snapshot_upsert = self
.snapshot_table_name
.map(|table| SnapshotUpsert { table });
let (events_ty, snapshot_param, query, snap_gather, snap_arg_adds) = match &snapshot_upsert
{
Some(su) => {
let head_p = n_base_params + 1;
let fp_p = head_p + 1;
let snap_p = fp_p + 1;
let first_p = snap_p + 1;
let cte = su.cte_per_entity("$1", "", head_p, fp_p, snap_p, first_p, 2);
let query = format!("WITH {cte} {}", events_insert.sql(&source, 2, 4));
let gather = quote! {
let __snapshot_json = snapshot.map(|s| {
es_entity::prelude::serde_json::to_value(s).expect("Failed to serialize snapshot")
});
let __snapshot_head = (events.len_persisted() + events.len_new()) as i32;
let __snapshot_first = events.entity_first_persisted_at();
};
let arg_adds = quote! {
__snapshot_head,
<<#entity as es_entity::EsEntity>::Snapshot as es_entity::EsSnapshot>::FINGERPRINT,
__snapshot_json as Option<es_entity::prelude::serde_json::Value>,
__snapshot_first,
};
(
quote! { es_entity::EntityEvents<#event_type, <#entity as es_entity::EsEntity>::Snapshot> },
quote! { , snapshot: Option<&<#entity as es_entity::EsEntity>::Snapshot> },
query,
gather,
arg_adds,
)
}
None => (
quote! { es_entity::EntityEvents<#event_type> },
quote! {},
events_insert.sql(&source, 2, 4),
quote! {},
quote! {},
),
};
let forgettable_code = if let Some(forgettable_tbl) = self.forgettable_table_name {
let conflict_clause = if snapshot_upsert.is_some() {
" ON CONFLICT (entity_id, sequence) DO UPDATE SET payload = EXCLUDED.payload"
} else {
""
};
let payload_insert_query = format!(
"INSERT INTO {forgettable_tbl} (entity_id, sequence, payload) SELECT $1, unnested.sequence, unnested.payload FROM UNNEST($2::INT[], $3::JSONB[]) AS unnested(sequence, payload){conflict_clause}"
);
let (snap_push, snap_delete) = if snapshot_upsert.is_some() {
let delete_query =
format!("DELETE FROM {forgettable_tbl} WHERE entity_id = $1 AND sequence = 0");
(
quote! {
if let Some(payload) = snapshot.and_then(es_entity::EsSnapshot::extract_forgettable_payloads) {
payload_sequences.push(0);
payload_values.push(payload);
}
},
quote! {
if snapshot.is_some()
&& snapshot.and_then(es_entity::EsSnapshot::extract_forgettable_payloads).is_none()
{
sqlx::query!(#delete_query, id as &#id_type)
.execute(op.as_executor())
.await?;
}
},
)
} else {
(quote! {}, quote! {})
};
quote! {
let mut payload_sequences: Vec<i32> = Vec::new();
let mut payload_values: Vec<es_entity::prelude::serde_json::Value> = Vec::new();
#snap_push
for (idx, event_with_ctx) in events.iter_new_events().enumerate() {
if let Some(payload) = #event_type::extract_forgettable_payloads(&event_with_ctx.event) {
payload_sequences.push((offset + 1 + idx) as i32);
payload_values.push(payload);
}
}
if !payload_sequences.is_empty() {
sqlx::query!(
#payload_insert_query,
id as &#id_type,
&payload_sequences,
&payload_values,
)
.execute(op.as_executor())
.await?;
}
#snap_delete
}
} else {
quote! {}
};
tokens.append_all(quote! {
async fn persist_events<OP>(
&self,
op: &mut OP,
events: &mut #events_ty
#snapshot_param
) -> Result<usize, sqlx::Error>
where
OP: es_entity::AtomicOperation + ?Sized,
{
let id = events.id();
if !events.any_new() {
return Ok(0);
}
let offset = events.len_persisted();
let events_types = events.new_event_types();
let serialized_events = events.serialize_new_events();
#ctx_var
#snap_gather
#forgettable_code
let now = op.maybe_now();
let rows = sqlx::query!(
#query,
#id_tokens,
now,
offset as i32,
&events_types,
&serialized_events,
#ctx_arg
#snap_arg_adds
).fetch_all(op.as_executor()).await?;
let recorded_at = rows
.first()
.map(|row| row.recorded_at)
.ok_or(sqlx::Error::RowNotFound)?;
let n_events = events.mark_new_events_persisted_at(recorded_at);
Ok(n_events)
}
});
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn persist_events_fn() {
let id = syn::parse_str("EntityId").unwrap();
let event = syn::Ident::new("EntityEvent", proc_macro2::Span::call_site());
let entity = syn::Ident::new("Entity", proc_macro2::Span::call_site());
let persist_fn = PersistEventsFn {
entity: &entity,
id: &id,
event: &event,
events_table_name: "entity_events",
event_ctx: true,
forgettable_table_name: None,
snapshot_table_name: None,
};
let mut tokens = TokenStream::new();
persist_fn.to_tokens(&mut tokens);
let expected = quote! {
async fn persist_events<OP>(
&self,
op: &mut OP,
events: &mut es_entity::EntityEvents<EntityEvent>
) -> Result<usize, sqlx::Error>
where
OP: es_entity::AtomicOperation + ?Sized,
{
let id = events.id();
if !events.any_new() {
return Ok(0);
}
let offset = events.len_persisted();
let events_types = events.new_event_types();
let serialized_events = events.serialize_new_events();
let contexts = events.serialize_new_event_contexts();
let now = op.maybe_now();
let rows = sqlx::query!(
"INSERT INTO entity_events (id, recorded_at, sequence, event_type, event, context) SELECT $1, COALESCE($2, NOW()), ROW_NUMBER() OVER () + $3, unnested.event_type, unnested.event, unnested.context FROM UNNEST($4::TEXT[], $5::JSONB[], $6::JSONB[]) AS unnested(event_type, event, context) RETURNING recorded_at",
id as &EntityId,
now,
offset as i32,
&events_types,
&serialized_events,
contexts.as_deref() as Option<&[es_entity::ContextData]>,
).fetch_all(op.as_executor()).await?;
let recorded_at = rows
.first()
.map(|row| row.recorded_at)
.ok_or(sqlx::Error::RowNotFound)?;
let n_events = events.mark_new_events_persisted_at(recorded_at);
Ok(n_events)
}
};
assert_eq!(tokens.to_string(), expected.to_string());
}
#[test]
fn persist_events_fn_without_event_context() {
let id = syn::parse_str("EntityId").unwrap();
let event = syn::Ident::new("EntityEvent", proc_macro2::Span::call_site());
let entity = syn::Ident::new("Entity", proc_macro2::Span::call_site());
let persist_fn = PersistEventsFn {
entity: &entity,
id: &id,
event: &event,
events_table_name: "entity_events",
event_ctx: false,
forgettable_table_name: None,
snapshot_table_name: None,
};
let mut tokens = TokenStream::new();
persist_fn.to_tokens(&mut tokens);
let expected = quote! {
async fn persist_events<OP>(
&self,
op: &mut OP,
events: &mut es_entity::EntityEvents<EntityEvent>
) -> Result<usize, sqlx::Error>
where
OP: es_entity::AtomicOperation + ?Sized,
{
let id = events.id();
if !events.any_new() {
return Ok(0);
}
let offset = events.len_persisted();
let events_types = events.new_event_types();
let serialized_events = events.serialize_new_events();
let now = op.maybe_now();
let rows = sqlx::query!(
"INSERT INTO entity_events (id, recorded_at, sequence, event_type, event) SELECT $1, COALESCE($2, NOW()), ROW_NUMBER() OVER () + $3, unnested.event_type, unnested.event FROM UNNEST($4::TEXT[], $5::JSONB[]) AS unnested(event_type, event) RETURNING recorded_at",
id as &EntityId,
now,
offset as i32,
&events_types,
&serialized_events,
).fetch_all(op.as_executor()).await?;
let recorded_at = rows
.first()
.map(|row| row.recorded_at)
.ok_or(sqlx::Error::RowNotFound)?;
let n_events = events.mark_new_events_persisted_at(recorded_at);
Ok(n_events)
}
};
assert_eq!(tokens.to_string(), expected.to_string());
}
}