use std::{
collections::HashSet,
sync::{
atomic::{AtomicUsize, Ordering},
LazyLock, Mutex, RwLock,
},
};
#[cfg(feature = "sync")]
use std::sync::Arc;
#[cfg(feature = "sync")]
use crossbeam::channel;
use crate::messages::{OutgoingMessages, ResponseMessage};
use crate::transport::routing::{classify_error, determine_routing, ErrorDisposition, RoutingDecision};
use crate::transport::RoutedItem;
use crate::Error;
#[cfg(feature = "sync")]
use crate::transport::{InternalSubscription, MessageBus, SubscriptionBuilder};
#[cfg(feature = "async")]
use {
crate::transport::{
r#async::{AsyncInternalSubscription, CleanupSignal},
AsyncMessageBus,
},
async_trait::async_trait,
tokio::sync::broadcast,
};
#[cfg(feature = "async")]
const TEST_BROADCAST_CAPACITY: usize = 1024;
pub(crate) struct MessageBusStub {
pub request_messages: RwLock<Vec<Vec<u8>>>,
pub response_messages: Vec<String>,
pub ordered_responses: Vec<ResponseMessage>,
connection_resets: AtomicUsize,
}
static ORDER_UPDATE_SUBSCRIPTION_TRACKER: LazyLock<Mutex<HashSet<usize>>> = LazyLock::new(|| Mutex::new(HashSet::new()));
impl Default for MessageBusStub {
fn default() -> Self {
Self {
request_messages: RwLock::new(vec![]),
response_messages: vec![],
ordered_responses: vec![],
connection_resets: AtomicUsize::new(0),
}
}
}
impl Drop for MessageBusStub {
fn drop(&mut self) {
let stub_id = self as *const _ as usize;
ORDER_UPDATE_SUBSCRIPTION_TRACKER.lock().unwrap().remove(&stub_id);
}
}
impl MessageBusStub {
pub fn with_responses(response_messages: Vec<String>) -> Self {
Self {
request_messages: RwLock::new(vec![]),
response_messages,
ordered_responses: vec![],
connection_resets: AtomicUsize::new(0),
}
}
pub fn with_ordered_responses(ordered_responses: Vec<ResponseMessage>) -> Self {
Self {
request_messages: RwLock::new(vec![]),
response_messages: vec![],
ordered_responses,
connection_resets: AtomicUsize::new(0),
}
}
pub fn with_connection_resets(self, count: usize) -> Self {
self.connection_resets.store(count, Ordering::SeqCst);
self
}
pub fn request_messages(&self) -> Vec<Vec<u8>> {
self.request_messages.read().unwrap().clone()
}
pub(crate) fn response_messages_decoded(&self) -> Vec<ResponseMessage> {
if !self.ordered_responses.is_empty() {
return self.ordered_responses.clone();
}
self.response_messages
.iter()
.map(|m| ResponseMessage::from(&m.replace('|', "\0")))
.collect()
}
pub(crate) fn routed_items(&self) -> Vec<RoutedItem> {
self.response_messages_decoded().into_iter().map(classify_like_dispatcher).collect()
}
fn routed_items_for_request(&self) -> Vec<RoutedItem> {
let remaining = self.connection_resets.load(Ordering::SeqCst);
if remaining > 0 {
self.connection_resets.store(remaining - 1, Ordering::SeqCst);
return vec![RoutedItem::Error(Error::ConnectionReset)];
}
self.routed_items()
}
#[cfg(feature = "async")]
fn seeded_subscription(&self, message: Vec<u8>) -> AsyncInternalSubscription {
self.request_messages.write().unwrap().push(message);
let (sender, receiver) = broadcast::channel(TEST_BROADCAST_CAPACITY);
for item in self.routed_items_for_request() {
sender.send(item).unwrap();
}
AsyncInternalSubscription::new(receiver)
}
}
fn classify_like_dispatcher(message: ResponseMessage) -> RoutedItem {
match determine_routing(&message) {
RoutingDecision::Error(payload) => match classify_error(payload) {
ErrorDisposition::Route(_, item) => item,
ErrorDisposition::NoticeOnly(notice) => RoutedItem::Notice(notice),
ErrorDisposition::NoticeAndFailOneShots(_, error) => RoutedItem::Error(error),
},
RoutingDecision::Shutdown => RoutedItem::Error(Error::Shutdown),
_ => message.into(),
}
}
#[cfg(feature = "sync")]
impl MessageBus for MessageBusStub {
fn send_request(&self, request_id: i32, message: &[u8]) -> Result<InternalSubscription, Error> {
Ok(mock_request(self, Some(request_id), None, message))
}
fn cancel_subscription(&self, _request_id: i32, packet: &[u8]) -> Result<(), Error> {
self.request_messages.write().unwrap().push(packet.to_vec());
Ok(())
}
fn send_order_request(&self, request_id: i32, message: &[u8]) -> Result<InternalSubscription, Error> {
Ok(mock_request(self, Some(request_id), None, message))
}
fn send_message(&self, message: &[u8]) -> Result<(), Error> {
self.request_messages.write().unwrap().push(message.to_vec());
Ok(())
}
fn create_order_update_subscription(&self) -> Result<InternalSubscription, Error> {
let stub_id = self as *const _ as usize;
let mut tracker = ORDER_UPDATE_SUBSCRIPTION_TRACKER.lock().unwrap();
if !tracker.insert(stub_id) {
return Err(Error::AlreadySubscribed);
}
drop(tracker);
let (sender, receiver) = channel::unbounded();
let (signaler, _) = channel::unbounded();
for item in self.routed_items() {
sender.send(item).unwrap();
}
let subscription = SubscriptionBuilder::new().receiver(receiver).signaler(signaler).build();
Ok(subscription)
}
fn cancel_order_subscription(&self, _request_id: i32, packet: &[u8]) -> Result<(), Error> {
self.request_messages.write().unwrap().push(packet.to_vec());
let stub_id = self as *const _ as usize;
ORDER_UPDATE_SUBSCRIPTION_TRACKER.lock().unwrap().remove(&stub_id);
Ok(())
}
fn send_shared_request(&self, message_type: OutgoingMessages, message: &[u8]) -> Result<InternalSubscription, Error> {
Ok(mock_request(self, None, Some(message_type), message))
}
fn cancel_shared_subscription(&self, _message_type: OutgoingMessages, packet: &[u8]) -> Result<(), Error> {
self.request_messages.write().unwrap().push(packet.to_vec());
Ok(())
}
fn notice_subscribe(&self) -> crate::subscriptions::notice_stream::sync_impl::NoticeStream {
let (_sender, receiver) = channel::unbounded();
crate::subscriptions::notice_stream::sync_impl::NoticeStream::new(receiver)
}
fn ensure_shutdown(&self) {}
fn is_connected(&self) -> bool {
true }
}
#[cfg(feature = "sync")]
fn mock_request(stub: &MessageBusStub, request_id: Option<i32>, message_type: Option<OutgoingMessages>, message: &[u8]) -> InternalSubscription {
stub.request_messages.write().unwrap().push(message.to_vec());
let (sender, receiver) = channel::unbounded();
let (s1, _r1) = channel::unbounded();
for item in stub.routed_items_for_request() {
sender.send(item).unwrap();
}
let mut subscription = SubscriptionBuilder::new().signaler(s1);
if let Some(request_id) = request_id {
subscription = subscription.receiver(receiver).request_id(request_id);
} else if let Some(message_type) = message_type {
subscription = subscription.shared_receiver(Arc::new(receiver)).message_type(message_type);
}
subscription.build()
}
#[cfg(feature = "async")]
#[async_trait]
impl AsyncMessageBus for MessageBusStub {
async fn send_request(&self, _request_id: i32, message: Vec<u8>) -> Result<AsyncInternalSubscription, Error> {
Ok(self.seeded_subscription(message))
}
async fn send_order_request(&self, _order_id: i32, message: Vec<u8>) -> Result<AsyncInternalSubscription, Error> {
Ok(self.seeded_subscription(message))
}
async fn send_shared_request(&self, _message_type: OutgoingMessages, message: Vec<u8>) -> Result<AsyncInternalSubscription, Error> {
Ok(self.seeded_subscription(message))
}
async fn send_message(&self, message: Vec<u8>) -> Result<(), Error> {
self.request_messages.write().unwrap().push(message);
Ok(())
}
async fn cancel_subscription(&self, _request_id: i32, message: Vec<u8>) -> Result<(), Error> {
self.request_messages.write().unwrap().push(message);
Ok(())
}
async fn cancel_order_subscription(&self, _order_id: i32, _message: Vec<u8>) -> Result<(), Error> {
Ok(())
}
async fn create_order_update_subscription(&self) -> Result<AsyncInternalSubscription, Error> {
let stub_id = self as *const _ as usize;
let mut tracker = ORDER_UPDATE_SUBSCRIPTION_TRACKER.lock().unwrap();
if !tracker.insert(stub_id) {
return Err(Error::AlreadySubscribed);
}
drop(tracker);
let (sender, receiver) = broadcast::channel(TEST_BROADCAST_CAPACITY);
for item in self.routed_items() {
sender.send(item).unwrap();
}
let (cleanup_sender, mut cleanup_receiver) = tokio::sync::mpsc::unbounded_channel();
tokio::spawn(async move {
while let Some(signal) = cleanup_receiver.recv().await {
if matches!(signal, CleanupSignal::OrderUpdateStream) {
ORDER_UPDATE_SUBSCRIPTION_TRACKER.lock().unwrap().remove(&stub_id);
break;
}
}
});
Ok(AsyncInternalSubscription::with_cleanup(
receiver,
cleanup_sender,
CleanupSignal::OrderUpdateStream,
))
}
fn notice_subscribe(&self) -> crate::subscriptions::notice_stream::async_impl::NoticeStream {
let (_sender, receiver) = broadcast::channel(1);
crate::subscriptions::notice_stream::async_impl::NoticeStream::new(receiver)
}
async fn ensure_shutdown(&self) {
}
fn request_shutdown_sync(&self) {
}
fn is_connected(&self) -> bool {
true }
}
#[cfg(test)]
#[path = "stubs_tests.rs"]
mod tests;