1use std::sync::{Arc, RwLock};
9
10use crate::atom::Atom;
11
12#[derive(Copy, Clone, Debug, Eq, PartialEq)]
14pub struct HookEvent {
15 pub pid: u64,
17 pub module: Atom,
19 pub function: Atom,
21 pub arity: u8,
23 pub reductions_consumed: u32,
25}
26
27#[derive(Copy, Clone, Debug, Eq, PartialEq)]
29pub enum HookDecision {
30 Continue,
32 Suspend,
34}
35
36type HookCallback = dyn Fn(HookEvent) -> HookDecision + Send + Sync + 'static;
37
38#[derive(Clone, Default)]
40pub struct Hook {
41 callback: Arc<RwLock<Option<Arc<HookCallback>>>>,
42}
43
44impl Hook {
45 #[must_use]
47 pub fn new() -> Self {
48 Self::default()
49 }
50
51 pub fn register<F>(&self, callback: F)
53 where
54 F: Fn(HookEvent) -> HookDecision + Send + Sync + 'static,
55 {
56 let mut slot = self
57 .callback
58 .write()
59 .unwrap_or_else(|error| error.into_inner());
60 *slot = Some(Arc::new(callback));
61 }
62
63 pub fn unregister(&self) {
65 let mut slot = self
66 .callback
67 .write()
68 .unwrap_or_else(|error| error.into_inner());
69 *slot = None;
70 }
71
72 #[must_use]
74 pub fn is_registered(&self) -> bool {
75 self.callback
76 .read()
77 .unwrap_or_else(|error| error.into_inner())
78 .is_some()
79 }
80
81 #[must_use]
85 pub fn invoke(&self, event: HookEvent) -> HookDecision {
86 let callback = self
87 .callback
88 .read()
89 .unwrap_or_else(|error| error.into_inner())
90 .clone();
91 match callback {
92 Some(callback) => callback(event),
93 None => HookDecision::Continue,
94 }
95 }
96}
97
98#[cfg(test)]
99mod tests {
100 use std::sync::{Arc, Mutex};
101
102 use super::{Hook, HookDecision, HookEvent};
103 use crate::atom::Atom;
104
105 fn event() -> HookEvent {
106 HookEvent {
107 pid: 7,
108 module: Atom::OK,
109 function: Atom::ERROR,
110 arity: 2,
111 reductions_consumed: 42,
112 }
113 }
114
115 #[test]
116 fn hook_register_replace_and_unregister_hook() {
117 let hook = Hook::new();
118 assert!(!hook.is_registered());
119 assert_eq!(hook.invoke(event()), HookDecision::Continue);
120
121 hook.register(|_| HookDecision::Suspend);
122 assert!(hook.is_registered());
123 assert_eq!(hook.invoke(event()), HookDecision::Suspend);
124
125 hook.register(|_| HookDecision::Continue);
126 assert_eq!(hook.invoke(event()), HookDecision::Continue);
127
128 hook.unregister();
129 assert!(!hook.is_registered());
130 assert_eq!(hook.invoke(event()), HookDecision::Continue);
131 }
132
133 #[test]
134 fn hook_receives_copied_metadata_at_yield() {
135 let hook = Hook::new();
136 let seen = Arc::new(Mutex::new(Vec::new()));
137 let seen_by_hook = Arc::clone(&seen);
138 hook.register(move |event| {
139 seen_by_hook
140 .lock()
141 .unwrap_or_else(|error| error.into_inner())
142 .push(event);
143 HookDecision::Continue
144 });
145
146 assert_eq!(hook.invoke(event()), HookDecision::Continue);
147
148 assert_eq!(
149 seen.lock()
150 .unwrap_or_else(|error| error.into_inner())
151 .as_slice(),
152 &[event()]
153 );
154 }
155
156 #[test]
157 fn unregistered_hook_does_not_call_previous_callback() {
158 let hook = Hook::new();
159 let calls = Arc::new(Mutex::new(0_u32));
160 let calls_by_hook = Arc::clone(&calls);
161 hook.register(move |_| {
162 *calls_by_hook
163 .lock()
164 .unwrap_or_else(|error| error.into_inner()) += 1;
165 HookDecision::Continue
166 });
167 hook.unregister();
168
169 assert_eq!(hook.invoke(event()), HookDecision::Continue);
170 assert_eq!(*calls.lock().unwrap_or_else(|error| error.into_inner()), 0);
171 }
172}