1use crate::DbError;
24use crate::Value;
25use std::collections::HashMap;
26use std::sync::{Arc, RwLock};
27
28#[derive(Debug, Clone, Default)]
36pub struct HookContext {
37 pub tenant_id: Option<i64>,
39 pub operator_id: Option<i64>,
41 pub timestamp: u64,
43 pub metadata: HashMap<String, String>,
45}
46
47impl HookContext {
48 pub fn new() -> Self {
50 Self::default()
51 }
52
53 pub fn with_tenant(mut self, tenant_id: i64) -> Self {
55 self.tenant_id = Some(tenant_id);
56 self
57 }
58
59 pub fn with_operator(mut self, operator_id: i64) -> Self {
61 self.operator_id = Some(operator_id);
62 self
63 }
64
65 pub fn with_timestamp(mut self, ts: u64) -> Self {
67 self.timestamp = ts;
68 self
69 }
70
71 pub fn set_meta(&mut self, key: impl Into<String>, value: impl Into<String>) {
73 self.metadata.insert(key.into(), value.into());
74 }
75
76 pub fn get_meta(&self, key: &str) -> Option<&String> {
78 self.metadata.get(key)
79 }
80}
81
82#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
104pub enum HookEvent {
105 BeforeInsert,
106 AfterInsert,
107 BeforeUpdate,
108 AfterUpdate,
109 BeforeDelete,
110 AfterDelete,
111 BeforeWrite,
113 AfterWrite,
115 BeforeSave,
117 AfterSave,
119 BeforeRestore,
121 AfterRestore,
123 BeforeFind,
125 AfterFind,
127 BeforeValidate,
129 AfterValidate,
131}
132
133impl HookEvent {
134 pub fn is_before(&self) -> bool {
136 matches!(
137 self,
138 HookEvent::BeforeInsert
139 | HookEvent::BeforeUpdate
140 | HookEvent::BeforeDelete
141 | HookEvent::BeforeWrite
142 | HookEvent::BeforeSave
143 | HookEvent::BeforeRestore
144 | HookEvent::BeforeFind
145 | HookEvent::BeforeValidate
146 )
147 }
148
149 pub fn is_after(&self) -> bool {
151 matches!(
152 self,
153 HookEvent::AfterInsert
154 | HookEvent::AfterUpdate
155 | HookEvent::AfterDelete
156 | HookEvent::AfterWrite
157 | HookEvent::AfterSave
158 | HookEvent::AfterRestore
159 | HookEvent::AfterFind
160 | HookEvent::AfterValidate
161 )
162 }
163
164 pub fn is_write_level(&self) -> bool {
166 matches!(
167 self,
168 HookEvent::BeforeWrite
169 | HookEvent::AfterWrite
170 | HookEvent::BeforeSave
171 | HookEvent::AfterSave
172 )
173 }
174
175 pub fn is_find_level(&self) -> bool {
177 matches!(self, HookEvent::BeforeFind | HookEvent::AfterFind)
178 }
179
180 pub fn is_validate_level(&self) -> bool {
182 matches!(self, HookEvent::BeforeValidate | HookEvent::AfterValidate)
183 }
184
185 pub fn is_fine_grained(&self) -> bool {
187 self.is_write_level()
188 || self.is_find_level()
189 || self.is_validate_level()
190 || matches!(self, HookEvent::BeforeRestore | HookEvent::AfterRestore)
191 }
192}
193
194pub type HookResult<T> = Result<T, DbError>;
200
201pub trait Hookable: crate::model::Model {
220 fn before_insert(_ctx: &mut HookContext) -> HookResult<()> {
222 Ok(())
223 }
224
225 fn after_insert(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
227 Ok(())
228 }
229
230 fn before_update(_ctx: &mut HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
232 Ok(())
233 }
234
235 fn after_update(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
237 Ok(())
238 }
239
240 fn before_delete(_ctx: &mut HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
242 Ok(())
243 }
244
245 fn after_delete(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
247 Ok(())
248 }
249
250 fn before_write(_ctx: &mut HookContext) -> HookResult<()> {
254 Ok(())
255 }
256
257 fn after_write(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
259 Ok(())
260 }
261
262 fn before_save(_ctx: &mut HookContext) -> HookResult<()> {
264 Ok(())
265 }
266
267 fn after_save(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
269 Ok(())
270 }
271
272 fn before_restore(_ctx: &mut HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
276 Ok(())
277 }
278
279 fn after_restore(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
281 Ok(())
282 }
283
284 fn before_find(_ctx: &mut HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
289 Ok(())
290 }
291
292 fn after_find(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
297 Ok(())
298 }
299
300 fn before_validate(_ctx: &mut HookContext) -> HookResult<()> {
308 Ok(())
309 }
310
311 fn validate(_ctx: &mut HookContext) -> HookResult<()> {
317 Ok(())
318 }
319
320 fn after_validate(_ctx: &HookContext) -> HookResult<()> {
324 Ok(())
325 }
326}
327
328pub struct HookDispatcher;
332
333impl HookDispatcher {
334 pub fn insert<M, F>(ctx: &mut HookContext, f: F) -> HookResult<M::PrimaryKey>
340 where
341 M: Hookable,
342 F: FnOnce(&mut HookContext) -> HookResult<M::PrimaryKey>,
343 {
344 M::before_write(ctx)?;
345 M::before_save(ctx)?;
346 M::before_validate(ctx)?;
347 M::validate(ctx)?;
348 M::after_validate(ctx)?;
349 M::before_insert(ctx)?;
350 let id = f(ctx)?;
351 M::after_insert(ctx, &id)?;
352 M::after_save(ctx, &id)?;
353 M::after_write(ctx, &id)?;
354 Ok(id)
355 }
356
357 pub fn update<M, F>(ctx: &mut HookContext, id: &M::PrimaryKey, f: F) -> HookResult<()>
359 where
360 M: Hookable,
361 F: FnOnce(&mut HookContext) -> HookResult<()>,
362 {
363 M::before_write(ctx)?;
364 M::before_save(ctx)?;
365 M::before_validate(ctx)?;
366 M::validate(ctx)?;
367 M::after_validate(ctx)?;
368 M::before_update(ctx, id)?;
369 f(ctx)?;
370 M::after_update(ctx, id)?;
371 M::after_save(ctx, id)?;
372 M::after_write(ctx, id)?;
373 Ok(())
374 }
375
376 pub fn delete<M, F>(ctx: &mut HookContext, id: &M::PrimaryKey, f: F) -> HookResult<()>
378 where
379 M: Hookable,
380 F: FnOnce(&mut HookContext) -> HookResult<()>,
381 {
382 M::before_delete(ctx, id)?;
383 f(ctx)?;
384 M::after_delete(ctx, id)?;
385 Ok(())
386 }
387
388 pub fn restore<M, F>(ctx: &mut HookContext, id: &M::PrimaryKey, f: F) -> HookResult<()>
390 where
391 M: Hookable,
392 F: FnOnce(&mut HookContext) -> HookResult<()>,
393 {
394 M::before_restore(ctx, id)?;
395 f(ctx)?;
396 M::after_restore(ctx, id)?;
397 Ok(())
398 }
399
400 pub fn find<M, F>(ctx: &mut HookContext, id: &M::PrimaryKey, f: F) -> HookResult<()>
405 where
406 M: Hookable,
407 F: FnOnce(&mut HookContext) -> HookResult<()>,
408 {
409 M::before_find(ctx, id)?;
410 f(ctx)?;
411 M::after_find(ctx, id)?;
412 Ok(())
413 }
414
415 pub fn validate<M>(ctx: &mut HookContext) -> HookResult<()>
420 where
421 M: Hookable,
422 {
423 M::before_validate(ctx)?;
424 M::validate(ctx)?;
425 M::after_validate(ctx)?;
426 Ok(())
427 }
428}
429
430pub trait SoftDelete: crate::model::Model {
439 fn soft_delete_field() -> &'static str;
441
442 fn is_deleted(&self) -> bool;
444}
445
446pub trait GlobalScope {
458 fn scope_name() -> &'static str;
460
461 fn apply_scope(ctx: &HookContext) -> Option<(String, Vec<Value>)>;
466}
467
468pub struct SoftDeleteScope;
477
478impl<M: SoftDelete> GlobalScope for (SoftDeleteScope, M) {
479 fn scope_name() -> &'static str {
480 "soft_delete"
481 }
482
483 fn apply_scope(_ctx: &HookContext) -> Option<(String, Vec<Value>)> {
484 let field = <M as SoftDelete>::soft_delete_field();
486 Some((format!("{} IS NULL", field), vec![]))
487 }
488}
489
490pub struct TenantScope;
499
500pub trait TenantModel: crate::model::Model {
504 fn tenant_field() -> &'static str {
506 "tenant_id"
507 }
508
509 fn tenant_id(&self) -> i64;
511
512 fn set_tenant_id(&mut self, tenant_id: i64);
514}
515
516impl<M: TenantModel> GlobalScope for (TenantScope, M) {
517 fn scope_name() -> &'static str {
518 "tenant"
519 }
520
521 fn apply_scope(ctx: &HookContext) -> Option<(String, Vec<Value>)> {
522 ctx.tenant_id.map(|tid| {
523 (
524 format!("{} = ?", <M as TenantModel>::tenant_field()),
525 vec![Value::I64(tid)],
526 )
527 })
528 }
529}
530
531pub type HookFn = Arc<dyn Fn(&HookContext) -> HookResult<()> + Send + Sync>;
537
538pub struct HookRegistry {
543 hooks: RwLock<HashMap<HookEvent, Vec<HookFn>>>,
544}
545
546impl Default for HookRegistry {
547 fn default() -> Self {
548 Self::new()
549 }
550}
551
552impl HookRegistry {
553 pub fn new() -> Self {
555 Self {
556 hooks: RwLock::new(HashMap::new()),
557 }
558 }
559
560 pub fn register(&self, event: HookEvent, hook: HookFn) {
562 if let Ok(mut hooks) = self.hooks.write() {
564 hooks.entry(event).or_default().push(hook);
565 }
566 }
567
568 pub fn dispatch(&self, event: HookEvent, ctx: &HookContext) -> HookResult<()> {
572 let hooks = match self.hooks.read() {
573 Ok(h) => h,
574 Err(_) => return Ok(()), };
576 if let Some(fns) = hooks.get(&event) {
577 for f in fns {
578 f(ctx)?;
579 }
580 }
581 Ok(())
582 }
583
584 pub fn clear(&self, event: HookEvent) {
586 if let Ok(mut hooks) = self.hooks.write() {
587 hooks.remove(&event);
588 }
589 }
590
591 pub fn clear_all(&self) {
593 if let Ok(mut hooks) = self.hooks.write() {
594 hooks.clear();
595 }
596 }
597
598 pub fn count(&self, event: HookEvent) -> usize {
600 self.hooks
601 .read()
602 .map(|h| h.get(&event).map(|v| v.len()).unwrap_or(0))
603 .unwrap_or(0)
604 }
605}
606
607pub struct ScopeRegistry {
616 disabled: RwLock<Vec<String>>,
617}
618
619impl Default for ScopeRegistry {
620 fn default() -> Self {
621 Self::new()
622 }
623}
624
625impl ScopeRegistry {
626 pub fn new() -> Self {
628 Self {
629 disabled: RwLock::new(Vec::new()),
630 }
631 }
632
633 pub fn disable(&self, scope_name: impl Into<String>) {
635 if let Ok(mut disabled) = self.disabled.write() {
636 let name = scope_name.into();
637 if !disabled.contains(&name) {
638 disabled.push(name);
639 }
640 }
641 }
642
643 pub fn enable(&self, scope_name: &str) {
645 if let Ok(mut disabled) = self.disabled.write() {
646 disabled.retain(|n| n != scope_name);
647 }
648 }
649
650 pub fn is_enabled(&self, scope_name: &str) -> bool {
652 self.disabled
653 .read()
654 .map(|d| !d.iter().any(|n| n == scope_name))
655 .unwrap_or(true)
656 }
657
658 pub fn without_scope<F, R>(&self, scope_name: &str, f: F) -> R
668 where
669 F: FnOnce() -> R,
670 {
671 self.disable(scope_name);
672 let result = f();
673 self.enable(scope_name);
674 result
675 }
676}
677
678#[cfg(test)]
683mod tests {
684 use super::*;
685
686 #[test]
687 fn hook_context_builder() {
688 let ctx = HookContext::new()
689 .with_tenant(42)
690 .with_operator(1)
691 .with_timestamp(1700000000);
692
693 assert_eq!(ctx.tenant_id, Some(42));
694 assert_eq!(ctx.operator_id, Some(1));
695 assert_eq!(ctx.timestamp, 1700000000);
696 }
697
698 #[test]
699 fn hook_context_metadata() {
700 let mut ctx = HookContext::new();
701 ctx.set_meta("source", "api");
702 ctx.set_meta("ip", "127.0.0.1");
703
704 assert_eq!(ctx.get_meta("source"), Some(&"api".to_string()));
705 assert_eq!(ctx.get_meta("ip"), Some(&"127.0.0.1".to_string()));
706 assert_eq!(ctx.get_meta("missing"), None);
707 }
708
709 #[test]
710 fn hook_event_is_before_after() {
711 assert!(HookEvent::BeforeInsert.is_before());
712 assert!(!HookEvent::BeforeInsert.is_after());
713 assert!(HookEvent::AfterInsert.is_after());
714 assert!(!HookEvent::AfterInsert.is_before());
715 }
716
717 #[test]
718 fn hook_registry_register_and_dispatch() {
719 let registry = HookRegistry::new();
720 let counter = Arc::new(std::sync::atomic::AtomicU32::new(0));
721
722 let c = Arc::clone(&counter);
723 registry.register(
724 HookEvent::BeforeInsert,
725 Arc::new(move |_ctx| {
726 c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
727 Ok(())
728 }),
729 );
730
731 let ctx = HookContext::new();
732 registry.dispatch(HookEvent::BeforeInsert, &ctx).unwrap();
733 registry.dispatch(HookEvent::BeforeInsert, &ctx).unwrap();
734
735 assert_eq!(counter.load(std::sync::atomic::Ordering::SeqCst), 2);
736 }
737
738 #[test]
739 fn hook_registry_dispatch_no_hooks() {
740 let registry = HookRegistry::new();
741 let ctx = HookContext::new();
742 assert!(registry.dispatch(HookEvent::BeforeInsert, &ctx).is_ok());
744 }
745
746 #[test]
747 fn hook_registry_clear() {
748 let registry = HookRegistry::new();
749 registry.register(HookEvent::BeforeInsert, Arc::new(|_ctx| Ok(())));
750 assert_eq!(registry.count(HookEvent::BeforeInsert), 1);
751
752 registry.clear(HookEvent::BeforeInsert);
753 assert_eq!(registry.count(HookEvent::BeforeInsert), 0);
754 }
755
756 #[test]
757 fn hook_registry_clear_all() {
758 let registry = HookRegistry::new();
759 registry.register(HookEvent::BeforeInsert, Arc::new(|_ctx| Ok(())));
760 registry.register(HookEvent::AfterInsert, Arc::new(|_ctx| Ok(())));
761 registry.register(HookEvent::BeforeUpdate, Arc::new(|_ctx| Ok(())));
762
763 registry.clear_all();
764 assert_eq!(registry.count(HookEvent::BeforeInsert), 0);
765 assert_eq!(registry.count(HookEvent::AfterInsert), 0);
766 assert_eq!(registry.count(HookEvent::BeforeUpdate), 0);
767 }
768
769 #[test]
770 fn scope_registry_enable_disable() {
771 let registry = ScopeRegistry::new();
772
773 assert!(registry.is_enabled("soft_delete"));
774 assert!(registry.is_enabled("tenant"));
775
776 registry.disable("soft_delete");
777 assert!(!registry.is_enabled("soft_delete"));
778 assert!(registry.is_enabled("tenant"));
779
780 registry.enable("soft_delete");
781 assert!(registry.is_enabled("soft_delete"));
782 }
783
784 #[test]
785 fn scope_registry_without_scope() {
786 let registry = ScopeRegistry::new();
787 assert!(registry.is_enabled("soft_delete"));
788
789 let result = registry.without_scope("soft_delete", || {
790 assert!(!registry.is_enabled("soft_delete"));
791 42
792 });
793
794 assert_eq!(result, 42);
795 assert!(registry.is_enabled("soft_delete"));
796 }
797
798 #[test]
799 fn hook_registry_short_circuit_on_error() {
800 let registry = HookRegistry::new();
801 let called = Arc::new(std::sync::atomic::AtomicU32::new(0));
802
803 let c1 = Arc::clone(&called);
804 registry.register(
805 HookEvent::BeforeInsert,
806 Arc::new(move |_ctx| {
807 c1.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
808 Ok(())
809 }),
810 );
811
812 registry.register(
813 HookEvent::BeforeInsert,
814 Arc::new(|_ctx| Err(DbError::Hook("second hook failed".into()))),
815 );
816
817 let c3 = Arc::clone(&called);
818 registry.register(
819 HookEvent::BeforeInsert,
820 Arc::new(move |_ctx| {
821 c3.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
822 Ok(())
823 }),
824 );
825
826 let ctx = HookContext::new();
827 let result = registry.dispatch(HookEvent::BeforeInsert, &ctx);
828
829 assert!(result.is_err());
830 assert_eq!(called.load(std::sync::atomic::Ordering::SeqCst), 1);
832 }
833
834 #[test]
837 fn hook_event_is_write_level() {
838 assert!(HookEvent::BeforeWrite.is_write_level());
839 assert!(HookEvent::AfterWrite.is_write_level());
840 assert!(HookEvent::BeforeSave.is_write_level());
841 assert!(HookEvent::AfterSave.is_write_level());
842 assert!(!HookEvent::BeforeInsert.is_write_level());
843 assert!(!HookEvent::AfterDelete.is_write_level());
844 assert!(!HookEvent::BeforeRestore.is_write_level());
845 }
846
847 #[test]
848 fn hook_event_before_after_covers_new_variants() {
849 assert!(HookEvent::BeforeWrite.is_before());
850 assert!(HookEvent::BeforeSave.is_before());
851 assert!(HookEvent::BeforeRestore.is_before());
852 assert!(HookEvent::AfterWrite.is_after());
853 assert!(HookEvent::AfterSave.is_after());
854 assert!(HookEvent::AfterRestore.is_after());
855 assert!(!HookEvent::AfterWrite.is_before());
856 assert!(!HookEvent::BeforeWrite.is_after());
857 }
858
859 #[test]
860 fn hook_registry_supports_new_events() {
861 let registry = HookRegistry::new();
862 let counter = Arc::new(std::sync::atomic::AtomicU32::new(0));
863
864 for event in [
865 HookEvent::BeforeWrite,
866 HookEvent::AfterWrite,
867 HookEvent::BeforeSave,
868 HookEvent::AfterSave,
869 HookEvent::BeforeRestore,
870 HookEvent::AfterRestore,
871 ] {
872 let c = Arc::clone(&counter);
873 registry.register(
874 event,
875 Arc::new(move |_ctx| {
876 c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
877 Ok(())
878 }),
879 );
880 }
881
882 let ctx = HookContext::new();
883 for event in [
884 HookEvent::BeforeWrite,
885 HookEvent::AfterWrite,
886 HookEvent::BeforeSave,
887 HookEvent::AfterSave,
888 HookEvent::BeforeRestore,
889 HookEvent::AfterRestore,
890 ] {
891 registry.dispatch(event, &ctx).unwrap();
892 }
893
894 assert_eq!(
895 counter.load(std::sync::atomic::Ordering::SeqCst),
896 6,
897 "所有细粒度事件均应被正确注册与触发"
898 );
899 }
900
901 struct DispatchTestModel;
904 impl crate::model::Model for DispatchTestModel {
905 type PrimaryKey = i64;
906 fn table_name() -> &'static str {
907 "dispatch_test"
908 }
909 fn pk(&self) -> Self::PrimaryKey {
910 0
911 }
912 fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
913 }
914
915 static DISPATCH_CALLS: std::sync::OnceLock<Arc<std::sync::atomic::AtomicU32>> =
917 std::sync::OnceLock::new();
918
919 fn dispatch_calls() -> Arc<std::sync::atomic::AtomicU32> {
920 DISPATCH_CALLS
921 .get_or_init(|| Arc::new(std::sync::atomic::AtomicU32::new(0)))
922 .clone()
923 }
924
925 impl Hookable for DispatchTestModel {
926 fn before_write(ctx: &mut HookContext) -> HookResult<()> {
927 ctx.set_meta("before_write", "1");
928 Ok(())
929 }
930 fn before_save(ctx: &mut HookContext) -> HookResult<()> {
931 ctx.set_meta("before_save", "1");
932 Ok(())
933 }
934 fn before_validate(ctx: &mut HookContext) -> HookResult<()> {
935 ctx.set_meta("before_validate", "1");
936 Ok(())
937 }
938 fn after_validate(ctx: &HookContext) -> HookResult<()> {
939 assert_eq!(ctx.get_meta("before_validate"), Some(&"1".to_string()));
940 ctx_set_meta_for_after(ctx, "after_validate", "1");
941 Ok(())
942 }
943 fn before_insert(ctx: &mut HookContext) -> HookResult<()> {
944 ctx.set_meta("before_insert", "1");
945 Ok(())
946 }
947 fn after_insert(ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
948 assert_eq!(ctx.get_meta("before_write"), Some(&"1".to_string()));
949 assert_eq!(ctx.get_meta("before_save"), Some(&"1".to_string()));
950 assert_eq!(ctx.get_meta("before_insert"), Some(&"1".to_string()));
951 assert_eq!(ctx.get_meta("before_validate"), Some(&"1".to_string()));
952 Ok(())
953 }
954 fn after_save(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
955 dispatch_calls().fetch_add(1, std::sync::atomic::Ordering::SeqCst);
956 Ok(())
957 }
958 fn after_write(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
959 dispatch_calls().fetch_add(1, std::sync::atomic::Ordering::SeqCst);
960 Ok(())
961 }
962 fn before_find(ctx: &mut HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
963 ctx.set_meta("before_find", "1");
964 Ok(())
965 }
966 fn after_find(ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
967 assert_eq!(ctx.get_meta("before_find"), Some(&"1".to_string()));
968 ctx_set_meta_for_after(ctx, "after_find", "1");
969 Ok(())
970 }
971 }
972
973 static AFTER_VALIDATE_COUNT: std::sync::atomic::AtomicU32 =
976 std::sync::atomic::AtomicU32::new(0);
977 static AFTER_FIND_COUNT: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
978
979 fn ctx_set_meta_for_after(_ctx: &HookContext, key: &str, _value: &str) {
980 match key {
981 "after_validate" => {
982 AFTER_VALIDATE_COUNT.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
983 }
984 "after_find" => {
985 AFTER_FIND_COUNT.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
986 }
987 _ => {}
988 }
989 }
990
991 fn after_call_was(key: &str) -> bool {
992 match key {
993 "after_validate" => AFTER_VALIDATE_COUNT.load(std::sync::atomic::Ordering::SeqCst) > 0,
994 "after_find" => AFTER_FIND_COUNT.load(std::sync::atomic::Ordering::SeqCst) > 0,
995 _ => false,
996 }
997 }
998
999 fn reset_after_calls() {
1000 AFTER_VALIDATE_COUNT.store(0, std::sync::atomic::Ordering::SeqCst);
1001 AFTER_FIND_COUNT.store(0, std::sync::atomic::Ordering::SeqCst);
1002 }
1003
1004 static HOOK_TEST_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
1006
1007 #[test]
1008 fn hook_dispatcher_insert_full_sequence() {
1009 let _guard = HOOK_TEST_LOCK.lock().unwrap();
1010 dispatch_calls().store(0, std::sync::atomic::Ordering::SeqCst);
1011 reset_after_calls();
1012 let mut ctx = HookContext::new();
1013 let id = HookDispatcher::insert::<DispatchTestModel, _>(&mut ctx, |_ctx| Ok(42_i64));
1014 assert!(id.is_ok());
1015 assert_eq!(id.unwrap(), 42);
1016 assert_eq!(ctx.get_meta("before_write"), Some(&"1".to_string()));
1018 assert_eq!(ctx.get_meta("before_save"), Some(&"1".to_string()));
1019 assert_eq!(ctx.get_meta("before_insert"), Some(&"1".to_string()));
1020 assert_eq!(ctx.get_meta("before_validate"), Some(&"1".to_string()));
1022 assert!(after_call_was("after_validate"));
1023 assert_eq!(
1025 dispatch_calls().load(std::sync::atomic::Ordering::SeqCst),
1026 2
1027 );
1028 }
1029
1030 #[test]
1031 fn hook_dispatcher_insert_short_circuit_on_before_write_error() {
1032 struct ErrorModel;
1033 impl crate::model::Model for ErrorModel {
1034 type PrimaryKey = i64;
1035 fn table_name() -> &'static str {
1036 "error_model"
1037 }
1038 fn pk(&self) -> Self::PrimaryKey {
1039 0
1040 }
1041 fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
1042 }
1043 impl Hookable for ErrorModel {
1044 fn before_write(_ctx: &mut HookContext) -> HookResult<()> {
1045 Err(DbError::Hook("before_write failed".into()))
1046 }
1047 }
1048
1049 let mut ctx = HookContext::new();
1050 let result = HookDispatcher::insert::<ErrorModel, _>(&mut ctx, |_ctx| Ok(1_i64));
1051 assert!(result.is_err());
1052 }
1054
1055 #[test]
1056 fn hook_dispatcher_insert_short_circuit_on_before_validate_error() {
1057 struct ValidationFailModel;
1058 impl crate::model::Model for ValidationFailModel {
1059 type PrimaryKey = i64;
1060 fn table_name() -> &'static str {
1061 "validation_fail"
1062 }
1063 fn pk(&self) -> Self::PrimaryKey {
1064 0
1065 }
1066 fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
1067 }
1068 impl Hookable for ValidationFailModel {
1069 fn before_validate(_ctx: &mut HookContext) -> HookResult<()> {
1070 Err(DbError::Validation("name is required".into()))
1071 }
1072 }
1073
1074 let mut ctx = HookContext::new();
1075 let called = Arc::new(std::sync::atomic::AtomicU32::new(0));
1076 let c = Arc::clone(&called);
1077 let result = HookDispatcher::insert::<ValidationFailModel, _>(&mut ctx, move |_ctx| {
1078 c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1079 Ok(1_i64)
1080 });
1081 assert!(result.is_err());
1082 assert_eq!(
1084 called.load(std::sync::atomic::Ordering::SeqCst),
1085 0,
1086 "before_validate 失败应短路 INSERT 操作"
1087 );
1088 match result.unwrap_err() {
1090 DbError::Validation(msg) => assert_eq!(msg, "name is required"),
1091 other => panic!("期望 Validation 错误,得到 {:?}", other),
1092 }
1093 }
1094
1095 #[test]
1096 fn hook_dispatcher_update_full_sequence() {
1097 let _guard = HOOK_TEST_LOCK.lock().unwrap();
1098 dispatch_calls().store(0, std::sync::atomic::Ordering::SeqCst);
1099 reset_after_calls();
1100 let mut ctx = HookContext::new();
1101 let result =
1102 HookDispatcher::update::<DispatchTestModel, _>(&mut ctx, &42_i64, |_ctx| Ok(()));
1103 assert!(result.is_ok());
1104 assert_eq!(
1106 dispatch_calls().load(std::sync::atomic::Ordering::SeqCst),
1107 2
1108 );
1109 assert_eq!(ctx.get_meta("before_validate"), Some(&"1".to_string()));
1111 assert!(after_call_was("after_validate"));
1112 }
1113
1114 #[test]
1115 fn hook_dispatcher_delete_full_sequence() {
1116 let mut ctx = HookContext::new();
1117 let result =
1118 HookDispatcher::delete::<DispatchTestModel, _>(&mut ctx, &42_i64, |_ctx| Ok(()));
1119 assert!(result.is_ok());
1120 }
1121
1122 #[test]
1123 fn hook_dispatcher_restore_full_sequence() {
1124 let mut ctx = HookContext::new();
1125 let result =
1126 HookDispatcher::restore::<DispatchTestModel, _>(&mut ctx, &42_i64, |_ctx| Ok(()));
1127 assert!(result.is_ok());
1128 }
1129
1130 #[test]
1131 fn hook_dispatcher_find_full_sequence() {
1132 let _guard = HOOK_TEST_LOCK.lock().unwrap();
1133 dispatch_calls().store(0, std::sync::atomic::Ordering::SeqCst);
1134 reset_after_calls();
1135 let mut ctx = HookContext::new();
1136 let called = Arc::new(std::sync::atomic::AtomicU32::new(0));
1137 let c = Arc::clone(&called);
1138 let result = HookDispatcher::find::<DispatchTestModel, _>(&mut ctx, &42_i64, move |_ctx| {
1139 c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1140 Ok(())
1141 });
1142 assert!(result.is_ok());
1143 assert_eq!(
1144 called.load(std::sync::atomic::Ordering::SeqCst),
1145 1,
1146 "SELECT 操作应执行一次"
1147 );
1148 assert_eq!(ctx.get_meta("before_find"), Some(&"1".to_string()));
1150 assert!(after_call_was("after_find"));
1151 }
1152
1153 #[test]
1154 fn hook_dispatcher_find_short_circuit_on_before_find_error() {
1155 struct FindFailModel;
1156 impl crate::model::Model for FindFailModel {
1157 type PrimaryKey = i64;
1158 fn table_name() -> &'static str {
1159 "find_fail"
1160 }
1161 fn pk(&self) -> Self::PrimaryKey {
1162 0
1163 }
1164 fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
1165 }
1166 impl Hookable for FindFailModel {
1167 fn before_find(_ctx: &mut HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
1168 Err(DbError::Hook("before_find blocked".into()))
1169 }
1170 }
1171
1172 let mut ctx = HookContext::new();
1173 let called = Arc::new(std::sync::atomic::AtomicU32::new(0));
1174 let c = Arc::clone(&called);
1175 let result = HookDispatcher::find::<FindFailModel, _>(&mut ctx, &1_i64, move |_ctx| {
1176 c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1177 Ok(())
1178 });
1179 assert!(result.is_err());
1180 assert_eq!(
1181 called.load(std::sync::atomic::Ordering::SeqCst),
1182 0,
1183 "before_find 失败应短路 SELECT"
1184 );
1185 }
1186
1187 #[test]
1188 fn hook_dispatcher_validate_standalone() {
1189 reset_after_calls();
1190 let mut ctx = HookContext::new();
1191 let result = HookDispatcher::validate::<DispatchTestModel>(&mut ctx);
1192 assert!(result.is_ok());
1193 assert_eq!(ctx.get_meta("before_validate"), Some(&"1".to_string()));
1194 assert!(after_call_was("after_validate"));
1195 }
1196
1197 #[test]
1198 fn hook_event_is_find_level_and_is_validate_level() {
1199 assert!(HookEvent::BeforeFind.is_find_level());
1200 assert!(HookEvent::AfterFind.is_find_level());
1201 assert!(HookEvent::BeforeValidate.is_validate_level());
1202 assert!(HookEvent::AfterValidate.is_validate_level());
1203 assert!(!HookEvent::BeforeInsert.is_find_level());
1204 assert!(!HookEvent::BeforeInsert.is_validate_level());
1205 assert!(!HookEvent::BeforeWrite.is_find_level());
1206 assert!(!HookEvent::BeforeWrite.is_validate_level());
1207 }
1208
1209 #[test]
1210 fn hook_event_is_fine_grained_covers_all_v02_events() {
1211 assert!(HookEvent::BeforeWrite.is_fine_grained());
1213 assert!(HookEvent::AfterWrite.is_fine_grained());
1214 assert!(HookEvent::BeforeSave.is_fine_grained());
1215 assert!(HookEvent::AfterSave.is_fine_grained());
1216 assert!(HookEvent::BeforeRestore.is_fine_grained());
1217 assert!(HookEvent::AfterRestore.is_fine_grained());
1218 assert!(HookEvent::BeforeFind.is_fine_grained());
1219 assert!(HookEvent::AfterFind.is_fine_grained());
1220 assert!(HookEvent::BeforeValidate.is_fine_grained());
1221 assert!(HookEvent::AfterValidate.is_fine_grained());
1222 assert!(!HookEvent::BeforeInsert.is_fine_grained());
1224 assert!(!HookEvent::AfterInsert.is_fine_grained());
1225 assert!(!HookEvent::BeforeUpdate.is_fine_grained());
1226 assert!(!HookEvent::AfterUpdate.is_fine_grained());
1227 assert!(!HookEvent::BeforeDelete.is_fine_grained());
1228 assert!(!HookEvent::AfterDelete.is_fine_grained());
1229 }
1230
1231 #[test]
1232 fn hook_registry_supports_find_and_validate_events() {
1233 let registry = HookRegistry::new();
1234 let counter = Arc::new(std::sync::atomic::AtomicU32::new(0));
1235
1236 for event in [
1237 HookEvent::BeforeFind,
1238 HookEvent::AfterFind,
1239 HookEvent::BeforeValidate,
1240 HookEvent::AfterValidate,
1241 ] {
1242 let c = Arc::clone(&counter);
1243 registry.register(
1244 event,
1245 Arc::new(move |_ctx| {
1246 c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
1247 Ok(())
1248 }),
1249 );
1250 }
1251
1252 let ctx = HookContext::new();
1253 for event in [
1254 HookEvent::BeforeFind,
1255 HookEvent::AfterFind,
1256 HookEvent::BeforeValidate,
1257 HookEvent::AfterValidate,
1258 ] {
1259 registry.dispatch(event, &ctx).unwrap();
1260 }
1261
1262 assert_eq!(
1263 counter.load(std::sync::atomic::Ordering::SeqCst),
1264 4,
1265 "find/validate 钩子应能被注册与触发"
1266 );
1267 }
1268
1269 #[test]
1270 fn db_error_validation_error_code_and_display() {
1271 let err = DbError::Validation("name required".into());
1272 assert_eq!(err.error_code(), "DB021");
1273 assert_eq!(format!("{}", err), "Validation error: name required");
1274 }
1275}