use anyhow::Result;
use async_trait::async_trait;
use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use sqlx::PgPool;
use uuid::Uuid;
#[derive(Debug, Clone, Serialize, Deserialize, sqlx::FromRow)]
pub struct Snapshot {
pub aggregate_id: Uuid,
pub aggregate_type: String,
pub version: i64,
pub state: serde_json::Value,
pub created_at: DateTime<Utc>,
}
#[async_trait]
pub trait SnapshotStore: Send + Sync {
async fn save(&self, snapshot: &Snapshot) -> Result<()>;
async fn load(&self, aggregate_id: Uuid) -> Result<Option<Snapshot>>;
async fn delete(&self, aggregate_id: Uuid) -> Result<()>;
}
pub struct PostgresSnapshotStore {
pool: PgPool,
table_name: String,
}
impl PostgresSnapshotStore {
pub fn new(pool: PgPool) -> Self {
Self {
pool,
table_name: "catalog.aggregate_snapshots".to_string(),
}
}
pub fn with_table_name(pool: PgPool, table_name: impl Into<String>) -> Self {
Self {
pool,
table_name: table_name.into(),
}
}
}
#[async_trait]
impl SnapshotStore for PostgresSnapshotStore {
async fn save(&self, snapshot: &Snapshot) -> Result<()> {
let query = format!(
"INSERT INTO {} (aggregate_id, aggregate_type, version, state, created_at)
VALUES ($1, $2, $3, $4, $5)
ON CONFLICT (aggregate_id) DO UPDATE SET
version = EXCLUDED.version, state = EXCLUDED.state, created_at = EXCLUDED.created_at",
self.table_name
);
sqlx::query(&query)
.bind(snapshot.aggregate_id)
.bind(&snapshot.aggregate_type)
.bind(snapshot.version)
.bind(&snapshot.state)
.bind(snapshot.created_at)
.execute(&self.pool)
.await?;
Ok(())
}
async fn load(&self, aggregate_id: Uuid) -> Result<Option<Snapshot>> {
let query = format!(
"SELECT * FROM {} WHERE aggregate_id = $1",
self.table_name
);
let snapshot = sqlx::query_as::<_, Snapshot>(&query)
.bind(aggregate_id)
.fetch_optional(&self.pool)
.await?;
Ok(snapshot)
}
async fn delete(&self, aggregate_id: Uuid) -> Result<()> {
let query = format!(
"DELETE FROM {} WHERE aggregate_id = $1",
self.table_name
);
sqlx::query(&query)
.bind(aggregate_id)
.execute(&self.pool)
.await?;
Ok(())
}
}
#[derive(Debug, Clone)]
pub struct SnapshotStrategy {
pub enabled: bool,
pub every_n_events: u32,
pub max_age_seconds: Option<u64>,
pub storage: Option<String>,
}
impl Default for SnapshotStrategy {
fn default() -> Self {
Self {
enabled: true,
every_n_events: 100,
max_age_seconds: Some(86400), storage: None,
}
}
}
impl SnapshotStrategy {
pub fn new(enabled: bool, every_n_events: u32) -> Self {
Self {
enabled,
every_n_events,
max_age_seconds: None,
storage: None,
}
}
pub fn with_max_age(mut self, seconds: u64) -> Self {
self.max_age_seconds = Some(seconds);
self
}
pub fn with_storage(mut self, storage: impl Into<String>) -> Self {
self.storage = Some(storage.into());
self
}
pub fn should_snapshot(&self, current_version: i64, last_snapshot_version: i64) -> bool {
self.enabled && (current_version - last_snapshot_version) as u32 >= self.every_n_events
}
}