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