Skip to main content

nautilus_common/msgbus/
stubs.rs

1// -------------------------------------------------------------------------------------------------
2//  Copyright (C) 2015-2026 Nautech Systems Pty Ltd. All rights reserved.
3//  https://nautechsystems.io
4//
5//  Licensed under the GNU Lesser General Public License Version 3.0 (the "License");
6//  You may not use this file except in compliance with the License.
7//  You may obtain a copy of the License at https://www.gnu.org/licenses/lgpl-3.0.en.html
8//
9//  Unless required by applicable law or agreed to in writing, software
10//  distributed under the License is distributed on an "AS IS" BASIS,
11//  WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
12//  See the License for the specific language governing permissions and
13//  limitations under the License.
14// -------------------------------------------------------------------------------------------------
15
16use std::{
17    any::Any,
18    cell::RefCell,
19    rc::Rc,
20    sync::{
21        Arc,
22        atomic::{AtomicBool, Ordering},
23    },
24};
25
26use ahash::AHashMap;
27use nautilus_core::UUID4;
28use ustr::Ustr;
29
30use crate::msgbus::{
31    Handler, IntoHandler, ShareableMessageHandler, TypedHandler, TypedIntoHandler,
32    typed_handler::shareable_handler,
33};
34
35/// Minimal message payload for message bus tests.
36#[derive(Debug, Clone)]
37pub struct StubMessage;
38
39/// No-op handler for message bus tests.
40#[derive(Debug, Clone)]
41pub struct StubMessageHandler {
42    id: Ustr,
43}
44
45impl Handler<dyn Any> for StubMessageHandler {
46    fn id(&self) -> Ustr {
47        self.id
48    }
49
50    fn handle(&self, _message: &dyn Any) {}
51}
52
53#[must_use]
54pub fn get_stub_shareable_handler(id: Option<Ustr>) -> ShareableMessageHandler {
55    let unique_id = id.unwrap_or_else(|| Ustr::from(UUID4::new().as_str()));
56    shareable_handler(Rc::new(StubMessageHandler { id: unique_id }))
57}
58
59/// Handler that tracks whether it has been called.
60#[derive(Debug, Clone)]
61pub struct CallCheckHandler {
62    id: Ustr,
63    called: Arc<AtomicBool>,
64}
65
66impl CallCheckHandler {
67    #[must_use]
68    pub fn new(id: Option<Ustr>) -> Self {
69        let unique_id = id.unwrap_or_else(|| Ustr::from(UUID4::new().as_str()));
70        Self {
71            id: unique_id,
72            called: Arc::new(AtomicBool::new(false)),
73        }
74    }
75
76    #[must_use]
77    pub fn was_called(&self) -> bool {
78        self.called.load(Ordering::SeqCst)
79    }
80
81    /// Returns a `ShareableMessageHandler` for registration.
82    #[must_use]
83    pub fn handler(&self) -> ShareableMessageHandler {
84        shareable_handler(Rc::new(self.clone()))
85    }
86}
87
88impl Handler<dyn Any> for CallCheckHandler {
89    fn id(&self) -> Ustr {
90        self.id
91    }
92
93    fn handle(&self, _message: &dyn Any) {
94        self.called.store(true, Ordering::SeqCst);
95    }
96}
97
98/// Creates a call-checking handler and returns both the handler for registration
99/// and a clone that can be used to check if it was called.
100#[must_use]
101pub fn get_call_check_handler(id: Option<Ustr>) -> (ShareableMessageHandler, CallCheckHandler) {
102    let checker = CallCheckHandler::new(id);
103    let handler = checker.handler();
104    (handler, checker)
105}
106
107/// Handler that saves messages it receives (for Any-based routing).
108#[derive(Debug, Clone)]
109pub struct AnySavingHandler<T> {
110    id: Ustr,
111    messages: Rc<RefCell<Vec<T>>>,
112}
113
114impl<T: Clone + 'static> AnySavingHandler<T> {
115    #[must_use]
116    pub fn new(id: Option<Ustr>) -> Self {
117        let unique_id = id.unwrap_or_else(|| Ustr::from(UUID4::new().as_str()));
118        Self {
119            id: unique_id,
120            messages: Rc::new(RefCell::new(Vec::new())),
121        }
122    }
123
124    #[must_use]
125    pub fn get_messages(&self) -> Vec<T> {
126        self.messages.borrow().clone()
127    }
128
129    pub fn clear(&self) {
130        self.messages.borrow_mut().clear();
131    }
132
133    /// Returns a `ShareableMessageHandler` for registration.
134    #[must_use]
135    pub fn handler(&self) -> ShareableMessageHandler {
136        shareable_handler(Rc::new(self.clone()))
137    }
138}
139
140impl<T: Clone + 'static> Handler<dyn Any> for AnySavingHandler<T> {
141    fn id(&self) -> Ustr {
142        self.id
143    }
144
145    fn handle(&self, message: &dyn Any) {
146        if let Some(m) = message.downcast_ref::<T>() {
147            self.messages.borrow_mut().push(m.clone());
148        } else {
149            log::error!(
150                "AnySavingHandler: expected {} got {:?}",
151                std::any::type_name::<T>(),
152                message.type_id()
153            );
154        }
155    }
156}
157
158/// Creates an Any-based message saving handler and returns both the handler
159/// for registration and a clone that can be used to retrieve messages.
160#[must_use]
161pub fn get_any_saving_handler<T: Clone + 'static>(
162    id: Option<Ustr>,
163) -> (ShareableMessageHandler, AnySavingHandler<T>) {
164    let saver = AnySavingHandler::new(id);
165    let handler = saver.handler();
166    (handler, saver)
167}
168
169// Type alias for backward compatibility
170pub type MessageSavingHandler<T> = AnySavingHandler<T>;
171
172/// Typed handler which saves the messages it receives (no downcast needed).
173#[derive(Debug, Clone)]
174pub struct TypedMessageSavingHandler<T> {
175    id: Ustr,
176    messages: Rc<RefCell<Vec<T>>>,
177}
178
179impl<T: Clone + 'static> TypedMessageSavingHandler<T> {
180    #[must_use]
181    pub fn new(id: Option<Ustr>) -> Self {
182        let unique_id = id.unwrap_or_else(|| Ustr::from(UUID4::new().as_str()));
183        Self {
184            id: unique_id,
185            messages: Rc::new(RefCell::new(Vec::new())),
186        }
187    }
188
189    #[must_use]
190    pub fn get_messages(&self) -> Vec<T> {
191        self.messages.borrow().clone()
192    }
193
194    /// Returns a `TypedHandler` that can be used for subscriptions.
195    #[must_use]
196    pub fn handler(&self) -> TypedHandler<T> {
197        TypedHandler::new(self.clone())
198    }
199}
200
201impl<T: Clone + 'static> Handler<T> for TypedMessageSavingHandler<T> {
202    fn id(&self) -> Ustr {
203        self.id
204    }
205
206    fn handle(&self, message: &T) {
207        self.messages.borrow_mut().push(message.clone());
208    }
209}
210
211/// Creates a typed message saving handler and returns both the handler for subscriptions
212/// and a clone that can be used to retrieve messages.
213#[must_use]
214pub fn get_typed_message_saving_handler<T: Clone + 'static>(
215    id: Option<Ustr>,
216) -> (TypedHandler<T>, TypedMessageSavingHandler<T>) {
217    let saving_handler = TypedMessageSavingHandler::new(id);
218    let typed_handler = saving_handler.handler();
219    (typed_handler, saving_handler)
220}
221
222/// Ownership-based typed handler which saves the messages it receives.
223///
224/// Unlike [`TypedMessageSavingHandler`] which borrows messages, this handler
225/// takes ownership which is required for `IntoEndpointMap` endpoints.
226#[derive(Debug, Clone)]
227pub struct TypedIntoMessageSavingHandler<T> {
228    id: Ustr,
229    messages: Rc<RefCell<Vec<T>>>,
230}
231
232impl<T: 'static> TypedIntoMessageSavingHandler<T> {
233    #[must_use]
234    pub fn new(id: Option<Ustr>) -> Self {
235        let unique_id = id.unwrap_or_else(|| Ustr::from(UUID4::new().as_str()));
236        Self {
237            id: unique_id,
238            messages: Rc::new(RefCell::new(Vec::new())),
239        }
240    }
241
242    /// Creates a handler backed by an existing shared messages vec.
243    #[must_use]
244    pub fn new_with_messages(id: Option<Ustr>, messages: Rc<RefCell<Vec<T>>>) -> Self {
245        let unique_id = id.unwrap_or_else(|| Ustr::from(UUID4::new().as_str()));
246        Self {
247            id: unique_id,
248            messages,
249        }
250    }
251
252    #[must_use]
253    pub fn get_messages(&self) -> Vec<T>
254    where
255        T: Clone,
256    {
257        self.messages.borrow().clone()
258    }
259
260    /// Returns a `TypedIntoHandler` that can be used for endpoint registration.
261    #[must_use]
262    pub fn handler(&self) -> TypedIntoHandler<T> {
263        TypedIntoHandler::new(Self {
264            id: self.id,
265            messages: self.messages.clone(),
266        })
267    }
268
269    pub fn clear(&self) {
270        self.messages.borrow_mut().clear();
271    }
272}
273
274impl<T: 'static> IntoHandler<T> for TypedIntoMessageSavingHandler<T> {
275    fn id(&self) -> Ustr {
276        self.id
277    }
278
279    fn handle(&self, message: T) {
280        self.messages.borrow_mut().push(message);
281    }
282}
283
284/// Creates an ownership-based typed message saving handler and returns both the handler
285/// for endpoint registration and a clone that can be used to retrieve messages.
286#[must_use]
287pub fn get_typed_into_message_saving_handler<T: 'static>(
288    id: Option<Ustr>,
289) -> (TypedIntoHandler<T>, TypedIntoMessageSavingHandler<T>) {
290    let saving_handler = TypedIntoMessageSavingHandler::new(id);
291    let typed_handler = saving_handler.handler();
292    (typed_handler, saving_handler)
293}
294
295// Legacy API for tests that use the old pattern with thread_local storage.
296// These wrap AnySavingHandler in thread_local for simpler test usage.
297
298thread_local! {
299    static SAVING_HANDLERS: RefCell<AHashMap<Ustr, Box<dyn std::any::Any>>> = RefCell::new(AHashMap::new());
300}
301
302/// Creates a message saving handler and stores it for later retrieval.
303#[must_use]
304pub fn get_message_saving_handler<T: Clone + 'static>(id: Option<Ustr>) -> ShareableMessageHandler {
305    let (handler, saver) = get_any_saving_handler::<T>(id);
306    let handler_id = handler.0.id();
307    SAVING_HANDLERS.with(|handlers| {
308        handlers.borrow_mut().insert(handler_id, Box::new(saver));
309    });
310    handler
311}
312
313/// Retrieves saved messages from a handler created by `get_message_saving_handler`.
314#[must_use]
315pub fn get_saved_messages<T: Clone + 'static>(handler: &ShareableMessageHandler) -> Vec<T> {
316    let handler_id = handler.0.id();
317    SAVING_HANDLERS.with(|handlers| {
318        let handlers = handlers.borrow();
319        if let Some(saver) = handlers.get(&handler_id)
320            && let Some(saver) = saver.downcast_ref::<AnySavingHandler<T>>()
321        {
322            return saver.get_messages();
323        }
324        Vec::new()
325    })
326}