use crate::error::{DbError, DbResult};
use crate::queue::{Job, JobStatus};
use crate::storage::{Document, StorageEngine};
use serde::{Deserialize, Serialize};
use serde_json::Value as JsonValue;
use std::collections::HashMap;
use std::sync::Arc;
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>,
}
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 mut triggers = Vec::new();
for doc in triggers_coll.scan(None) {
if let Ok(trigger) = serde_json::from_value::<Trigger>(doc.to_value()) {
if trigger.collection == collection_name && trigger.enabled {
triggers.push(trigger);
}
}
}
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));
}
}