use async_trait::async_trait;
use futures_util::stream::SplitSink;
use futures_util::{SinkExt, StreamExt};
use serde::Serialize;
use serde_json::{Map, to_string, to_value, Value};
use thiserror::Error;
use tokio::sync::broadcast::{channel, Receiver, Sender};
use tokio::task::JoinHandle;
use tokio_tungstenite::connect_async;
use tokio_tungstenite::tungstenite::Message;
use tracing::{debug, error, info};
use url::Url;
use crate::schema::{StreamDataMessage, SubscribeRequest, UnsubscribeRequest};
use crate::listener::{listen_for_stream_data, Stream, StreamDataMessageHandler};
#[async_trait]
pub trait XtbStreamConnection {
type MessageStream: MessageStream;
type Error;
async fn subscribe(&mut self, command: &str, arguments: Option<Value>) -> Result<(), Self::Error>;
async fn unsubscribe(&mut self, command: &str, arguments: Option<Value>) -> Result<(), Self::Error>;
async fn make_message_stream(&mut self, filter: DataMessageFilter) -> Self::MessageStream;
}
#[derive(Debug)]
pub struct BasicXtbStreamConnection {
stream_session_id: String,
sender: Sender<StreamDataMessage>,
sink: SplitSink<Stream, Message>,
listener_join: JoinHandle<()>,
}
impl BasicXtbStreamConnection {
pub async fn new(url: Url, stream_session_id: String) -> Result<Self, BasicXtbStreamConnectionError> {
let (sender, _) = channel(64usize);
let host_clone = url.as_str().to_owned();
let (conn, _) = connect_async(url).await.map_err(|_| BasicXtbStreamConnectionError::CannotConnect(host_clone))?;
let (sink, stream) = conn.split();
let listener_join = listen_for_stream_data(stream, MessageHandler::new(sender.clone()));
Ok(Self {
stream_session_id,
sender,
sink,
listener_join,
})
}
async fn assemble_and_send<T: Serialize>(&mut self, request: T, arguments: Option<Value>) -> Result<(), BasicXtbStreamConnectionError> {
let mut obj = to_value(request).map_err(|err| BasicXtbStreamConnectionError::SerializationFailed(err))?;
let prepared_arguments = Self::prepare_arguments(arguments)?;
if let Some(mut prepared_obj) = prepared_arguments {
obj.as_object_mut().unwrap().append(&mut prepared_obj);
}
let serialized = to_string(&obj).map_err(|err| BasicXtbStreamConnectionError::SerializationFailed(err))?;
let message = Message::text(serialized);
self.sink.send(message).await.map_err(|err| BasicXtbStreamConnectionError::CannotSend(err))
}
fn prepare_arguments(arguments: Option<Value>) -> Result<Option<Map<String, Value>>, BasicXtbStreamConnectionError> {
match arguments {
None => Ok(None),
Some(Value::Object(obj)) => Ok(Some(obj)),
Some(Value::Null) => Ok(None),
_ => Err(BasicXtbStreamConnectionError::InvalidArgumentsType)
}
}
}
impl Drop for BasicXtbStreamConnection {
fn drop(&mut self) {
self.listener_join.abort();
}
}
#[async_trait]
impl XtbStreamConnection for BasicXtbStreamConnection {
type MessageStream = BasicMessageStream;
type Error = BasicXtbStreamConnectionError;
async fn subscribe(&mut self, command: &str, arguments: Option<Value>) -> Result<(), Self::Error> {
let request = SubscribeRequest::default()
.with_command(command)
.with_stream_session_id(&self.stream_session_id);
info!("Subscribing for {command}");
debug!("Subscription arguments are {arguments:?}");
self.assemble_and_send(request, arguments).await
}
async fn unsubscribe(&mut self, command: &str, arguments: Option<Value>) -> Result<(), Self::Error> {
let request = UnsubscribeRequest::default().with_command(command);
info!("Unsubscribing from {command}");
debug!("Unsubscription arguments are {arguments:?}");
self.assemble_and_send(request, arguments).await
}
async fn make_message_stream(&mut self, filter: DataMessageFilter) -> Self::MessageStream {
BasicMessageStream::new(filter, self.sender.subscribe())
}
}
struct MessageHandler {
sender: Sender<StreamDataMessage>,
}
impl MessageHandler {
pub fn new(sender: Sender<StreamDataMessage>) -> Self {
Self { sender }
}
}
#[async_trait]
impl StreamDataMessageHandler for MessageHandler {
async fn handle_message(&self, message: StreamDataMessage) {
let cmd = message.command.to_owned();
info!("Handling incoming message {cmd}");
debug!("Incoming message: {message:?}");
match self.sender.send(message) {
Err(err) => error!("Cannot broadcast message: {}", err),
_ => debug!("Message {cmd} was broadcast to the {} receivers", self.sender.len())
}
}
}
#[derive(Default)]
pub enum DataMessageFilter {
#[default]
Always,
Never,
Command(String),
FieldValue { name: String, value: Value },
Custom(Box<dyn Fn(&StreamDataMessage) -> bool + Send + Sync>),
All(Vec<DataMessageFilter>),
Any(Vec<DataMessageFilter>),
}
impl DataMessageFilter {
pub fn test_message(&self, msg: &StreamDataMessage) -> bool {
match self {
Self::Always => Self::resolve_always(msg),
Self::Never => Self::resolve_never(msg),
Self::Command(cmd) => Self::resolve_command(msg, cmd),
Self::All(ops) => Self::resolve_all(msg, ops),
Self::Any(ops) => Self::resolve_any(msg, ops),
Self::FieldValue { name, value } => Self::resolve_field_value(msg, name, value),
Self::Custom(cbk) => Self::resolve_custom(msg, cbk),
}
}
fn resolve_always(_: &StreamDataMessage) -> bool {
true
}
fn resolve_never(_: &StreamDataMessage) -> bool {
false
}
fn resolve_command(msg: &StreamDataMessage, command: &str) -> bool {
return msg.command.as_str() == command;
}
fn resolve_all(msg: &StreamDataMessage, ops: &Vec<DataMessageFilter>) -> bool {
ops.iter().all(|f| f.test_message(msg))
}
fn resolve_any(msg: &StreamDataMessage, ops: &Vec<DataMessageFilter>) -> bool {
ops.iter().any(|f| f.test_message(msg))
}
fn resolve_field_value(msg: &StreamDataMessage, field_name: &str, field_value: &Value) -> bool {
match &msg.data {
Value::Object(data_obj) => {
if let Some(field_content) = data_obj.get(field_name) {
field_content == field_value
} else {
false
}
}
_ => false
}
}
fn resolve_custom(msg: &StreamDataMessage, cbk: &Box<dyn Fn(&StreamDataMessage) -> bool + Send + Sync>) -> bool {
(*cbk)(msg)
}
}
#[derive(Debug, Error)]
pub enum BasicXtbStreamConnectionError {
#[error("Cannot connect to server ({0}")]
CannotConnect(String),
#[error("Cannot send message")]
CannotSend(tokio_tungstenite::tungstenite::Error),
#[error("Cannot serialize data")]
SerializationFailed(serde_json::Error),
#[error("Only Value::Object can be used for the arguments")]
InvalidArgumentsType,
}
#[async_trait]
pub trait MessageStream {
async fn next(&mut self) -> Option<StreamDataMessage>;
}
pub struct BasicMessageStream {
filter: DataMessageFilter,
stream: Receiver<StreamDataMessage>,
}
impl BasicMessageStream {
pub fn new(filter: DataMessageFilter, stream: Receiver<StreamDataMessage>) -> Self {
BasicMessageStream {
filter,
stream,
}
}
}
#[async_trait]
impl MessageStream for BasicMessageStream {
async fn next(&mut self) -> Option<StreamDataMessage> {
while let Some(msg) = self.stream.recv().await.ok() {
if self.filter.test_message(&msg) {
return Some(msg);
}
}
None
}
}
#[cfg(test)]
mod tests {
mod data_message_filter {
use rstest::rstest;
use serde_json::{from_str, Value};
use crate::schema::StreamDataMessage;
use crate::DataMessageFilter;
#[test]
fn always() {
let msg = StreamDataMessage::default();
assert!(DataMessageFilter::Always.test_message(&msg));
}
#[test]
fn never() {
let msg = StreamDataMessage::default();
assert!(!DataMessageFilter::Never.test_message(&msg));
}
#[rstest]
#[case("command", true)]
#[case("other_command", false)]
fn command(#[case] cmd: &str, #[case] expected_result: bool) {
let msg = StreamDataMessage { command: "command".to_string(), data: Value::Null };
assert_eq!(DataMessageFilter::Command(cmd.to_string()).test_message(&msg), expected_result);
}
#[rstest]
#[case(vec ! [], true)]
#[case(vec ! [DataMessageFilter::Always], true)]
#[case(vec ! [DataMessageFilter::Never], false)]
#[case(vec ! [DataMessageFilter::Command("command".to_string()), DataMessageFilter::Always], true)]
#[case(vec ! [DataMessageFilter::Command("command".to_string()), DataMessageFilter::Never], false)]
#[case(vec ! [DataMessageFilter::Command("other_command".to_string()), DataMessageFilter::Never], false)]
#[case(vec ! [DataMessageFilter::Command("other_command".to_string()), DataMessageFilter::Always], false)]
fn all(#[case] filters: Vec<DataMessageFilter>, #[case] expected_result: bool) {
let msg = StreamDataMessage { command: "command".to_owned(), data: Value::Null };
let f = DataMessageFilter::All(filters);
assert_eq!(f.test_message(&msg), expected_result);
}
#[rstest]
#[case(vec ! [], false)]
#[case(vec ! [DataMessageFilter::Always], true)]
#[case(vec ! [DataMessageFilter::Never], false)]
#[case(vec ! [DataMessageFilter::Command("command".to_string()), DataMessageFilter::Always], true)]
#[case(vec ! [DataMessageFilter::Command("command".to_string()), DataMessageFilter::Never], true)]
#[case(vec ! [DataMessageFilter::Command("other_command".to_string()), DataMessageFilter::Never], false)]
#[case(vec ! [DataMessageFilter::Command("other_command".to_string()), DataMessageFilter::Always], true)]
fn any(#[case] filters: Vec<DataMessageFilter>, #[case] expected_result: bool) {
let msg = StreamDataMessage { command: "command".to_owned(), data: Value::Null };
let f = DataMessageFilter::Any(filters);
assert_eq!(f.test_message(&msg), expected_result);
}
#[rstest]
#[case(r#"{"field": "value"}"#, true)]
#[case(r#"{"field": 10}"#, false)]
#[case(r#"{"other_field": 10}"#, false)]
#[case(r#"null"#, false)]
fn filed_value(#[case] source_data: &str, #[case] expected_value: bool) {
let data: Value = from_str(source_data).unwrap();
let msg = StreamDataMessage { data, command: "".to_owned() };
let f = DataMessageFilter::FieldValue { name: "field".to_owned(), value: Value::String("value".to_owned()) };
assert_eq!(f.test_message(&msg), expected_value)
}
#[test]
fn custom_true() {
let msg = StreamDataMessage::default();
let f = DataMessageFilter::Custom(Box::new(|msg| true));
assert_eq!(f.test_message(&msg), true)
}
#[test]
fn custom_false() {
let msg = StreamDataMessage::default();
let f = DataMessageFilter::Custom(Box::new(|msg| false));
assert_eq!(f.test_message(&msg), false)
}
}
}