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(test)]
688mod tests {
689 use super::*;
690
691 #[test]
692 fn hook_context_builder() {
693 let ctx = HookContext::new()
694 .with_tenant(42)
695 .with_operator(1)
696 .with_timestamp(1700000000);
697
698 assert_eq!(ctx.tenant_id, Some(42));
699 assert_eq!(ctx.operator_id, Some(1));
700 assert_eq!(ctx.timestamp, 1700000000);
701 }
702
703 #[test]
704 fn hook_context_metadata() {
705 let mut ctx = HookContext::new();
706 ctx.set_meta("source", "api");
707 ctx.set_meta("ip", "127.0.0.1");
708
709 assert_eq!(ctx.get_meta("source"), Some(&"api".to_string()));
710 assert_eq!(ctx.get_meta("ip"), Some(&"127.0.0.1".to_string()));
711 assert_eq!(ctx.get_meta("missing"), None);
712 }
713
714 #[test]
715 fn hook_event_is_before_after() {
716 assert!(HookEvent::BeforeInsert.is_before());
717 assert!(!HookEvent::BeforeInsert.is_after());
718 assert!(HookEvent::AfterInsert.is_after());
719 assert!(!HookEvent::AfterInsert.is_before());
720 }
721
722 #[test]
723 fn hook_registry_register_and_dispatch() {
724 let registry = HookRegistry::new();
725 let counter = Arc::new(std::sync::atomic::AtomicU32::new(0));
726
727 let c = Arc::clone(&counter);
728 registry.register(
729 HookEvent::BeforeInsert,
730 Arc::new(move |_ctx| {
731 c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
732 Ok(())
733 }),
734 );
735
736 let ctx = HookContext::new();
737 registry.dispatch(HookEvent::BeforeInsert, &ctx).unwrap();
738 registry.dispatch(HookEvent::BeforeInsert, &ctx).unwrap();
739
740 assert_eq!(counter.load(std::sync::atomic::Ordering::SeqCst), 2);
741 }
742
743 #[test]
744 fn hook_registry_dispatch_no_hooks() {
745 let registry = HookRegistry::new();
746 let ctx = HookContext::new();
747 assert!(registry.dispatch(HookEvent::BeforeInsert, &ctx).is_ok());
749 }
750
751 #[test]
752 fn hook_registry_clear() {
753 let registry = HookRegistry::new();
754 registry.register(HookEvent::BeforeInsert, Arc::new(|_ctx| Ok(())));
755 assert_eq!(registry.count(HookEvent::BeforeInsert), 1);
756
757 registry.clear(HookEvent::BeforeInsert);
758 assert_eq!(registry.count(HookEvent::BeforeInsert), 0);
759 }
760
761 #[test]
762 fn hook_registry_clear_all() {
763 let registry = HookRegistry::new();
764 registry.register(HookEvent::BeforeInsert, Arc::new(|_ctx| Ok(())));
765 registry.register(HookEvent::AfterInsert, Arc::new(|_ctx| Ok(())));
766 registry.register(HookEvent::BeforeUpdate, Arc::new(|_ctx| Ok(())));
767
768 registry.clear_all();
769 assert_eq!(registry.count(HookEvent::BeforeInsert), 0);
770 assert_eq!(registry.count(HookEvent::AfterInsert), 0);
771 assert_eq!(registry.count(HookEvent::BeforeUpdate), 0);
772 }
773
774 #[test]
775 fn scope_registry_enable_disable() {
776 let registry = ScopeRegistry::new();
777
778 assert!(registry.is_enabled("soft_delete"));
779 assert!(registry.is_enabled("tenant"));
780
781 registry.disable("soft_delete");
782 assert!(!registry.is_enabled("soft_delete"));
783 assert!(registry.is_enabled("tenant"));
784
785 registry.enable("soft_delete");
786 assert!(registry.is_enabled("soft_delete"));
787 }
788
789 #[test]
790 fn scope_registry_without_scope() {
791 let registry = ScopeRegistry::new();
792 assert!(registry.is_enabled("soft_delete"));
793
794 let result = registry.without_scope("soft_delete", || {
795 assert!(!registry.is_enabled("soft_delete"));
796 42
797 });
798
799 assert_eq!(result, 42);
800 assert!(registry.is_enabled("soft_delete"));
801 }
802
803 #[test]
804 fn hook_registry_short_circuit_on_error() {
805 let registry = HookRegistry::new();
806 let called = Arc::new(std::sync::atomic::AtomicU32::new(0));
807
808 let c1 = Arc::clone(&called);
809 registry.register(
810 HookEvent::BeforeInsert,
811 Arc::new(move |_ctx| {
812 c1.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
813 Ok(())
814 }),
815 );
816
817 registry.register(
818 HookEvent::BeforeInsert,
819 Arc::new(|_ctx| Err(DbError::Hook("second hook failed".into()))),
820 );
821
822 let c3 = Arc::clone(&called);
823 registry.register(
824 HookEvent::BeforeInsert,
825 Arc::new(move |_ctx| {
826 c3.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
827 Ok(())
828 }),
829 );
830
831 let ctx = HookContext::new();
832 let result = registry.dispatch(HookEvent::BeforeInsert, &ctx);
833
834 assert!(result.is_err());
835 assert_eq!(called.load(std::sync::atomic::Ordering::SeqCst), 1);
837 }
838
839 #[test]
842 fn hook_event_is_write_level() {
843 assert!(HookEvent::BeforeWrite.is_write_level());
844 assert!(HookEvent::AfterWrite.is_write_level());
845 assert!(HookEvent::BeforeSave.is_write_level());
846 assert!(HookEvent::AfterSave.is_write_level());
847 assert!(!HookEvent::BeforeInsert.is_write_level());
848 assert!(!HookEvent::AfterDelete.is_write_level());
849 assert!(!HookEvent::BeforeRestore.is_write_level());
850 }
851
852 #[test]
853 fn hook_event_before_after_covers_new_variants() {
854 assert!(HookEvent::BeforeWrite.is_before());
855 assert!(HookEvent::BeforeSave.is_before());
856 assert!(HookEvent::BeforeRestore.is_before());
857 assert!(HookEvent::AfterWrite.is_after());
858 assert!(HookEvent::AfterSave.is_after());
859 assert!(HookEvent::AfterRestore.is_after());
860 assert!(!HookEvent::AfterWrite.is_before());
861 assert!(!HookEvent::BeforeWrite.is_after());
862 }
863
864 #[test]
865 fn hook_registry_supports_new_events() {
866 let registry = HookRegistry::new();
867 let counter = Arc::new(std::sync::atomic::AtomicU32::new(0));
868
869 for event in [
870 HookEvent::BeforeWrite,
871 HookEvent::AfterWrite,
872 HookEvent::BeforeSave,
873 HookEvent::AfterSave,
874 HookEvent::BeforeRestore,
875 HookEvent::AfterRestore,
876 ] {
877 let c = Arc::clone(&counter);
878 registry.register(
879 event,
880 Arc::new(move |_ctx| {
881 c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
882 Ok(())
883 }),
884 );
885 }
886
887 let ctx = HookContext::new();
888 for event in [
889 HookEvent::BeforeWrite,
890 HookEvent::AfterWrite,
891 HookEvent::BeforeSave,
892 HookEvent::AfterSave,
893 HookEvent::BeforeRestore,
894 HookEvent::AfterRestore,
895 ] {
896 registry.dispatch(event, &ctx).unwrap();
897 }
898
899 assert_eq!(
900 counter.load(std::sync::atomic::Ordering::SeqCst),
901 6,
902 "所有细粒度事件均应被正确注册与触发"
903 );
904 }
905
906 struct DispatchTestModel;
909 impl crate::model::Model for DispatchTestModel {
910 type PrimaryKey = i64;
911 fn table_name() -> &'static str {
912 "dispatch_test"
913 }
914 fn pk(&self) -> Self::PrimaryKey {
915 0
916 }
917 fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
918 }
919
920 static DISPATCH_CALLS: std::sync::OnceLock<Arc<std::sync::atomic::AtomicU32>> =
922 std::sync::OnceLock::new();
923
924 fn dispatch_calls() -> Arc<std::sync::atomic::AtomicU32> {
925 DISPATCH_CALLS
926 .get_or_init(|| Arc::new(std::sync::atomic::AtomicU32::new(0)))
927 .clone()
928 }
929
930 impl Hookable for DispatchTestModel {
931 fn before_write(ctx: &mut HookContext) -> HookResult<()> {
932 ctx.set_meta("before_write", "1");
933 Ok(())
934 }
935 fn before_save(ctx: &mut HookContext) -> HookResult<()> {
936 ctx.set_meta("before_save", "1");
937 Ok(())
938 }
939 fn before_validate(ctx: &mut HookContext) -> HookResult<()> {
940 ctx.set_meta("before_validate", "1");
941 Ok(())
942 }
943 fn after_validate(ctx: &HookContext) -> HookResult<()> {
944 assert_eq!(ctx.get_meta("before_validate"), Some(&"1".to_string()));
945 ctx_set_meta_for_after(ctx, "after_validate", "1");
946 Ok(())
947 }
948 fn before_insert(ctx: &mut HookContext) -> HookResult<()> {
949 ctx.set_meta("before_insert", "1");
950 Ok(())
951 }
952 fn after_insert(ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
953 assert_eq!(ctx.get_meta("before_write"), Some(&"1".to_string()));
954 assert_eq!(ctx.get_meta("before_save"), Some(&"1".to_string()));
955 assert_eq!(ctx.get_meta("before_insert"), Some(&"1".to_string()));
956 assert_eq!(ctx.get_meta("before_validate"), Some(&"1".to_string()));
957 Ok(())
958 }
959 fn after_save(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
960 dispatch_calls().fetch_add(1, std::sync::atomic::Ordering::SeqCst);
961 Ok(())
962 }
963 fn after_write(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
964 dispatch_calls().fetch_add(1, std::sync::atomic::Ordering::SeqCst);
965 Ok(())
966 }
967 fn before_find(ctx: &mut HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
968 ctx.set_meta("before_find", "1");
969 Ok(())
970 }
971 fn after_find(ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
972 assert_eq!(ctx.get_meta("before_find"), Some(&"1".to_string()));
973 ctx_set_meta_for_after(ctx, "after_find", "1");
974 Ok(())
975 }
976 }
977
978 static AFTER_VALIDATE_COUNT: std::sync::atomic::AtomicU32 =
981 std::sync::atomic::AtomicU32::new(0);
982 static AFTER_FIND_COUNT: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
983
984 fn ctx_set_meta_for_after(_ctx: &HookContext, key: &str, _value: &str) {
985 match key {
986 "after_validate" => {
987 AFTER_VALIDATE_COUNT.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
988 }
989 "after_find" => {
990 AFTER_FIND_COUNT.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
991 }
992 _ => {}
993 }
994 }
995
996 fn after_call_was(key: &str) -> bool {
997 match key {
998 "after_validate" => AFTER_VALIDATE_COUNT.load(std::sync::atomic::Ordering::SeqCst) > 0,
999 "after_find" => AFTER_FIND_COUNT.load(std::sync::atomic::Ordering::SeqCst) > 0,
1000 _ => false,
1001 }
1002 }
1003
1004 fn reset_after_calls() {
1005 AFTER_VALIDATE_COUNT.store(0, std::sync::atomic::Ordering::SeqCst);
1006 AFTER_FIND_COUNT.store(0, std::sync::atomic::Ordering::SeqCst);
1007 }
1008
1009 static HOOK_TEST_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
1011
1012 #[test]
1013 fn hook_dispatcher_insert_full_sequence() {
1014 let _guard = HOOK_TEST_LOCK.lock().unwrap();
1015 dispatch_calls().store(0, std::sync::atomic::Ordering::SeqCst);
1016 reset_after_calls();
1017 let mut ctx = HookContext::new();
1018 let id = HookDispatcher::insert::<DispatchTestModel, _>(&mut ctx, |_ctx| Ok(42_i64));
1019 assert!(id.is_ok());
1020 assert_eq!(id.unwrap(), 42);
1021 assert_eq!(ctx.get_meta("before_write"), Some(&"1".to_string()));
1023 assert_eq!(ctx.get_meta("before_save"), Some(&"1".to_string()));
1024 assert_eq!(ctx.get_meta("before_insert"), Some(&"1".to_string()));
1025 assert_eq!(ctx.get_meta("before_validate"), Some(&"1".to_string()));
1027 assert!(after_call_was("after_validate"));
1028 assert_eq!(
1030 dispatch_calls().load(std::sync::atomic::Ordering::SeqCst),
1031 2
1032 );
1033 }
1034
1035 #[test]
1036 fn hook_dispatcher_insert_short_circuit_on_before_write_error() {
1037 struct ErrorModel;
1038 impl crate::model::Model for ErrorModel {
1039 type PrimaryKey = i64;
1040 fn table_name() -> &'static str {
1041 "error_model"
1042 }
1043 fn pk(&self) -> Self::PrimaryKey {
1044 0
1045 }
1046 fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
1047 }
1048 impl Hookable for ErrorModel {
1049 fn before_write(_ctx: &mut HookContext) -> HookResult<()> {
1050 Err(DbError::Hook("before_write failed".into()))
1051 }
1052 }
1053
1054 let mut ctx = HookContext::new();
1055 let result = HookDispatcher::insert::<ErrorModel, _>(&mut ctx, |_ctx| Ok(1_i64));
1056 assert!(result.is_err());
1057 }
1059
1060 #[test]
1061 fn hook_dispatcher_insert_short_circuit_on_before_validate_error() {
1062 struct ValidationFailModel;
1063 impl crate::model::Model for ValidationFailModel {
1064 type PrimaryKey = i64;
1065 fn table_name() -> &'static str {
1066 "validation_fail"
1067 }
1068 fn pk(&self) -> Self::PrimaryKey {
1069 0
1070 }
1071 fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
1072 }
1073 impl Hookable for ValidationFailModel {
1074 fn before_validate(_ctx: &mut HookContext) -> HookResult<()> {
1075 Err(DbError::Validation("name is required".into()))
1076 }
1077 }
1078
1079 let mut ctx = HookContext::new();
1080 let called = Arc::new(std::sync::atomic::AtomicU32::new(0));
1081 let c = Arc::clone(&called);
1082 let result = HookDispatcher::insert::<ValidationFailModel, _>(&mut ctx, move |_ctx| {
1083 c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1084 Ok(1_i64)
1085 });
1086 assert!(result.is_err());
1087 assert_eq!(
1089 called.load(std::sync::atomic::Ordering::SeqCst),
1090 0,
1091 "before_validate 失败应短路 INSERT 操作"
1092 );
1093 match result.unwrap_err() {
1095 DbError::Validation(msg) => assert_eq!(msg, "name is required"),
1096 other => panic!("期望 Validation 错误,得到 {:?}", other),
1097 }
1098 }
1099
1100 #[test]
1101 fn hook_dispatcher_update_full_sequence() {
1102 let _guard = HOOK_TEST_LOCK.lock().unwrap();
1103 dispatch_calls().store(0, std::sync::atomic::Ordering::SeqCst);
1104 reset_after_calls();
1105 let mut ctx = HookContext::new();
1106 let result =
1107 HookDispatcher::update::<DispatchTestModel, _>(&mut ctx, &42_i64, |_ctx| Ok(()));
1108 assert!(result.is_ok());
1109 assert_eq!(
1111 dispatch_calls().load(std::sync::atomic::Ordering::SeqCst),
1112 2
1113 );
1114 assert_eq!(ctx.get_meta("before_validate"), Some(&"1".to_string()));
1116 assert!(after_call_was("after_validate"));
1117 }
1118
1119 #[test]
1120 fn hook_dispatcher_delete_full_sequence() {
1121 let mut ctx = HookContext::new();
1122 let result =
1123 HookDispatcher::delete::<DispatchTestModel, _>(&mut ctx, &42_i64, |_ctx| Ok(()));
1124 assert!(result.is_ok());
1125 }
1126
1127 #[test]
1128 fn hook_dispatcher_restore_full_sequence() {
1129 let mut ctx = HookContext::new();
1130 let result =
1131 HookDispatcher::restore::<DispatchTestModel, _>(&mut ctx, &42_i64, |_ctx| Ok(()));
1132 assert!(result.is_ok());
1133 }
1134
1135 #[test]
1136 fn hook_dispatcher_find_full_sequence() {
1137 let _guard = HOOK_TEST_LOCK.lock().unwrap();
1138 dispatch_calls().store(0, std::sync::atomic::Ordering::SeqCst);
1139 reset_after_calls();
1140 let mut ctx = HookContext::new();
1141 let called = Arc::new(std::sync::atomic::AtomicU32::new(0));
1142 let c = Arc::clone(&called);
1143 let result = HookDispatcher::find::<DispatchTestModel, _>(&mut ctx, &42_i64, move |_ctx| {
1144 c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1145 Ok(())
1146 });
1147 assert!(result.is_ok());
1148 assert_eq!(
1149 called.load(std::sync::atomic::Ordering::SeqCst),
1150 1,
1151 "SELECT 操作应执行一次"
1152 );
1153 assert_eq!(ctx.get_meta("before_find"), Some(&"1".to_string()));
1155 assert!(after_call_was("after_find"));
1156 }
1157
1158 #[test]
1159 fn hook_dispatcher_find_short_circuit_on_before_find_error() {
1160 struct FindFailModel;
1161 impl crate::model::Model for FindFailModel {
1162 type PrimaryKey = i64;
1163 fn table_name() -> &'static str {
1164 "find_fail"
1165 }
1166 fn pk(&self) -> Self::PrimaryKey {
1167 0
1168 }
1169 fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
1170 }
1171 impl Hookable for FindFailModel {
1172 fn before_find(_ctx: &mut HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
1173 Err(DbError::Hook("before_find blocked".into()))
1174 }
1175 }
1176
1177 let mut ctx = HookContext::new();
1178 let called = Arc::new(std::sync::atomic::AtomicU32::new(0));
1179 let c = Arc::clone(&called);
1180 let result = HookDispatcher::find::<FindFailModel, _>(&mut ctx, &1_i64, move |_ctx| {
1181 c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1182 Ok(())
1183 });
1184 assert!(result.is_err());
1185 assert_eq!(
1186 called.load(std::sync::atomic::Ordering::SeqCst),
1187 0,
1188 "before_find 失败应短路 SELECT"
1189 );
1190 }
1191
1192 #[test]
1193 fn hook_dispatcher_validate_standalone() {
1194 reset_after_calls();
1195 let mut ctx = HookContext::new();
1196 let result = HookDispatcher::validate::<DispatchTestModel>(&mut ctx);
1197 assert!(result.is_ok());
1198 assert_eq!(ctx.get_meta("before_validate"), Some(&"1".to_string()));
1199 assert!(after_call_was("after_validate"));
1200 }
1201
1202 #[test]
1203 fn hook_event_is_find_level_and_is_validate_level() {
1204 assert!(HookEvent::BeforeFind.is_find_level());
1205 assert!(HookEvent::AfterFind.is_find_level());
1206 assert!(HookEvent::BeforeValidate.is_validate_level());
1207 assert!(HookEvent::AfterValidate.is_validate_level());
1208 assert!(!HookEvent::BeforeInsert.is_find_level());
1209 assert!(!HookEvent::BeforeInsert.is_validate_level());
1210 assert!(!HookEvent::BeforeWrite.is_find_level());
1211 assert!(!HookEvent::BeforeWrite.is_validate_level());
1212 }
1213
1214 #[test]
1215 fn hook_event_is_fine_grained_covers_all_v02_events() {
1216 assert!(HookEvent::BeforeWrite.is_fine_grained());
1218 assert!(HookEvent::AfterWrite.is_fine_grained());
1219 assert!(HookEvent::BeforeSave.is_fine_grained());
1220 assert!(HookEvent::AfterSave.is_fine_grained());
1221 assert!(HookEvent::BeforeRestore.is_fine_grained());
1222 assert!(HookEvent::AfterRestore.is_fine_grained());
1223 assert!(HookEvent::BeforeFind.is_fine_grained());
1224 assert!(HookEvent::AfterFind.is_fine_grained());
1225 assert!(HookEvent::BeforeValidate.is_fine_grained());
1226 assert!(HookEvent::AfterValidate.is_fine_grained());
1227 assert!(!HookEvent::BeforeInsert.is_fine_grained());
1229 assert!(!HookEvent::AfterInsert.is_fine_grained());
1230 assert!(!HookEvent::BeforeUpdate.is_fine_grained());
1231 assert!(!HookEvent::AfterUpdate.is_fine_grained());
1232 assert!(!HookEvent::BeforeDelete.is_fine_grained());
1233 assert!(!HookEvent::AfterDelete.is_fine_grained());
1234 }
1235
1236 #[test]
1237 fn hook_registry_supports_find_and_validate_events() {
1238 let registry = HookRegistry::new();
1239 let counter = Arc::new(std::sync::atomic::AtomicU32::new(0));
1240
1241 for event in [
1242 HookEvent::BeforeFind,
1243 HookEvent::AfterFind,
1244 HookEvent::BeforeValidate,
1245 HookEvent::AfterValidate,
1246 ] {
1247 let c = Arc::clone(&counter);
1248 registry.register(
1249 event,
1250 Arc::new(move |_ctx| {
1251 c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1252 Ok(())
1253 }),
1254 );
1255 }
1256
1257 let ctx = HookContext::new();
1258 for event in [
1259 HookEvent::BeforeFind,
1260 HookEvent::AfterFind,
1261 HookEvent::BeforeValidate,
1262 HookEvent::AfterValidate,
1263 ] {
1264 registry.dispatch(event, &ctx).unwrap();
1265 }
1266
1267 assert_eq!(
1268 counter.load(std::sync::atomic::Ordering::SeqCst),
1269 4,
1270 "find/validate 钩子应能被注册与触发"
1271 );
1272 }
1273
1274 #[test]
1275 fn db_error_validation_error_code_and_display() {
1276 let err = DbError::Validation("name required".into());
1277 assert_eq!(err.error_code(), "DB021");
1278 assert_eq!(format!("{}", err), "Validation error: name required");
1279 }
1280}