1use crate::error::DbError;
24use std::collections::HashMap;
25use std::sync::{Arc, RwLock};
26
27#[derive(Debug, Clone, Default)]
35pub struct HookContext {
36 pub tenant_id: Option<i64>,
38 pub operator_id: Option<i64>,
40 pub timestamp: u64,
42 pub metadata: HashMap<String, String>,
44}
45
46impl HookContext {
47 pub fn new() -> Self {
49 Self::default()
50 }
51
52 pub fn with_tenant(mut self, tenant_id: i64) -> Self {
54 self.tenant_id = Some(tenant_id);
55 self
56 }
57
58 pub fn with_operator(mut self, operator_id: i64) -> Self {
60 self.operator_id = Some(operator_id);
61 self
62 }
63
64 pub fn with_timestamp(mut self, ts: u64) -> Self {
66 self.timestamp = ts;
67 self
68 }
69
70 pub fn set_meta(&mut self, key: impl Into<String>, value: impl Into<String>) {
72 self.metadata.insert(key.into(), value.into());
73 }
74
75 pub fn get_meta(&self, key: &str) -> Option<&String> {
77 self.metadata.get(key)
78 }
79}
80
81#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
103pub enum HookEvent {
104 BeforeInsert,
106 AfterInsert,
108 BeforeUpdate,
110 AfterUpdate,
112 BeforeDelete,
114 AfterDelete,
116 BeforeWrite,
118 AfterWrite,
120 BeforeSave,
122 AfterSave,
124 BeforeRestore,
126 AfterRestore,
128 BeforeFind,
130 AfterFind,
132 BeforeValidate,
134 AfterValidate,
136}
137
138impl HookEvent {
139 pub fn is_before(&self) -> bool {
141 matches!(
142 self,
143 HookEvent::BeforeInsert
144 | HookEvent::BeforeUpdate
145 | HookEvent::BeforeDelete
146 | HookEvent::BeforeWrite
147 | HookEvent::BeforeSave
148 | HookEvent::BeforeRestore
149 | HookEvent::BeforeFind
150 | HookEvent::BeforeValidate
151 )
152 }
153
154 pub fn is_after(&self) -> bool {
156 matches!(
157 self,
158 HookEvent::AfterInsert
159 | HookEvent::AfterUpdate
160 | HookEvent::AfterDelete
161 | HookEvent::AfterWrite
162 | HookEvent::AfterSave
163 | HookEvent::AfterRestore
164 | HookEvent::AfterFind
165 | HookEvent::AfterValidate
166 )
167 }
168
169 pub fn is_write_level(&self) -> bool {
171 matches!(
172 self,
173 HookEvent::BeforeWrite
174 | HookEvent::AfterWrite
175 | HookEvent::BeforeSave
176 | HookEvent::AfterSave
177 )
178 }
179
180 pub fn is_find_level(&self) -> bool {
182 matches!(self, HookEvent::BeforeFind | HookEvent::AfterFind)
183 }
184
185 pub fn is_validate_level(&self) -> bool {
187 matches!(self, HookEvent::BeforeValidate | HookEvent::AfterValidate)
188 }
189
190 pub fn is_fine_grained(&self) -> bool {
192 self.is_write_level()
193 || self.is_find_level()
194 || self.is_validate_level()
195 || matches!(self, HookEvent::BeforeRestore | HookEvent::AfterRestore)
196 }
197}
198
199pub type HookResult<T> = Result<T, DbError>;
205
206pub trait Hookable: crate::model::Model {
225 fn before_insert(_ctx: &mut HookContext) -> HookResult<()> {
227 Ok(())
228 }
229
230 fn after_insert(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
232 Ok(())
233 }
234
235 fn before_update(_ctx: &mut HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
237 Ok(())
238 }
239
240 fn after_update(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
242 Ok(())
243 }
244
245 fn before_delete(_ctx: &mut HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
247 Ok(())
248 }
249
250 fn after_delete(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
252 Ok(())
253 }
254
255 fn before_write(_ctx: &mut HookContext) -> HookResult<()> {
259 Ok(())
260 }
261
262 fn after_write(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
264 Ok(())
265 }
266
267 fn before_save(_ctx: &mut HookContext) -> HookResult<()> {
269 Ok(())
270 }
271
272 fn after_save(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
274 Ok(())
275 }
276
277 fn before_restore(_ctx: &mut HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
281 Ok(())
282 }
283
284 fn after_restore(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
286 Ok(())
287 }
288
289 fn before_find(_ctx: &mut HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
294 Ok(())
295 }
296
297 fn after_find(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
302 Ok(())
303 }
304
305 fn before_validate(_ctx: &mut HookContext) -> HookResult<()> {
313 Ok(())
314 }
315
316 fn validate(_ctx: &mut HookContext) -> HookResult<()> {
322 Ok(())
323 }
324
325 fn after_validate(_ctx: &HookContext) -> HookResult<()> {
329 Ok(())
330 }
331}
332
333pub struct HookDispatcher;
337
338impl HookDispatcher {
339 pub fn insert<M, F>(ctx: &mut HookContext, f: F) -> HookResult<M::PrimaryKey>
345 where
346 M: Hookable,
347 F: FnOnce(&mut HookContext) -> HookResult<M::PrimaryKey>,
348 {
349 M::before_write(ctx)?;
350 M::before_save(ctx)?;
351 M::before_validate(ctx)?;
352 M::validate(ctx)?;
353 M::after_validate(ctx)?;
354 M::before_insert(ctx)?;
355 let id = f(ctx)?;
356 M::after_insert(ctx, &id)?;
357 M::after_save(ctx, &id)?;
358 M::after_write(ctx, &id)?;
359 Ok(id)
360 }
361
362 pub fn update<M, F>(ctx: &mut HookContext, id: &M::PrimaryKey, f: F) -> HookResult<()>
364 where
365 M: Hookable,
366 F: FnOnce(&mut HookContext) -> HookResult<()>,
367 {
368 M::before_write(ctx)?;
369 M::before_save(ctx)?;
370 M::before_validate(ctx)?;
371 M::validate(ctx)?;
372 M::after_validate(ctx)?;
373 M::before_update(ctx, id)?;
374 f(ctx)?;
375 M::after_update(ctx, id)?;
376 M::after_save(ctx, id)?;
377 M::after_write(ctx, id)?;
378 Ok(())
379 }
380
381 pub fn delete<M, F>(ctx: &mut HookContext, id: &M::PrimaryKey, f: F) -> HookResult<()>
383 where
384 M: Hookable,
385 F: FnOnce(&mut HookContext) -> HookResult<()>,
386 {
387 M::before_delete(ctx, id)?;
388 f(ctx)?;
389 M::after_delete(ctx, id)?;
390 Ok(())
391 }
392
393 pub fn restore<M, F>(ctx: &mut HookContext, id: &M::PrimaryKey, f: F) -> HookResult<()>
395 where
396 M: Hookable,
397 F: FnOnce(&mut HookContext) -> HookResult<()>,
398 {
399 M::before_restore(ctx, id)?;
400 f(ctx)?;
401 M::after_restore(ctx, id)?;
402 Ok(())
403 }
404
405 pub fn find<M, F>(ctx: &mut HookContext, id: &M::PrimaryKey, f: F) -> HookResult<()>
410 where
411 M: Hookable,
412 F: FnOnce(&mut HookContext) -> HookResult<()>,
413 {
414 M::before_find(ctx, id)?;
415 f(ctx)?;
416 M::after_find(ctx, id)?;
417 Ok(())
418 }
419
420 pub fn validate<M>(ctx: &mut HookContext) -> HookResult<()>
425 where
426 M: Hookable,
427 {
428 M::before_validate(ctx)?;
429 M::validate(ctx)?;
430 M::after_validate(ctx)?;
431 Ok(())
432 }
433}
434
435pub trait SoftDelete: crate::model::Model {
444 fn soft_delete_field() -> &'static str;
446
447 fn is_deleted(&self) -> bool;
449}
450
451pub trait GlobalScope {
463 fn scope_name() -> &'static str;
465
466 fn apply_scope(ctx: &HookContext) -> Option<(String, Vec<crate::value::Value>)>;
471}
472
473pub struct SoftDeleteScope;
482
483impl<M: SoftDelete> GlobalScope for (SoftDeleteScope, M) {
484 fn scope_name() -> &'static str {
485 "soft_delete"
486 }
487
488 fn apply_scope(_ctx: &HookContext) -> Option<(String, Vec<crate::value::Value>)> {
489 let field = <M as SoftDelete>::soft_delete_field();
491 Some((format!("{} IS NULL", field), vec![]))
492 }
493}
494
495pub struct TenantScope;
504
505pub trait TenantModel: crate::model::Model {
509 fn tenant_field() -> &'static str {
511 "tenant_id"
512 }
513
514 fn tenant_id(&self) -> i64;
516
517 fn set_tenant_id(&mut self, tenant_id: i64);
519}
520
521impl<M: TenantModel> GlobalScope for (TenantScope, M) {
522 fn scope_name() -> &'static str {
523 "tenant"
524 }
525
526 fn apply_scope(ctx: &HookContext) -> Option<(String, Vec<crate::value::Value>)> {
527 ctx.tenant_id.map(|tid| {
528 (
529 format!("{} = ?", <M as TenantModel>::tenant_field()),
530 vec![crate::value::Value::I64(tid)],
531 )
532 })
533 }
534}
535
536pub type HookFn = Arc<dyn Fn(&HookContext) -> HookResult<()> + Send + Sync>;
542
543pub struct HookRegistry {
548 hooks: RwLock<HashMap<HookEvent, Vec<HookFn>>>,
549}
550
551impl Default for HookRegistry {
552 fn default() -> Self {
553 Self::new()
554 }
555}
556
557impl HookRegistry {
558 pub fn new() -> Self {
560 Self {
561 hooks: RwLock::new(HashMap::new()),
562 }
563 }
564
565 pub fn register(&self, event: HookEvent, hook: HookFn) {
567 if let Ok(mut hooks) = self.hooks.write() {
569 hooks.entry(event).or_default().push(hook);
570 }
571 }
572
573 pub fn dispatch(&self, event: HookEvent, ctx: &HookContext) -> HookResult<()> {
577 let hooks = match self.hooks.read() {
578 Ok(h) => h,
579 Err(_) => return Ok(()), };
581 if let Some(fns) = hooks.get(&event) {
582 for f in fns {
583 f(ctx)?;
584 }
585 }
586 Ok(())
587 }
588
589 pub fn clear(&self, event: HookEvent) {
591 if let Ok(mut hooks) = self.hooks.write() {
592 hooks.remove(&event);
593 }
594 }
595
596 pub fn clear_all(&self) {
598 if let Ok(mut hooks) = self.hooks.write() {
599 hooks.clear();
600 }
601 }
602
603 pub fn count(&self, event: HookEvent) -> usize {
605 self.hooks
606 .read()
607 .map(|h| h.get(&event).map(|v| v.len()).unwrap_or(0))
608 .unwrap_or(0)
609 }
610}
611
612pub struct ScopeRegistry {
621 disabled: RwLock<Vec<String>>,
622}
623
624impl Default for ScopeRegistry {
625 fn default() -> Self {
626 Self::new()
627 }
628}
629
630impl ScopeRegistry {
631 pub fn new() -> Self {
633 Self {
634 disabled: RwLock::new(Vec::new()),
635 }
636 }
637
638 pub fn disable(&self, scope_name: impl Into<String>) {
640 if let Ok(mut disabled) = self.disabled.write() {
641 let name = scope_name.into();
642 if !disabled.contains(&name) {
643 disabled.push(name);
644 }
645 }
646 }
647
648 pub fn enable(&self, scope_name: &str) {
650 if let Ok(mut disabled) = self.disabled.write() {
651 disabled.retain(|n| n != scope_name);
652 }
653 }
654
655 pub fn is_enabled(&self, scope_name: &str) -> bool {
657 self.disabled
658 .read()
659 .map(|d| !d.iter().any(|n| n == scope_name))
660 .unwrap_or(true)
661 }
662
663 pub fn without_scope<F, R>(&self, scope_name: &str, f: F) -> R
673 where
674 F: FnOnce() -> R,
675 {
676 self.disable(scope_name);
677 let result = f();
678 self.enable(scope_name);
679 result
680 }
681}
682
683#[cfg(feature = "composable-plugin")]
688mod extension_point {
689 use std::collections::HashMap;
690 use std::sync::Arc;
691
692 use parking_lot::RwLock;
693
694 use super::{HookContext, HookResult};
695
696 #[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
700 pub enum ExtensionPoint {
701 BeforeConnect,
703 AfterSqlGen,
705 BeforeResultMap,
707 BeforeInsert,
709 BeforeUpdate,
711 BeforeDelete,
713 AfterInsert,
715 AfterUpdate,
717 AfterDelete,
719 }
720
721 pub trait ExtensionHandler: Send + Sync {
723 fn name(&self) -> &str;
725
726 fn handle(&self, ctx: &mut HookContext) -> HookResult<()>;
728 }
729
730 pub struct ExtensionPointRegistry {
734 handlers: RwLock<HashMap<ExtensionPoint, Vec<Arc<dyn ExtensionHandler>>>>,
735 }
736
737 impl Default for ExtensionPointRegistry {
738 fn default() -> Self {
739 Self::new()
740 }
741 }
742
743 impl ExtensionPointRegistry {
744 pub fn new() -> Self {
746 Self {
747 handlers: RwLock::new(HashMap::new()),
748 }
749 }
750
751 pub fn register(&self, point: ExtensionPoint, handler: Arc<dyn ExtensionHandler>) {
753 self.handlers
754 .write()
755 .entry(point)
756 .or_default()
757 .push(handler);
758 }
759
760 pub fn trigger(&self, point: ExtensionPoint, ctx: &mut HookContext) -> HookResult<()> {
762 let handlers = self.handlers.read();
763 if let Some(fns) = handlers.get(&point) {
764 let fns = fns.clone();
765 drop(handlers);
766 for f in &fns {
767 f.handle(ctx)?;
768 }
769 }
770 Ok(())
771 }
772
773 pub fn count(&self, point: ExtensionPoint) -> usize {
775 self.handlers
776 .read()
777 .get(&point)
778 .map(|v| v.len())
779 .unwrap_or(0)
780 }
781
782 pub fn clear(&self, point: ExtensionPoint) {
784 self.handlers.write().remove(&point);
785 }
786 }
787
788 #[cfg(test)]
789 mod tests {
790 use super::*;
791 use std::sync::atomic::{AtomicU32, Ordering};
792
793 struct CounterHandler {
794 name: String,
795 counter: Arc<AtomicU32>,
796 }
797
798 impl ExtensionHandler for CounterHandler {
799 fn name(&self) -> &str {
800 &self.name
801 }
802 fn handle(&self, _ctx: &mut HookContext) -> HookResult<()> {
803 self.counter.fetch_add(1, Ordering::SeqCst);
804 Ok(())
805 }
806 }
807
808 #[test]
809 fn extension_point_register_and_trigger() {
810 let reg = ExtensionPointRegistry::new();
811 let counter = Arc::new(AtomicU32::new(0));
812
813 reg.register(
814 ExtensionPoint::BeforeInsert,
815 Arc::new(CounterHandler {
816 name: "h1".into(),
817 counter: counter.clone(),
818 }),
819 );
820
821 assert_eq!(reg.count(ExtensionPoint::BeforeInsert), 1);
822 let mut ctx = HookContext::new();
823 reg.trigger(ExtensionPoint::BeforeInsert, &mut ctx).unwrap();
824 assert_eq!(counter.load(Ordering::SeqCst), 1);
825 }
826
827 #[test]
828 fn extension_point_multiple_handlers_ordered() {
829 let reg = ExtensionPointRegistry::new();
830 let c1 = Arc::new(AtomicU32::new(0));
831 let c2 = Arc::new(AtomicU32::new(0));
832
833 reg.register(
834 ExtensionPoint::AfterUpdate,
835 Arc::new(CounterHandler {
836 name: "first".into(),
837 counter: c1.clone(),
838 }),
839 );
840 reg.register(
841 ExtensionPoint::AfterUpdate,
842 Arc::new(CounterHandler {
843 name: "second".into(),
844 counter: c2.clone(),
845 }),
846 );
847
848 let mut ctx = HookContext::new();
849 reg.trigger(ExtensionPoint::AfterUpdate, &mut ctx).unwrap();
850 assert_eq!(c1.load(Ordering::SeqCst), 1);
851 assert_eq!(c2.load(Ordering::SeqCst), 1);
852 }
853
854 #[test]
855 fn extension_point_unregistered_returns_ok() {
856 let reg = ExtensionPointRegistry::new();
857 let mut ctx = HookContext::new();
858 assert!(reg.trigger(ExtensionPoint::BeforeConnect, &mut ctx).is_ok());
859 }
860
861 #[test]
862 fn extension_point_clear() {
863 let reg = ExtensionPointRegistry::new();
864 let counter = Arc::new(AtomicU32::new(0));
865 reg.register(
866 ExtensionPoint::BeforeDelete,
867 Arc::new(CounterHandler {
868 name: "h".into(),
869 counter,
870 }),
871 );
872 assert_eq!(reg.count(ExtensionPoint::BeforeDelete), 1);
873 reg.clear(ExtensionPoint::BeforeDelete);
874 assert_eq!(reg.count(ExtensionPoint::BeforeDelete), 0);
875 }
876 }
877}
878
879#[cfg(feature = "composable-plugin")]
880pub use extension_point::{ExtensionHandler, ExtensionPoint, ExtensionPointRegistry};
881
882#[cfg(test)]
887mod tests {
888 use super::*;
889
890 #[test]
891 fn hook_context_builder() {
892 let ctx = HookContext::new()
893 .with_tenant(42)
894 .with_operator(1)
895 .with_timestamp(1700000000);
896
897 assert_eq!(ctx.tenant_id, Some(42));
898 assert_eq!(ctx.operator_id, Some(1));
899 assert_eq!(ctx.timestamp, 1700000000);
900 }
901
902 #[test]
903 fn hook_context_metadata() {
904 let mut ctx = HookContext::new();
905 ctx.set_meta("source", "api");
906 ctx.set_meta("ip", "127.0.0.1");
907
908 assert_eq!(ctx.get_meta("source"), Some(&"api".to_string()));
909 assert_eq!(ctx.get_meta("ip"), Some(&"127.0.0.1".to_string()));
910 assert_eq!(ctx.get_meta("missing"), None);
911 }
912
913 #[test]
914 fn hook_event_is_before_after() {
915 assert!(HookEvent::BeforeInsert.is_before());
916 assert!(!HookEvent::BeforeInsert.is_after());
917 assert!(HookEvent::AfterInsert.is_after());
918 assert!(!HookEvent::AfterInsert.is_before());
919 }
920
921 #[test]
922 fn hook_registry_register_and_dispatch() {
923 let registry = HookRegistry::new();
924 let counter = Arc::new(std::sync::atomic::AtomicU32::new(0));
925
926 let c = Arc::clone(&counter);
927 registry.register(
928 HookEvent::BeforeInsert,
929 Arc::new(move |_ctx| {
930 c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
931 Ok(())
932 }),
933 );
934
935 let ctx = HookContext::new();
936 registry.dispatch(HookEvent::BeforeInsert, &ctx).unwrap();
937 registry.dispatch(HookEvent::BeforeInsert, &ctx).unwrap();
938
939 assert_eq!(counter.load(std::sync::atomic::Ordering::SeqCst), 2);
940 }
941
942 #[test]
943 fn hook_registry_dispatch_no_hooks() {
944 let registry = HookRegistry::new();
945 let ctx = HookContext::new();
946 assert!(registry.dispatch(HookEvent::BeforeInsert, &ctx).is_ok());
948 }
949
950 #[test]
951 fn hook_registry_clear() {
952 let registry = HookRegistry::new();
953 registry.register(HookEvent::BeforeInsert, Arc::new(|_ctx| Ok(())));
954 assert_eq!(registry.count(HookEvent::BeforeInsert), 1);
955
956 registry.clear(HookEvent::BeforeInsert);
957 assert_eq!(registry.count(HookEvent::BeforeInsert), 0);
958 }
959
960 #[test]
961 fn hook_registry_clear_all() {
962 let registry = HookRegistry::new();
963 registry.register(HookEvent::BeforeInsert, Arc::new(|_ctx| Ok(())));
964 registry.register(HookEvent::AfterInsert, Arc::new(|_ctx| Ok(())));
965 registry.register(HookEvent::BeforeUpdate, Arc::new(|_ctx| Ok(())));
966
967 registry.clear_all();
968 assert_eq!(registry.count(HookEvent::BeforeInsert), 0);
969 assert_eq!(registry.count(HookEvent::AfterInsert), 0);
970 assert_eq!(registry.count(HookEvent::BeforeUpdate), 0);
971 }
972
973 #[test]
974 fn scope_registry_enable_disable() {
975 let registry = ScopeRegistry::new();
976
977 assert!(registry.is_enabled("soft_delete"));
978 assert!(registry.is_enabled("tenant"));
979
980 registry.disable("soft_delete");
981 assert!(!registry.is_enabled("soft_delete"));
982 assert!(registry.is_enabled("tenant"));
983
984 registry.enable("soft_delete");
985 assert!(registry.is_enabled("soft_delete"));
986 }
987
988 #[test]
989 fn scope_registry_without_scope() {
990 let registry = ScopeRegistry::new();
991 assert!(registry.is_enabled("soft_delete"));
992
993 let result = registry.without_scope("soft_delete", || {
994 assert!(!registry.is_enabled("soft_delete"));
995 42
996 });
997
998 assert_eq!(result, 42);
999 assert!(registry.is_enabled("soft_delete"));
1000 }
1001
1002 #[test]
1003 fn hook_registry_short_circuit_on_error() {
1004 let registry = HookRegistry::new();
1005 let called = Arc::new(std::sync::atomic::AtomicU32::new(0));
1006
1007 let c1 = Arc::clone(&called);
1008 registry.register(
1009 HookEvent::BeforeInsert,
1010 Arc::new(move |_ctx| {
1011 c1.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1012 Ok(())
1013 }),
1014 );
1015
1016 registry.register(
1017 HookEvent::BeforeInsert,
1018 Arc::new(|_ctx| Err(DbError::Hook("second hook failed".into()))),
1019 );
1020
1021 let c3 = Arc::clone(&called);
1022 registry.register(
1023 HookEvent::BeforeInsert,
1024 Arc::new(move |_ctx| {
1025 c3.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1026 Ok(())
1027 }),
1028 );
1029
1030 let ctx = HookContext::new();
1031 let result = registry.dispatch(HookEvent::BeforeInsert, &ctx);
1032
1033 assert!(result.is_err());
1034 assert_eq!(called.load(std::sync::atomic::Ordering::SeqCst), 1);
1036 }
1037
1038 #[test]
1041 fn hook_event_is_write_level() {
1042 assert!(HookEvent::BeforeWrite.is_write_level());
1043 assert!(HookEvent::AfterWrite.is_write_level());
1044 assert!(HookEvent::BeforeSave.is_write_level());
1045 assert!(HookEvent::AfterSave.is_write_level());
1046 assert!(!HookEvent::BeforeInsert.is_write_level());
1047 assert!(!HookEvent::AfterDelete.is_write_level());
1048 assert!(!HookEvent::BeforeRestore.is_write_level());
1049 }
1050
1051 #[test]
1052 fn hook_event_before_after_covers_new_variants() {
1053 assert!(HookEvent::BeforeWrite.is_before());
1054 assert!(HookEvent::BeforeSave.is_before());
1055 assert!(HookEvent::BeforeRestore.is_before());
1056 assert!(HookEvent::AfterWrite.is_after());
1057 assert!(HookEvent::AfterSave.is_after());
1058 assert!(HookEvent::AfterRestore.is_after());
1059 assert!(!HookEvent::AfterWrite.is_before());
1060 assert!(!HookEvent::BeforeWrite.is_after());
1061 }
1062
1063 #[test]
1064 fn hook_registry_supports_new_events() {
1065 let registry = HookRegistry::new();
1066 let counter = Arc::new(std::sync::atomic::AtomicU32::new(0));
1067
1068 for event in [
1069 HookEvent::BeforeWrite,
1070 HookEvent::AfterWrite,
1071 HookEvent::BeforeSave,
1072 HookEvent::AfterSave,
1073 HookEvent::BeforeRestore,
1074 HookEvent::AfterRestore,
1075 ] {
1076 let c = Arc::clone(&counter);
1077 registry.register(
1078 event,
1079 Arc::new(move |_ctx| {
1080 c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1081 Ok(())
1082 }),
1083 );
1084 }
1085
1086 let ctx = HookContext::new();
1087 for event in [
1088 HookEvent::BeforeWrite,
1089 HookEvent::AfterWrite,
1090 HookEvent::BeforeSave,
1091 HookEvent::AfterSave,
1092 HookEvent::BeforeRestore,
1093 HookEvent::AfterRestore,
1094 ] {
1095 registry.dispatch(event, &ctx).unwrap();
1096 }
1097
1098 assert_eq!(
1099 counter.load(std::sync::atomic::Ordering::SeqCst),
1100 6,
1101 "所有细粒度事件均应被正确注册与触发"
1102 );
1103 }
1104
1105 struct DispatchTestModel;
1108 impl crate::model::Model for DispatchTestModel {
1109 type PrimaryKey = i64;
1110 fn table_name() -> &'static str {
1111 "dispatch_test"
1112 }
1113 fn pk(&self) -> Self::PrimaryKey {
1114 0
1115 }
1116 fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
1117 }
1118
1119 static DISPATCH_CALLS: std::sync::OnceLock<Arc<std::sync::atomic::AtomicU32>> =
1121 std::sync::OnceLock::new();
1122
1123 fn dispatch_calls() -> Arc<std::sync::atomic::AtomicU32> {
1124 DISPATCH_CALLS
1125 .get_or_init(|| Arc::new(std::sync::atomic::AtomicU32::new(0)))
1126 .clone()
1127 }
1128
1129 impl Hookable for DispatchTestModel {
1130 fn before_write(ctx: &mut HookContext) -> HookResult<()> {
1131 ctx.set_meta("before_write", "1");
1132 Ok(())
1133 }
1134 fn before_save(ctx: &mut HookContext) -> HookResult<()> {
1135 ctx.set_meta("before_save", "1");
1136 Ok(())
1137 }
1138 fn before_validate(ctx: &mut HookContext) -> HookResult<()> {
1139 ctx.set_meta("before_validate", "1");
1140 Ok(())
1141 }
1142 fn after_validate(ctx: &HookContext) -> HookResult<()> {
1143 assert_eq!(ctx.get_meta("before_validate"), Some(&"1".to_string()));
1144 ctx_set_meta_for_after(ctx, "after_validate", "1");
1145 Ok(())
1146 }
1147 fn before_insert(ctx: &mut HookContext) -> HookResult<()> {
1148 ctx.set_meta("before_insert", "1");
1149 Ok(())
1150 }
1151 fn after_insert(ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
1152 assert_eq!(ctx.get_meta("before_write"), Some(&"1".to_string()));
1153 assert_eq!(ctx.get_meta("before_save"), Some(&"1".to_string()));
1154 assert_eq!(ctx.get_meta("before_insert"), Some(&"1".to_string()));
1155 assert_eq!(ctx.get_meta("before_validate"), Some(&"1".to_string()));
1156 Ok(())
1157 }
1158 fn after_save(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
1159 dispatch_calls().fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1160 Ok(())
1161 }
1162 fn after_write(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
1163 dispatch_calls().fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1164 Ok(())
1165 }
1166 fn before_find(ctx: &mut HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
1167 ctx.set_meta("before_find", "1");
1168 Ok(())
1169 }
1170 fn after_find(ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
1171 assert_eq!(ctx.get_meta("before_find"), Some(&"1".to_string()));
1172 ctx_set_meta_for_after(ctx, "after_find", "1");
1173 Ok(())
1174 }
1175 }
1176
1177 static AFTER_VALIDATE_COUNT: std::sync::atomic::AtomicU32 =
1180 std::sync::atomic::AtomicU32::new(0);
1181 static AFTER_FIND_COUNT: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
1182
1183 fn ctx_set_meta_for_after(_ctx: &HookContext, key: &str, _value: &str) {
1184 match key {
1185 "after_validate" => {
1186 AFTER_VALIDATE_COUNT.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1187 }
1188 "after_find" => {
1189 AFTER_FIND_COUNT.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1190 }
1191 _ => {}
1192 }
1193 }
1194
1195 fn after_call_was(key: &str) -> bool {
1196 match key {
1197 "after_validate" => AFTER_VALIDATE_COUNT.load(std::sync::atomic::Ordering::SeqCst) > 0,
1198 "after_find" => AFTER_FIND_COUNT.load(std::sync::atomic::Ordering::SeqCst) > 0,
1199 _ => false,
1200 }
1201 }
1202
1203 fn reset_after_calls() {
1204 AFTER_VALIDATE_COUNT.store(0, std::sync::atomic::Ordering::SeqCst);
1205 AFTER_FIND_COUNT.store(0, std::sync::atomic::Ordering::SeqCst);
1206 }
1207
1208 static HOOK_TEST_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
1210
1211 #[test]
1212 fn hook_dispatcher_insert_full_sequence() {
1213 let _guard = HOOK_TEST_LOCK.lock().unwrap();
1214 dispatch_calls().store(0, std::sync::atomic::Ordering::SeqCst);
1215 reset_after_calls();
1216 let mut ctx = HookContext::new();
1217 let id = HookDispatcher::insert::<DispatchTestModel, _>(&mut ctx, |_ctx| Ok(42_i64));
1218 assert!(id.is_ok());
1219 assert_eq!(id.unwrap(), 42);
1220 assert_eq!(ctx.get_meta("before_write"), Some(&"1".to_string()));
1222 assert_eq!(ctx.get_meta("before_save"), Some(&"1".to_string()));
1223 assert_eq!(ctx.get_meta("before_insert"), Some(&"1".to_string()));
1224 assert_eq!(ctx.get_meta("before_validate"), Some(&"1".to_string()));
1226 assert!(after_call_was("after_validate"));
1227 assert_eq!(
1229 dispatch_calls().load(std::sync::atomic::Ordering::SeqCst),
1230 2
1231 );
1232 }
1233
1234 #[test]
1235 fn hook_dispatcher_insert_short_circuit_on_before_write_error() {
1236 struct ErrorModel;
1237 impl crate::model::Model for ErrorModel {
1238 type PrimaryKey = i64;
1239 fn table_name() -> &'static str {
1240 "error_model"
1241 }
1242 fn pk(&self) -> Self::PrimaryKey {
1243 0
1244 }
1245 fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
1246 }
1247 impl Hookable for ErrorModel {
1248 fn before_write(_ctx: &mut HookContext) -> HookResult<()> {
1249 Err(DbError::Hook("before_write failed".into()))
1250 }
1251 }
1252
1253 let mut ctx = HookContext::new();
1254 let result = HookDispatcher::insert::<ErrorModel, _>(&mut ctx, |_ctx| Ok(1_i64));
1255 assert!(result.is_err());
1256 }
1258
1259 #[test]
1260 fn hook_dispatcher_insert_short_circuit_on_before_validate_error() {
1261 struct ValidationFailModel;
1262 impl crate::model::Model for ValidationFailModel {
1263 type PrimaryKey = i64;
1264 fn table_name() -> &'static str {
1265 "validation_fail"
1266 }
1267 fn pk(&self) -> Self::PrimaryKey {
1268 0
1269 }
1270 fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
1271 }
1272 impl Hookable for ValidationFailModel {
1273 fn before_validate(_ctx: &mut HookContext) -> HookResult<()> {
1274 Err(DbError::Validation("name is required".into()))
1275 }
1276 }
1277
1278 let mut ctx = HookContext::new();
1279 let called = Arc::new(std::sync::atomic::AtomicU32::new(0));
1280 let c = Arc::clone(&called);
1281 let result = HookDispatcher::insert::<ValidationFailModel, _>(&mut ctx, move |_ctx| {
1282 c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1283 Ok(1_i64)
1284 });
1285 assert!(result.is_err());
1286 assert_eq!(
1288 called.load(std::sync::atomic::Ordering::SeqCst),
1289 0,
1290 "before_validate 失败应短路 INSERT 操作"
1291 );
1292 match result.unwrap_err() {
1294 DbError::Validation(msg) => assert_eq!(msg, "name is required"),
1295 other => panic!("期望 Validation 错误,得到 {:?}", other),
1296 }
1297 }
1298
1299 #[test]
1300 fn hook_dispatcher_update_full_sequence() {
1301 let _guard = HOOK_TEST_LOCK.lock().unwrap();
1302 dispatch_calls().store(0, std::sync::atomic::Ordering::SeqCst);
1303 reset_after_calls();
1304 let mut ctx = HookContext::new();
1305 let result =
1306 HookDispatcher::update::<DispatchTestModel, _>(&mut ctx, &42_i64, |_ctx| Ok(()));
1307 assert!(result.is_ok());
1308 assert_eq!(
1310 dispatch_calls().load(std::sync::atomic::Ordering::SeqCst),
1311 2
1312 );
1313 assert_eq!(ctx.get_meta("before_validate"), Some(&"1".to_string()));
1315 assert!(after_call_was("after_validate"));
1316 }
1317
1318 #[test]
1319 fn hook_dispatcher_delete_full_sequence() {
1320 let mut ctx = HookContext::new();
1321 let result =
1322 HookDispatcher::delete::<DispatchTestModel, _>(&mut ctx, &42_i64, |_ctx| Ok(()));
1323 assert!(result.is_ok());
1324 }
1325
1326 #[test]
1327 fn hook_dispatcher_restore_full_sequence() {
1328 let mut ctx = HookContext::new();
1329 let result =
1330 HookDispatcher::restore::<DispatchTestModel, _>(&mut ctx, &42_i64, |_ctx| Ok(()));
1331 assert!(result.is_ok());
1332 }
1333
1334 #[test]
1335 fn hook_dispatcher_find_full_sequence() {
1336 let _guard = HOOK_TEST_LOCK.lock().unwrap();
1337 dispatch_calls().store(0, std::sync::atomic::Ordering::SeqCst);
1338 reset_after_calls();
1339 let mut ctx = HookContext::new();
1340 let called = Arc::new(std::sync::atomic::AtomicU32::new(0));
1341 let c = Arc::clone(&called);
1342 let result = HookDispatcher::find::<DispatchTestModel, _>(&mut ctx, &42_i64, move |_ctx| {
1343 c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1344 Ok(())
1345 });
1346 assert!(result.is_ok());
1347 assert_eq!(
1348 called.load(std::sync::atomic::Ordering::SeqCst),
1349 1,
1350 "SELECT 操作应执行一次"
1351 );
1352 assert_eq!(ctx.get_meta("before_find"), Some(&"1".to_string()));
1354 assert!(after_call_was("after_find"));
1355 }
1356
1357 #[test]
1358 fn hook_dispatcher_find_short_circuit_on_before_find_error() {
1359 struct FindFailModel;
1360 impl crate::model::Model for FindFailModel {
1361 type PrimaryKey = i64;
1362 fn table_name() -> &'static str {
1363 "find_fail"
1364 }
1365 fn pk(&self) -> Self::PrimaryKey {
1366 0
1367 }
1368 fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
1369 }
1370 impl Hookable for FindFailModel {
1371 fn before_find(_ctx: &mut HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
1372 Err(DbError::Hook("before_find blocked".into()))
1373 }
1374 }
1375
1376 let mut ctx = HookContext::new();
1377 let called = Arc::new(std::sync::atomic::AtomicU32::new(0));
1378 let c = Arc::clone(&called);
1379 let result = HookDispatcher::find::<FindFailModel, _>(&mut ctx, &1_i64, move |_ctx| {
1380 c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1381 Ok(())
1382 });
1383 assert!(result.is_err());
1384 assert_eq!(
1385 called.load(std::sync::atomic::Ordering::SeqCst),
1386 0,
1387 "before_find 失败应短路 SELECT"
1388 );
1389 }
1390
1391 #[test]
1392 fn hook_dispatcher_validate_standalone() {
1393 reset_after_calls();
1394 let mut ctx = HookContext::new();
1395 let result = HookDispatcher::validate::<DispatchTestModel>(&mut ctx);
1396 assert!(result.is_ok());
1397 assert_eq!(ctx.get_meta("before_validate"), Some(&"1".to_string()));
1398 assert!(after_call_was("after_validate"));
1399 }
1400
1401 #[test]
1402 fn hook_event_is_find_level_and_is_validate_level() {
1403 assert!(HookEvent::BeforeFind.is_find_level());
1404 assert!(HookEvent::AfterFind.is_find_level());
1405 assert!(HookEvent::BeforeValidate.is_validate_level());
1406 assert!(HookEvent::AfterValidate.is_validate_level());
1407 assert!(!HookEvent::BeforeInsert.is_find_level());
1408 assert!(!HookEvent::BeforeInsert.is_validate_level());
1409 assert!(!HookEvent::BeforeWrite.is_find_level());
1410 assert!(!HookEvent::BeforeWrite.is_validate_level());
1411 }
1412
1413 #[test]
1414 fn hook_event_is_fine_grained_covers_all_v02_events() {
1415 assert!(HookEvent::BeforeWrite.is_fine_grained());
1417 assert!(HookEvent::AfterWrite.is_fine_grained());
1418 assert!(HookEvent::BeforeSave.is_fine_grained());
1419 assert!(HookEvent::AfterSave.is_fine_grained());
1420 assert!(HookEvent::BeforeRestore.is_fine_grained());
1421 assert!(HookEvent::AfterRestore.is_fine_grained());
1422 assert!(HookEvent::BeforeFind.is_fine_grained());
1423 assert!(HookEvent::AfterFind.is_fine_grained());
1424 assert!(HookEvent::BeforeValidate.is_fine_grained());
1425 assert!(HookEvent::AfterValidate.is_fine_grained());
1426 assert!(!HookEvent::BeforeInsert.is_fine_grained());
1428 assert!(!HookEvent::AfterInsert.is_fine_grained());
1429 assert!(!HookEvent::BeforeUpdate.is_fine_grained());
1430 assert!(!HookEvent::AfterUpdate.is_fine_grained());
1431 assert!(!HookEvent::BeforeDelete.is_fine_grained());
1432 assert!(!HookEvent::AfterDelete.is_fine_grained());
1433 }
1434
1435 #[test]
1436 fn hook_registry_supports_find_and_validate_events() {
1437 let registry = HookRegistry::new();
1438 let counter = Arc::new(std::sync::atomic::AtomicU32::new(0));
1439
1440 for event in [
1441 HookEvent::BeforeFind,
1442 HookEvent::AfterFind,
1443 HookEvent::BeforeValidate,
1444 HookEvent::AfterValidate,
1445 ] {
1446 let c = Arc::clone(&counter);
1447 registry.register(
1448 event,
1449 Arc::new(move |_ctx| {
1450 c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1451 Ok(())
1452 }),
1453 );
1454 }
1455
1456 let ctx = HookContext::new();
1457 for event in [
1458 HookEvent::BeforeFind,
1459 HookEvent::AfterFind,
1460 HookEvent::BeforeValidate,
1461 HookEvent::AfterValidate,
1462 ] {
1463 registry.dispatch(event, &ctx).unwrap();
1464 }
1465
1466 assert_eq!(
1467 counter.load(std::sync::atomic::Ordering::SeqCst),
1468 4,
1469 "find/validate 钩子应能被注册与触发"
1470 );
1471 }
1472
1473 #[test]
1474 fn db_error_validation_error_code_and_display() {
1475 let err = DbError::Validation("name required".into());
1476 assert_eq!(err.error_code(), "DB021");
1477 assert_eq!(format!("{}", err), "Validation error: name required");
1478 }
1479}