use crate::DbError;
use crate::Value;
use std::collections::HashMap;
use std::sync::{Arc, RwLock};
#[derive(Debug, Clone, Default)]
pub struct HookContext {
pub tenant_id: Option<i64>,
pub operator_id: Option<i64>,
pub timestamp: u64,
pub metadata: HashMap<String, String>,
}
impl HookContext {
pub fn new() -> Self {
Self::default()
}
pub fn with_tenant(mut self, tenant_id: i64) -> Self {
self.tenant_id = Some(tenant_id);
self
}
pub fn with_operator(mut self, operator_id: i64) -> Self {
self.operator_id = Some(operator_id);
self
}
pub fn with_timestamp(mut self, ts: u64) -> Self {
self.timestamp = ts;
self
}
pub fn set_meta(&mut self, key: impl Into<String>, value: impl Into<String>) {
self.metadata.insert(key.into(), value.into());
}
pub fn get_meta(&self, key: &str) -> Option<&String> {
self.metadata.get(key)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum HookEvent {
BeforeInsert,
AfterInsert,
BeforeUpdate,
AfterUpdate,
BeforeDelete,
AfterDelete,
BeforeWrite,
AfterWrite,
BeforeSave,
AfterSave,
BeforeRestore,
AfterRestore,
BeforeFind,
AfterFind,
BeforeValidate,
AfterValidate,
}
impl HookEvent {
pub fn is_before(&self) -> bool {
matches!(
self,
HookEvent::BeforeInsert
| HookEvent::BeforeUpdate
| HookEvent::BeforeDelete
| HookEvent::BeforeWrite
| HookEvent::BeforeSave
| HookEvent::BeforeRestore
| HookEvent::BeforeFind
| HookEvent::BeforeValidate
)
}
pub fn is_after(&self) -> bool {
matches!(
self,
HookEvent::AfterInsert
| HookEvent::AfterUpdate
| HookEvent::AfterDelete
| HookEvent::AfterWrite
| HookEvent::AfterSave
| HookEvent::AfterRestore
| HookEvent::AfterFind
| HookEvent::AfterValidate
)
}
pub fn is_write_level(&self) -> bool {
matches!(
self,
HookEvent::BeforeWrite
| HookEvent::AfterWrite
| HookEvent::BeforeSave
| HookEvent::AfterSave
)
}
pub fn is_find_level(&self) -> bool {
matches!(self, HookEvent::BeforeFind | HookEvent::AfterFind)
}
pub fn is_validate_level(&self) -> bool {
matches!(self, HookEvent::BeforeValidate | HookEvent::AfterValidate)
}
pub fn is_fine_grained(&self) -> bool {
self.is_write_level()
|| self.is_find_level()
|| self.is_validate_level()
|| matches!(self, HookEvent::BeforeRestore | HookEvent::AfterRestore)
}
}
pub type HookResult<T> = Result<T, DbError>;
pub trait Hookable: crate::model::Model {
fn before_insert(_ctx: &mut HookContext) -> HookResult<()> {
Ok(())
}
fn after_insert(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
Ok(())
}
fn before_update(_ctx: &mut HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
Ok(())
}
fn after_update(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
Ok(())
}
fn before_delete(_ctx: &mut HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
Ok(())
}
fn after_delete(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
Ok(())
}
fn before_write(_ctx: &mut HookContext) -> HookResult<()> {
Ok(())
}
fn after_write(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
Ok(())
}
fn before_save(_ctx: &mut HookContext) -> HookResult<()> {
Ok(())
}
fn after_save(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
Ok(())
}
fn before_restore(_ctx: &mut HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
Ok(())
}
fn after_restore(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
Ok(())
}
fn before_find(_ctx: &mut HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
Ok(())
}
fn after_find(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
Ok(())
}
fn before_validate(_ctx: &mut HookContext) -> HookResult<()> {
Ok(())
}
fn validate(_ctx: &mut HookContext) -> HookResult<()> {
Ok(())
}
fn after_validate(_ctx: &HookContext) -> HookResult<()> {
Ok(())
}
}
pub struct HookDispatcher;
impl HookDispatcher {
pub fn insert<M, F>(ctx: &mut HookContext, f: F) -> HookResult<M::PrimaryKey>
where
M: Hookable,
F: FnOnce(&mut HookContext) -> HookResult<M::PrimaryKey>,
{
M::before_write(ctx)?;
M::before_save(ctx)?;
M::before_validate(ctx)?;
M::validate(ctx)?;
M::after_validate(ctx)?;
M::before_insert(ctx)?;
let id = f(ctx)?;
M::after_insert(ctx, &id)?;
M::after_save(ctx, &id)?;
M::after_write(ctx, &id)?;
Ok(id)
}
pub fn update<M, F>(ctx: &mut HookContext, id: &M::PrimaryKey, f: F) -> HookResult<()>
where
M: Hookable,
F: FnOnce(&mut HookContext) -> HookResult<()>,
{
M::before_write(ctx)?;
M::before_save(ctx)?;
M::before_validate(ctx)?;
M::validate(ctx)?;
M::after_validate(ctx)?;
M::before_update(ctx, id)?;
f(ctx)?;
M::after_update(ctx, id)?;
M::after_save(ctx, id)?;
M::after_write(ctx, id)?;
Ok(())
}
pub fn delete<M, F>(ctx: &mut HookContext, id: &M::PrimaryKey, f: F) -> HookResult<()>
where
M: Hookable,
F: FnOnce(&mut HookContext) -> HookResult<()>,
{
M::before_delete(ctx, id)?;
f(ctx)?;
M::after_delete(ctx, id)?;
Ok(())
}
pub fn restore<M, F>(ctx: &mut HookContext, id: &M::PrimaryKey, f: F) -> HookResult<()>
where
M: Hookable,
F: FnOnce(&mut HookContext) -> HookResult<()>,
{
M::before_restore(ctx, id)?;
f(ctx)?;
M::after_restore(ctx, id)?;
Ok(())
}
pub fn find<M, F>(ctx: &mut HookContext, id: &M::PrimaryKey, f: F) -> HookResult<()>
where
M: Hookable,
F: FnOnce(&mut HookContext) -> HookResult<()>,
{
M::before_find(ctx, id)?;
f(ctx)?;
M::after_find(ctx, id)?;
Ok(())
}
pub fn validate<M>(ctx: &mut HookContext) -> HookResult<()>
where
M: Hookable,
{
M::before_validate(ctx)?;
M::validate(ctx)?;
M::after_validate(ctx)?;
Ok(())
}
}
pub trait SoftDelete: crate::model::Model {
fn soft_delete_field() -> &'static str;
fn is_deleted(&self) -> bool;
}
pub trait GlobalScope {
fn scope_name() -> &'static str;
fn apply_scope(ctx: &HookContext) -> Option<(String, Vec<Value>)>;
}
pub struct SoftDeleteScope;
impl<M: SoftDelete> GlobalScope for (SoftDeleteScope, M) {
fn scope_name() -> &'static str {
"soft_delete"
}
fn apply_scope(_ctx: &HookContext) -> Option<(String, Vec<Value>)> {
let field = <M as SoftDelete>::soft_delete_field();
Some((format!("{} IS NULL", field), vec![]))
}
}
pub struct TenantScope;
pub trait TenantModel: crate::model::Model {
fn tenant_field() -> &'static str {
"tenant_id"
}
fn tenant_id(&self) -> i64;
fn set_tenant_id(&mut self, tenant_id: i64);
}
impl<M: TenantModel> GlobalScope for (TenantScope, M) {
fn scope_name() -> &'static str {
"tenant"
}
fn apply_scope(ctx: &HookContext) -> Option<(String, Vec<Value>)> {
ctx.tenant_id.map(|tid| {
(
format!("{} = ?", <M as TenantModel>::tenant_field()),
vec![Value::I64(tid)],
)
})
}
}
pub type HookFn = Arc<dyn Fn(&HookContext) -> HookResult<()> + Send + Sync>;
pub struct HookRegistry {
hooks: RwLock<HashMap<HookEvent, Vec<HookFn>>>,
}
impl Default for HookRegistry {
fn default() -> Self {
Self::new()
}
}
impl HookRegistry {
pub fn new() -> Self {
Self {
hooks: RwLock::new(HashMap::new()),
}
}
pub fn register(&self, event: HookEvent, hook: HookFn) {
if let Ok(mut hooks) = self.hooks.write() {
hooks.entry(event).or_default().push(hook);
}
}
pub fn dispatch(&self, event: HookEvent, ctx: &HookContext) -> HookResult<()> {
let hooks = match self.hooks.read() {
Ok(h) => h,
Err(_) => return Ok(()), };
if let Some(fns) = hooks.get(&event) {
for f in fns {
f(ctx)?;
}
}
Ok(())
}
pub fn clear(&self, event: HookEvent) {
if let Ok(mut hooks) = self.hooks.write() {
hooks.remove(&event);
}
}
pub fn clear_all(&self) {
if let Ok(mut hooks) = self.hooks.write() {
hooks.clear();
}
}
pub fn count(&self, event: HookEvent) -> usize {
self.hooks
.read()
.map(|h| h.get(&event).map(|v| v.len()).unwrap_or(0))
.unwrap_or(0)
}
}
pub struct ScopeRegistry {
disabled: RwLock<Vec<String>>,
}
impl Default for ScopeRegistry {
fn default() -> Self {
Self::new()
}
}
impl ScopeRegistry {
pub fn new() -> Self {
Self {
disabled: RwLock::new(Vec::new()),
}
}
pub fn disable(&self, scope_name: impl Into<String>) {
if let Ok(mut disabled) = self.disabled.write() {
let name = scope_name.into();
if !disabled.contains(&name) {
disabled.push(name);
}
}
}
pub fn enable(&self, scope_name: &str) {
if let Ok(mut disabled) = self.disabled.write() {
disabled.retain(|n| n != scope_name);
}
}
pub fn is_enabled(&self, scope_name: &str) -> bool {
self.disabled
.read()
.map(|d| !d.iter().any(|n| n == scope_name))
.unwrap_or(true)
}
pub fn without_scope<F, R>(&self, scope_name: &str, f: F) -> R
where
F: FnOnce() -> R,
{
self.disable(scope_name);
let result = f();
self.enable(scope_name);
result
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn hook_context_builder() {
let ctx = HookContext::new()
.with_tenant(42)
.with_operator(1)
.with_timestamp(1700000000);
assert_eq!(ctx.tenant_id, Some(42));
assert_eq!(ctx.operator_id, Some(1));
assert_eq!(ctx.timestamp, 1700000000);
}
#[test]
fn hook_context_metadata() {
let mut ctx = HookContext::new();
ctx.set_meta("source", "api");
ctx.set_meta("ip", "127.0.0.1");
assert_eq!(ctx.get_meta("source"), Some(&"api".to_string()));
assert_eq!(ctx.get_meta("ip"), Some(&"127.0.0.1".to_string()));
assert_eq!(ctx.get_meta("missing"), None);
}
#[test]
fn hook_event_is_before_after() {
assert!(HookEvent::BeforeInsert.is_before());
assert!(!HookEvent::BeforeInsert.is_after());
assert!(HookEvent::AfterInsert.is_after());
assert!(!HookEvent::AfterInsert.is_before());
}
#[test]
fn hook_registry_register_and_dispatch() {
let registry = HookRegistry::new();
let counter = Arc::new(std::sync::atomic::AtomicU32::new(0));
let c = Arc::clone(&counter);
registry.register(
HookEvent::BeforeInsert,
Arc::new(move |_ctx| {
c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Ok(())
}),
);
let ctx = HookContext::new();
registry.dispatch(HookEvent::BeforeInsert, &ctx).unwrap();
registry.dispatch(HookEvent::BeforeInsert, &ctx).unwrap();
assert_eq!(counter.load(std::sync::atomic::Ordering::SeqCst), 2);
}
#[test]
fn hook_registry_dispatch_no_hooks() {
let registry = HookRegistry::new();
let ctx = HookContext::new();
assert!(registry.dispatch(HookEvent::BeforeInsert, &ctx).is_ok());
}
#[test]
fn hook_registry_clear() {
let registry = HookRegistry::new();
registry.register(HookEvent::BeforeInsert, Arc::new(|_ctx| Ok(())));
assert_eq!(registry.count(HookEvent::BeforeInsert), 1);
registry.clear(HookEvent::BeforeInsert);
assert_eq!(registry.count(HookEvent::BeforeInsert), 0);
}
#[test]
fn hook_registry_clear_all() {
let registry = HookRegistry::new();
registry.register(HookEvent::BeforeInsert, Arc::new(|_ctx| Ok(())));
registry.register(HookEvent::AfterInsert, Arc::new(|_ctx| Ok(())));
registry.register(HookEvent::BeforeUpdate, Arc::new(|_ctx| Ok(())));
registry.clear_all();
assert_eq!(registry.count(HookEvent::BeforeInsert), 0);
assert_eq!(registry.count(HookEvent::AfterInsert), 0);
assert_eq!(registry.count(HookEvent::BeforeUpdate), 0);
}
#[test]
fn scope_registry_enable_disable() {
let registry = ScopeRegistry::new();
assert!(registry.is_enabled("soft_delete"));
assert!(registry.is_enabled("tenant"));
registry.disable("soft_delete");
assert!(!registry.is_enabled("soft_delete"));
assert!(registry.is_enabled("tenant"));
registry.enable("soft_delete");
assert!(registry.is_enabled("soft_delete"));
}
#[test]
fn scope_registry_without_scope() {
let registry = ScopeRegistry::new();
assert!(registry.is_enabled("soft_delete"));
let result = registry.without_scope("soft_delete", || {
assert!(!registry.is_enabled("soft_delete"));
42
});
assert_eq!(result, 42);
assert!(registry.is_enabled("soft_delete"));
}
#[test]
fn hook_registry_short_circuit_on_error() {
let registry = HookRegistry::new();
let called = Arc::new(std::sync::atomic::AtomicU32::new(0));
let c1 = Arc::clone(&called);
registry.register(
HookEvent::BeforeInsert,
Arc::new(move |_ctx| {
c1.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Ok(())
}),
);
registry.register(
HookEvent::BeforeInsert,
Arc::new(|_ctx| Err(DbError::Hook("second hook failed".into()))),
);
let c3 = Arc::clone(&called);
registry.register(
HookEvent::BeforeInsert,
Arc::new(move |_ctx| {
c3.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Ok(())
}),
);
let ctx = HookContext::new();
let result = registry.dispatch(HookEvent::BeforeInsert, &ctx);
assert!(result.is_err());
assert_eq!(called.load(std::sync::atomic::Ordering::SeqCst), 1);
}
#[test]
fn hook_event_is_write_level() {
assert!(HookEvent::BeforeWrite.is_write_level());
assert!(HookEvent::AfterWrite.is_write_level());
assert!(HookEvent::BeforeSave.is_write_level());
assert!(HookEvent::AfterSave.is_write_level());
assert!(!HookEvent::BeforeInsert.is_write_level());
assert!(!HookEvent::AfterDelete.is_write_level());
assert!(!HookEvent::BeforeRestore.is_write_level());
}
#[test]
fn hook_event_before_after_covers_new_variants() {
assert!(HookEvent::BeforeWrite.is_before());
assert!(HookEvent::BeforeSave.is_before());
assert!(HookEvent::BeforeRestore.is_before());
assert!(HookEvent::AfterWrite.is_after());
assert!(HookEvent::AfterSave.is_after());
assert!(HookEvent::AfterRestore.is_after());
assert!(!HookEvent::AfterWrite.is_before());
assert!(!HookEvent::BeforeWrite.is_after());
}
#[test]
fn hook_registry_supports_new_events() {
let registry = HookRegistry::new();
let counter = Arc::new(std::sync::atomic::AtomicU32::new(0));
for event in [
HookEvent::BeforeWrite,
HookEvent::AfterWrite,
HookEvent::BeforeSave,
HookEvent::AfterSave,
HookEvent::BeforeRestore,
HookEvent::AfterRestore,
] {
let c = Arc::clone(&counter);
registry.register(
event,
Arc::new(move |_ctx| {
c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Ok(())
}),
);
}
let ctx = HookContext::new();
for event in [
HookEvent::BeforeWrite,
HookEvent::AfterWrite,
HookEvent::BeforeSave,
HookEvent::AfterSave,
HookEvent::BeforeRestore,
HookEvent::AfterRestore,
] {
registry.dispatch(event, &ctx).unwrap();
}
assert_eq!(
counter.load(std::sync::atomic::Ordering::SeqCst),
6,
"所有细粒度事件均应被正确注册与触发"
);
}
struct DispatchTestModel;
impl crate::model::Model for DispatchTestModel {
type PrimaryKey = i64;
fn table_name() -> &'static str {
"dispatch_test"
}
fn pk(&self) -> Self::PrimaryKey {
0
}
fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
}
static DISPATCH_CALLS: std::sync::OnceLock<Arc<std::sync::atomic::AtomicU32>> =
std::sync::OnceLock::new();
fn dispatch_calls() -> Arc<std::sync::atomic::AtomicU32> {
DISPATCH_CALLS
.get_or_init(|| Arc::new(std::sync::atomic::AtomicU32::new(0)))
.clone()
}
impl Hookable for DispatchTestModel {
fn before_write(ctx: &mut HookContext) -> HookResult<()> {
ctx.set_meta("before_write", "1");
Ok(())
}
fn before_save(ctx: &mut HookContext) -> HookResult<()> {
ctx.set_meta("before_save", "1");
Ok(())
}
fn before_validate(ctx: &mut HookContext) -> HookResult<()> {
ctx.set_meta("before_validate", "1");
Ok(())
}
fn after_validate(ctx: &HookContext) -> HookResult<()> {
assert_eq!(ctx.get_meta("before_validate"), Some(&"1".to_string()));
ctx_set_meta_for_after(ctx, "after_validate", "1");
Ok(())
}
fn before_insert(ctx: &mut HookContext) -> HookResult<()> {
ctx.set_meta("before_insert", "1");
Ok(())
}
fn after_insert(ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
assert_eq!(ctx.get_meta("before_write"), Some(&"1".to_string()));
assert_eq!(ctx.get_meta("before_save"), Some(&"1".to_string()));
assert_eq!(ctx.get_meta("before_insert"), Some(&"1".to_string()));
assert_eq!(ctx.get_meta("before_validate"), Some(&"1".to_string()));
Ok(())
}
fn after_save(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
dispatch_calls().fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Ok(())
}
fn after_write(_ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
dispatch_calls().fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Ok(())
}
fn before_find(ctx: &mut HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
ctx.set_meta("before_find", "1");
Ok(())
}
fn after_find(ctx: &HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
assert_eq!(ctx.get_meta("before_find"), Some(&"1".to_string()));
ctx_set_meta_for_after(ctx, "after_find", "1");
Ok(())
}
}
static AFTER_VALIDATE_COUNT: std::sync::atomic::AtomicU32 =
std::sync::atomic::AtomicU32::new(0);
static AFTER_FIND_COUNT: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
fn ctx_set_meta_for_after(_ctx: &HookContext, key: &str, _value: &str) {
match key {
"after_validate" => {
AFTER_VALIDATE_COUNT.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
}
"after_find" => {
AFTER_FIND_COUNT.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
}
_ => {}
}
}
fn after_call_was(key: &str) -> bool {
match key {
"after_validate" => AFTER_VALIDATE_COUNT.load(std::sync::atomic::Ordering::SeqCst) > 0,
"after_find" => AFTER_FIND_COUNT.load(std::sync::atomic::Ordering::SeqCst) > 0,
_ => false,
}
}
fn reset_after_calls() {
AFTER_VALIDATE_COUNT.store(0, std::sync::atomic::Ordering::SeqCst);
AFTER_FIND_COUNT.store(0, std::sync::atomic::Ordering::SeqCst);
}
static HOOK_TEST_LOCK: std::sync::Mutex<()> = std::sync::Mutex::new(());
#[test]
fn hook_dispatcher_insert_full_sequence() {
let _guard = HOOK_TEST_LOCK.lock().unwrap();
dispatch_calls().store(0, std::sync::atomic::Ordering::SeqCst);
reset_after_calls();
let mut ctx = HookContext::new();
let id = HookDispatcher::insert::<DispatchTestModel, _>(&mut ctx, |_ctx| Ok(42_i64));
assert!(id.is_ok());
assert_eq!(id.unwrap(), 42);
assert_eq!(ctx.get_meta("before_write"), Some(&"1".to_string()));
assert_eq!(ctx.get_meta("before_save"), Some(&"1".to_string()));
assert_eq!(ctx.get_meta("before_insert"), Some(&"1".to_string()));
assert_eq!(ctx.get_meta("before_validate"), Some(&"1".to_string()));
assert!(after_call_was("after_validate"));
assert_eq!(
dispatch_calls().load(std::sync::atomic::Ordering::SeqCst),
2
);
}
#[test]
fn hook_dispatcher_insert_short_circuit_on_before_write_error() {
struct ErrorModel;
impl crate::model::Model for ErrorModel {
type PrimaryKey = i64;
fn table_name() -> &'static str {
"error_model"
}
fn pk(&self) -> Self::PrimaryKey {
0
}
fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
}
impl Hookable for ErrorModel {
fn before_write(_ctx: &mut HookContext) -> HookResult<()> {
Err(DbError::Hook("before_write failed".into()))
}
}
let mut ctx = HookContext::new();
let result = HookDispatcher::insert::<ErrorModel, _>(&mut ctx, |_ctx| Ok(1_i64));
assert!(result.is_err());
}
#[test]
fn hook_dispatcher_insert_short_circuit_on_before_validate_error() {
struct ValidationFailModel;
impl crate::model::Model for ValidationFailModel {
type PrimaryKey = i64;
fn table_name() -> &'static str {
"validation_fail"
}
fn pk(&self) -> Self::PrimaryKey {
0
}
fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
}
impl Hookable for ValidationFailModel {
fn before_validate(_ctx: &mut HookContext) -> HookResult<()> {
Err(DbError::Validation("name is required".into()))
}
}
let mut ctx = HookContext::new();
let called = Arc::new(std::sync::atomic::AtomicU32::new(0));
let c = Arc::clone(&called);
let result = HookDispatcher::insert::<ValidationFailModel, _>(&mut ctx, move |_ctx| {
c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Ok(1_i64)
});
assert!(result.is_err());
assert_eq!(
called.load(std::sync::atomic::Ordering::SeqCst),
0,
"before_validate 失败应短路 INSERT 操作"
);
match result.unwrap_err() {
DbError::Validation(msg) => assert_eq!(msg, "name is required"),
other => panic!("期望 Validation 错误,得到 {:?}", other),
}
}
#[test]
fn hook_dispatcher_update_full_sequence() {
let _guard = HOOK_TEST_LOCK.lock().unwrap();
dispatch_calls().store(0, std::sync::atomic::Ordering::SeqCst);
reset_after_calls();
let mut ctx = HookContext::new();
let result =
HookDispatcher::update::<DispatchTestModel, _>(&mut ctx, &42_i64, |_ctx| Ok(()));
assert!(result.is_ok());
assert_eq!(
dispatch_calls().load(std::sync::atomic::Ordering::SeqCst),
2
);
assert_eq!(ctx.get_meta("before_validate"), Some(&"1".to_string()));
assert!(after_call_was("after_validate"));
}
#[test]
fn hook_dispatcher_delete_full_sequence() {
let mut ctx = HookContext::new();
let result =
HookDispatcher::delete::<DispatchTestModel, _>(&mut ctx, &42_i64, |_ctx| Ok(()));
assert!(result.is_ok());
}
#[test]
fn hook_dispatcher_restore_full_sequence() {
let mut ctx = HookContext::new();
let result =
HookDispatcher::restore::<DispatchTestModel, _>(&mut ctx, &42_i64, |_ctx| Ok(()));
assert!(result.is_ok());
}
#[test]
fn hook_dispatcher_find_full_sequence() {
let _guard = HOOK_TEST_LOCK.lock().unwrap();
dispatch_calls().store(0, std::sync::atomic::Ordering::SeqCst);
reset_after_calls();
let mut ctx = HookContext::new();
let called = Arc::new(std::sync::atomic::AtomicU32::new(0));
let c = Arc::clone(&called);
let result = HookDispatcher::find::<DispatchTestModel, _>(&mut ctx, &42_i64, move |_ctx| {
c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Ok(())
});
assert!(result.is_ok());
assert_eq!(
called.load(std::sync::atomic::Ordering::SeqCst),
1,
"SELECT 操作应执行一次"
);
assert_eq!(ctx.get_meta("before_find"), Some(&"1".to_string()));
assert!(after_call_was("after_find"));
}
#[test]
fn hook_dispatcher_find_short_circuit_on_before_find_error() {
struct FindFailModel;
impl crate::model::Model for FindFailModel {
type PrimaryKey = i64;
fn table_name() -> &'static str {
"find_fail"
}
fn pk(&self) -> Self::PrimaryKey {
0
}
fn set_pk(&mut self, _pk: Self::PrimaryKey) {}
}
impl Hookable for FindFailModel {
fn before_find(_ctx: &mut HookContext, _id: &Self::PrimaryKey) -> HookResult<()> {
Err(DbError::Hook("before_find blocked".into()))
}
}
let mut ctx = HookContext::new();
let called = Arc::new(std::sync::atomic::AtomicU32::new(0));
let c = Arc::clone(&called);
let result = HookDispatcher::find::<FindFailModel, _>(&mut ctx, &1_i64, move |_ctx| {
c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Ok(())
});
assert!(result.is_err());
assert_eq!(
called.load(std::sync::atomic::Ordering::SeqCst),
0,
"before_find 失败应短路 SELECT"
);
}
#[test]
fn hook_dispatcher_validate_standalone() {
reset_after_calls();
let mut ctx = HookContext::new();
let result = HookDispatcher::validate::<DispatchTestModel>(&mut ctx);
assert!(result.is_ok());
assert_eq!(ctx.get_meta("before_validate"), Some(&"1".to_string()));
assert!(after_call_was("after_validate"));
}
#[test]
fn hook_event_is_find_level_and_is_validate_level() {
assert!(HookEvent::BeforeFind.is_find_level());
assert!(HookEvent::AfterFind.is_find_level());
assert!(HookEvent::BeforeValidate.is_validate_level());
assert!(HookEvent::AfterValidate.is_validate_level());
assert!(!HookEvent::BeforeInsert.is_find_level());
assert!(!HookEvent::BeforeInsert.is_validate_level());
assert!(!HookEvent::BeforeWrite.is_find_level());
assert!(!HookEvent::BeforeWrite.is_validate_level());
}
#[test]
fn hook_event_is_fine_grained_covers_all_v02_events() {
assert!(HookEvent::BeforeWrite.is_fine_grained());
assert!(HookEvent::AfterWrite.is_fine_grained());
assert!(HookEvent::BeforeSave.is_fine_grained());
assert!(HookEvent::AfterSave.is_fine_grained());
assert!(HookEvent::BeforeRestore.is_fine_grained());
assert!(HookEvent::AfterRestore.is_fine_grained());
assert!(HookEvent::BeforeFind.is_fine_grained());
assert!(HookEvent::AfterFind.is_fine_grained());
assert!(HookEvent::BeforeValidate.is_fine_grained());
assert!(HookEvent::AfterValidate.is_fine_grained());
assert!(!HookEvent::BeforeInsert.is_fine_grained());
assert!(!HookEvent::AfterInsert.is_fine_grained());
assert!(!HookEvent::BeforeUpdate.is_fine_grained());
assert!(!HookEvent::AfterUpdate.is_fine_grained());
assert!(!HookEvent::BeforeDelete.is_fine_grained());
assert!(!HookEvent::AfterDelete.is_fine_grained());
}
#[test]
fn hook_registry_supports_find_and_validate_events() {
let registry = HookRegistry::new();
let counter = Arc::new(std::sync::atomic::AtomicU32::new(0));
for event in [
HookEvent::BeforeFind,
HookEvent::AfterFind,
HookEvent::BeforeValidate,
HookEvent::AfterValidate,
] {
let c = Arc::clone(&counter);
registry.register(
event,
Arc::new(move |_ctx| {
c.fetch_add(1, std::sync::atomic::Ordering::SeqCst);
Ok(())
}),
);
}
let ctx = HookContext::new();
for event in [
HookEvent::BeforeFind,
HookEvent::AfterFind,
HookEvent::BeforeValidate,
HookEvent::AfterValidate,
] {
registry.dispatch(event, &ctx).unwrap();
}
assert_eq!(
counter.load(std::sync::atomic::Ordering::SeqCst),
4,
"find/validate 钩子应能被注册与触发"
);
}
#[test]
fn db_error_validation_error_code_and_display() {
let err = DbError::Validation("name required".into());
assert_eq!(err.error_code(), "DB021");
assert_eq!(format!("{}", err), "Validation error: name required");
}
}