use crate::error::{DbError, DbResult};
use crate::queue::{Job, JobStatus};
use crate::storage::collection::ChangeEvent;
use crate::storage::{Document, StorageEngine};
use serde::{Deserialize, Serialize};
use serde_json::Value as JsonValue;
use std::collections::HashMap;
use std::sync::{Arc, OnceLock};
use std::time::{Duration, Instant};
use tokio::sync::broadcast;
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq)]
#[serde(rename_all = "lowercase")]
pub enum TriggerEvent {
Insert,
Update,
Delete,
}
impl TriggerEvent {
pub fn as_str(&self) -> &'static str {
match self {
TriggerEvent::Insert => "insert",
TriggerEvent::Update => "update",
TriggerEvent::Delete => "delete",
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Trigger {
#[serde(rename = "_key")]
pub id: String,
#[serde(rename = "_rev", skip_serializing_if = "Option::is_none")]
pub revision: Option<String>,
pub name: String,
pub collection: String,
pub events: Vec<TriggerEvent>,
#[serde(default)]
pub script_path: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub webhook_url: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub webhook_secret: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub webhook_headers: Option<HashMap<String, String>>,
#[serde(default = "default_queue")]
pub queue: String,
#[serde(default)]
pub priority: i32,
#[serde(default = "default_max_retries")]
pub max_retries: i32,
#[serde(default = "default_enabled")]
pub enabled: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub filter: Option<String>,
pub created_at: u64,
pub updated_at: u64,
}
fn default_queue() -> String {
"default".to_string()
}
fn default_max_retries() -> i32 {
5
}
fn default_enabled() -> bool {
true
}
impl Trigger {
pub fn new(
name: String,
collection: String,
events: Vec<TriggerEvent>,
script_path: String,
) -> Self {
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs();
Self {
id: uuid::Uuid::new_v4().to_string(),
revision: None,
name,
collection,
events,
script_path,
webhook_url: None,
webhook_secret: None,
webhook_headers: None,
queue: default_queue(),
priority: 0,
max_retries: default_max_retries(),
enabled: default_enabled(),
filter: None,
created_at: now,
updated_at: now,
}
}
pub fn matches_event(&self, event: &TriggerEvent) -> bool {
self.enabled && self.events.contains(event)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct TriggerJobParams {
pub trigger_name: String,
pub event: String,
pub collection: String,
pub key: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub data: Option<JsonValue>,
#[serde(skip_serializing_if = "Option::is_none")]
pub old_data: Option<JsonValue>,
}
const TRIGGER_CACHE_MAX_AGE: Duration = Duration::from_secs(30);
struct TriggerCacheEntry {
source: Arc<broadcast::Sender<ChangeEvent>>,
changes: broadcast::Receiver<ChangeEvent>,
loaded_at: Instant,
by_collection: HashMap<String, Vec<Trigger>>,
}
impl TriggerCacheEntry {
fn is_fresh(&mut self, source: &Arc<broadcast::Sender<ChangeEvent>>) -> bool {
Arc::ptr_eq(&self.source, source)
&& self.loaded_at.elapsed() < TRIGGER_CACHE_MAX_AGE
&& matches!(
self.changes.try_recv(),
Err(broadcast::error::TryRecvError::Empty)
)
}
}
fn trigger_cache() -> &'static dashmap::DashMap<String, TriggerCacheEntry> {
static CACHE: OnceLock<dashmap::DashMap<String, TriggerCacheEntry>> = OnceLock::new();
CACHE.get_or_init(dashmap::DashMap::new)
}
pub fn invalidate_trigger_cache(db_name: &str) {
trigger_cache().remove(db_name);
}
pub struct TriggerManager {
storage: Arc<StorageEngine>,
notifier: Option<broadcast::Sender<()>>,
}
impl TriggerManager {
pub fn new(storage: Arc<StorageEngine>) -> Self {
Self {
storage,
notifier: None,
}
}
pub fn with_notifier(mut self, notifier: broadcast::Sender<()>) -> Self {
self.notifier = Some(notifier);
self
}
pub fn get_triggers_for_collection(
&self,
db_name: &str,
collection_name: &str,
) -> DbResult<Vec<Trigger>> {
let db = self.storage.get_database(db_name)?;
let triggers_coll = match db.get_collection("_triggers") {
Ok(coll) => coll,
Err(_) => return Ok(Vec::new()), };
let source = triggers_coll.change_sender.clone();
let cache = trigger_cache();
if let Some(mut entry) = cache.get_mut(db_name) {
if entry.is_fresh(&source) {
return Ok(entry
.by_collection
.get(collection_name)
.cloned()
.unwrap_or_default());
}
}
let changes = source.subscribe();
let mut by_collection: HashMap<String, Vec<Trigger>> = HashMap::new();
for doc in triggers_coll.scan(None) {
if let Ok(trigger) = serde_json::from_value::<Trigger>(doc.to_value()) {
if trigger.enabled {
by_collection
.entry(trigger.collection.clone())
.or_default()
.push(trigger);
}
}
}
let triggers = by_collection
.get(collection_name)
.cloned()
.unwrap_or_default();
cache.insert(
db_name.to_string(),
TriggerCacheEntry {
source,
changes,
loaded_at: Instant::now(),
by_collection,
},
);
Ok(triggers)
}
pub fn fire_triggers(
&self,
db_name: &str,
collection_name: &str,
event: TriggerEvent,
doc: &Document,
old_doc: Option<&JsonValue>,
) -> DbResult<Vec<String>> {
if collection_name.starts_with('_') {
return Ok(Vec::new());
}
let triggers = self.get_triggers_for_collection(db_name, collection_name)?;
let mut job_ids = Vec::new();
for trigger in triggers {
if !trigger.matches_event(&event) {
continue;
}
match self.create_trigger_job(db_name, &trigger, &event, doc, old_doc) {
Ok(job_id) => {
tracing::info!(
"Trigger '{}' fired for {} on {}/{}: created job {}",
trigger.name,
event.as_str(),
collection_name,
&doc.key,
job_id
);
job_ids.push(job_id);
}
Err(e) => {
tracing::error!("Failed to create job for trigger '{}': {}", trigger.name, e);
}
}
}
if !job_ids.is_empty() {
if let Some(ref notifier) = self.notifier {
let _ = notifier.send(());
}
}
Ok(job_ids)
}
fn create_trigger_job(
&self,
db_name: &str,
trigger: &Trigger,
event: &TriggerEvent,
doc: &Document,
old_doc: Option<&JsonValue>,
) -> DbResult<String> {
let db = self.storage.get_database(db_name)?;
if db.get_collection("_jobs").is_err() {
db.create_collection("_jobs".to_string(), None)?;
}
let jobs_coll = db.get_collection("_jobs")?;
let now = std::time::SystemTime::now()
.duration_since(std::time::UNIX_EPOCH)
.unwrap()
.as_secs();
let params = TriggerJobParams {
trigger_name: trigger.name.clone(),
event: event.as_str().to_string(),
collection: trigger.collection.clone(),
key: doc.key.clone(),
data: match event {
TriggerEvent::Delete => None,
_ => Some(doc.to_value()),
},
old_data: old_doc.cloned(),
};
let job = Job {
id: uuid::Uuid::new_v4().to_string(),
revision: None,
queue: trigger.queue.clone(),
priority: trigger.priority,
script_path: trigger.script_path.clone(),
webhook_url: trigger.webhook_url.clone(),
webhook_secret: trigger.webhook_secret.clone(),
webhook_headers: trigger.webhook_headers.clone(),
params: serde_json::to_value(¶ms).unwrap_or(JsonValue::Null),
status: JobStatus::Pending,
retry_count: 0,
max_retries: trigger.max_retries,
last_error: None,
cron_job_id: None,
run_at: now,
created_at: now,
started_at: None,
completed_at: None,
};
let job_id = job.id.clone();
let job_val = serde_json::to_value(&job)
.map_err(|e| DbError::InternalError(format!("Failed to serialize job: {}", e)))?;
jobs_coll.insert(job_val)?;
Ok(job_id)
}
}
pub fn fire_collection_triggers(
storage: &Arc<StorageEngine>,
notifier: Option<&broadcast::Sender<()>>,
db_name: &str,
collection_name: &str,
event: TriggerEvent,
doc: &Document,
old_doc: Option<&JsonValue>,
) -> DbResult<Vec<String>> {
let mut manager = TriggerManager::new(storage.clone());
if let Some(n) = notifier {
manager = manager.with_notifier(n.clone());
}
manager.fire_triggers(db_name, collection_name, event, doc, old_doc)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_trigger_event_serialization() {
let event = TriggerEvent::Insert;
let json = serde_json::to_string(&event).unwrap();
assert_eq!(json, "\"insert\"");
let parsed: TriggerEvent = serde_json::from_str(&json).unwrap();
assert_eq!(parsed, TriggerEvent::Insert);
}
#[test]
fn test_trigger_creation() {
let trigger = Trigger::new(
"test_trigger".to_string(),
"users".to_string(),
vec![TriggerEvent::Insert, TriggerEvent::Update],
"triggers/welcome.lua".to_string(),
);
assert_eq!(trigger.name, "test_trigger");
assert_eq!(trigger.collection, "users");
assert_eq!(trigger.events.len(), 2);
assert!(trigger.enabled);
assert_eq!(trigger.queue, "default");
assert_eq!(trigger.max_retries, 5);
}
#[test]
fn test_trigger_matches_event() {
let trigger = Trigger::new(
"test".to_string(),
"users".to_string(),
vec![TriggerEvent::Insert],
"test.lua".to_string(),
);
assert!(trigger.matches_event(&TriggerEvent::Insert));
assert!(!trigger.matches_event(&TriggerEvent::Update));
assert!(!trigger.matches_event(&TriggerEvent::Delete));
}
#[test]
fn test_trigger_disabled() {
let mut trigger = Trigger::new(
"test".to_string(),
"users".to_string(),
vec![TriggerEvent::Insert],
"test.lua".to_string(),
);
trigger.enabled = false;
assert!(!trigger.matches_event(&TriggerEvent::Insert));
}
#[test]
fn trigger_cache_sees_direct_writes() {
let dir = tempfile::TempDir::new().unwrap();
let storage = Arc::new(StorageEngine::new(dir.path()).unwrap());
storage.create_database("tdb".to_string()).unwrap();
let db = storage.get_database("tdb").unwrap();
db.create_collection("_triggers".to_string(), None).unwrap();
let coll = db.get_collection("_triggers").unwrap();
let manager = TriggerManager::new(storage.clone());
assert!(manager
.get_triggers_for_collection("tdb", "users")
.unwrap()
.is_empty());
let t = Trigger::new(
"t1".to_string(),
"users".to_string(),
vec![TriggerEvent::Insert],
"a.lua".to_string(),
);
let id = t.id.clone();
coll.insert(serde_json::to_value(&t).unwrap()).unwrap();
assert_eq!(
manager
.get_triggers_for_collection("tdb", "users")
.unwrap()
.len(),
1
);
assert_eq!(
manager
.get_triggers_for_collection("tdb", "users")
.unwrap()
.len(),
1
);
assert!(manager
.get_triggers_for_collection("tdb", "orders")
.unwrap()
.is_empty());
coll.delete(&id).unwrap();
assert!(manager
.get_triggers_for_collection("tdb", "users")
.unwrap()
.is_empty());
}
}