use std::collections::HashMap;
use std::marker::PhantomData;
use std::str::FromStr;
use std::sync::Arc;
use std::time::Duration;
use async_trait::async_trait;
use derive_setters::Setters;
use serde::{Deserialize, Serialize};
use serde_json::{from_value, to_value, Value};
use thiserror::Error;
use tokio::spawn;
use tokio::sync::Mutex;
use tokio::task::JoinHandle;
use tokio::time::sleep;
use tracing::{debug, error};
use url::Url;
use crate::{BasicMessageStream, BasicXtbConnection, BasicXtbStreamConnection, DataMessageFilter, MessageStream, ResponsePromise, XtbConnection, BasicXtbConnectionError, XtbStreamConnection, BasicXtbStreamConnectionError};
use crate::message_processing::ProcessedMessage;
use crate::schema::{COMMAND_GET_ALL_SYMBOLS, COMMAND_GET_CALENDAR, COMMAND_GET_CHART_LAST_REQUEST, COMMAND_GET_CHART_RANGE_REQUEST, COMMAND_GET_COMMISSION_DEF, COMMAND_GET_CURRENT_USER_DATA, COMMAND_GET_IBS_HISTORY, COMMAND_GET_MARGIN_LEVEL, COMMAND_GET_MARGIN_TRADE, COMMAND_GET_NEWS, COMMAND_GET_PROFIT_CALCULATION, COMMAND_GET_SERVER_TIME, COMMAND_GET_STEP_RULES, COMMAND_GET_SYMBOL, COMMAND_GET_TICK_PRICES, COMMAND_GET_TRADE_RECORDS, COMMAND_GET_TRADES, COMMAND_GET_TRADES_HISTORY, COMMAND_GET_TRADING_HOURS, COMMAND_GET_VERSION, COMMAND_LOGIN, COMMAND_PING, COMMAND_TRADE_TRANSACTION, COMMAND_TRADE_TRANSACTION_STATUS, ErrorResponse, GetAllSymbolsRequest, GetAllSymbolsResponse, GetCalendarRequest, GetCalendarResponse, GetChartLastRequestRequest, GetChartLastRequestResponse, GetChartRangeRequestRequest, GetChartRangeRequestResponse, GetCommissionDefRequest, GetCommissionDefResponse, GetCurrentUserDataRequest, GetCurrentUserDataResponse, GetIbsHistoryRequest, GetIbsHistoryResponse, GetMarginLevelRequest, GetMarginLevelResponse, GetMarginTradeRequest, GetMarginTradeResponse, GetNewsRequest, GetNewsResponse, GetProfitCalculationRequest, GetProfitCalculationResponse, GetServerTimeRequest, GetServerTimeResponse, GetStepRulesRequest, GetStepRulesResponse, GetSymbolRequest, GetSymbolResponse, GetTickPricesRequest, GetTickPricesResponse, GetTradeRecordsRequest, GetTradeRecordsResponse, GetTradesHistoryRequest, GetTradesHistoryResponse, GetTradesRequest, GetTradesResponse, GetTradingHoursRequest, GetTradingHoursResponse, GetVersionRequest, GetVersionResponse, LoginRequest, PingRequest, STREAM_BALANCE, STREAM_CANDLES, STREAM_BALANCE_SUBSCRIBE, STREAM_CANDLES_SUBSCRIBE, STREAM_KEEP_ALIVE_SUBSCRIBE, STREAM_NEWS_SUBSCRIBE, STREAM_PROFITS_SUBSCRIBE, STREAM_TICK_PRICES_SUBSCRIBE, STREAM_TRADE_STATUS_SUBSCRIBE, STREAM_TRADES_SUBSCRIBE, STREAM_KEEP_ALIVE, STREAM_NEWS, STREAM_PING, STREAM_PROFITS, STREAM_BALANCE_UNSUBSCRIBE, STREAM_CANDLES_UNSUBSCRIBE, STREAM_KEEP_ALIVE_UNSUBSCRIBE, STREAM_NEWS_UNSUBSCRIBE, STREAM_PROFITS_UNSUBSCRIBE, STREAM_TICK_PRICES_UNSUBSCRIBE, STREAM_TRADE_STATUS_UNSUBSCRIBE, STREAM_TRADES_UNSUBSCRIBE, STREAM_TICK_PRICES, STREAM_TRADE_STATUS, STREAM_TRADES, StreamDataMessage, StreamGetBalanceData, StreamGetBalanceSubscribe, StreamGetBalanceUnsubscribe, StreamGetCandlesData, StreamGetCandlesSubscribe, StreamGetCandlesUnsubscribe, StreamGetKeepAliveData, StreamGetKeepAliveSubscribe, StreamGetKeepAliveUnsubscribe, StreamGetNewsData, StreamGetNewsSubscribe, StreamGetNewsUnsubscribe, StreamGetProfitData, StreamGetProfitSubscribe, StreamGetProfitUnsubscribe, StreamGetTickPricesData, StreamGetTickPricesSubscribe, StreamGetTickPricesUnsubscribe, StreamGetTradesData, StreamGetTradesSubscribe, StreamGetTradeStatusData, StreamGetTradeStatusSubscribe, StreamGetTradeStatusUnsubscribe, StreamGetTradesUnsubscribe, StreamPingSubscribe, TradeTransactionRequest, TradeTransactionResponse, TradeTransactionStatusRequest, TradeTransactionStatusResponse};
#[derive(Default, Setters)]
#[setters(into, prefix = "with_", strip_option)]
pub struct XtbClientBuilder {
api_url: Option<String>,
stream_api_url: Option<String>,
app_id: Option<String>,
app_name: Option<String>,
ping_period: Option<u64>,
}
const DEFAULT_PING_INTERVAL_S: u64 = 30;
const DEFAULT_XTB_REAL: &'static str = "wss://ws.xtb.com/real";
const DEFAULT_XTB_REAL_STREAM: &'static str = "wss://ws.xtb.com/realStream";
const DEFAULT_XTB_DEMO: &'static str = "wss://ws.xtb.com/demo";
const DEFAULT_XTB_DEMO_STREAM: &'static str = "wss://ws.xtb.com/demoStream";
impl XtbClientBuilder {
pub fn new(api_url: &str, stream_api_url: &str) -> Self {
Self {
api_url: Some(api_url.to_string()),
stream_api_url: Some(stream_api_url.to_string()),
app_id: None,
app_name: None,
ping_period: None,
}
}
pub fn new_bare() -> Self {
return Self {
api_url: None,
stream_api_url: None,
app_id: None,
app_name: None,
ping_period: None,
}
}
pub fn new_real() -> Self {
Self::new(DEFAULT_XTB_REAL, DEFAULT_XTB_REAL_STREAM)
}
pub fn new_demo() -> Self {
Self::new(DEFAULT_XTB_DEMO, DEFAULT_XTB_DEMO_STREAM)
}
pub async fn build(self, user_id: &str, password: &str) -> Result<XtbClient, XtbClientBuilderError> {
let api_url = Self::make_url(self.api_url)?;
let stream_api_url = Self::make_url(self.stream_api_url)?;
let mut connection = BasicXtbConnection::new(api_url).await.map_err(|err| XtbClientBuilderError::CannotMakeConnection(err))?;
let mut login_request = LoginRequest::default().with_user_id(user_id).with_password(password);
if let Some(app_id) = self.app_id {
login_request = login_request.with_app_id(app_id);
}
if let Some(app_name) = self.app_name {
login_request = login_request.with_app_name(app_name);
}
let login_request_value = to_value(login_request).map_err(|err| XtbClientBuilderError::UnexpectedError(format!("{:?}", err)))?;
let response = connection
.send_command(COMMAND_LOGIN, Some(login_request_value)).await
.map_err(|err| XtbClientBuilderError::UnexpectedError(format!("{:?}", err)))?.await
.map_err(|err| XtbClientBuilderError::UnexpectedError(format!("{:?}", err)))?;
let stream_session_id = match response {
ProcessedMessage::ErrorResponse(msg) => return Err(XtbClientBuilderError::LoginFailed { user_id: user_id.to_string(), extra_info: format!("{:?}", msg) }),
ProcessedMessage::Response(response) => response.stream_session_id.unwrap(),
};
let stream_connection = BasicXtbStreamConnection::new(stream_api_url, stream_session_id).await.map_err(|err| XtbClientBuilderError::CannotMakeStreamConnection(err))?;
Ok(XtbClient::new(connection, stream_connection, self.ping_period.unwrap_or(DEFAULT_PING_INTERVAL_S)))
}
fn make_url(source: Option<String>) -> Result<Url, XtbClientBuilderError> {
let source_str = source.ok_or_else(|| XtbClientBuilderError::RequiredFieldMissing("api_url".to_owned()))?;
Url::from_str(&source_str).map_err(|err| XtbClientBuilderError::InvalidUrl(source_str, err))
}
}
#[derive(Debug, Error)]
pub enum XtbClientBuilderError {
#[error("Required configuration field is missing: {0}")]
RequiredFieldMissing(String),
#[error("Url is invalid or malformed: {0} ({1})")]
InvalidUrl(String, url::ParseError),
#[error("Cannot connect to server")]
CannotMakeConnection(BasicXtbConnectionError),
#[error("Cannot connect to stream server")]
CannotMakeStreamConnection(BasicXtbStreamConnectionError),
#[error("Login failed for user: {user_id} ({extra_info:?})")]
LoginFailed { user_id: String, extra_info: String },
#[error("Something gets horribly wrong: {0}")]
UnexpectedError(String),
}
#[async_trait]
pub trait RequestResponseApi {
type Error;
async fn get_all_symbols(&mut self, request: GetAllSymbolsRequest) -> Result<GetAllSymbolsResponse, Self::Error>;
async fn get_calendar(&mut self, request: GetCalendarRequest) -> Result<GetCalendarResponse, Self::Error>;
async fn get_chart_last_request(&mut self, request: GetChartLastRequestRequest) -> Result<GetChartLastRequestResponse, Self::Error>;
async fn get_chart_range_request(&mut self, request: GetChartRangeRequestRequest) -> Result<GetChartRangeRequestResponse, Self::Error>;
async fn get_commission_def(&mut self, request: GetCommissionDefRequest) -> Result<GetCommissionDefResponse, Self::Error>;
async fn get_current_user_data(&mut self, request: GetCurrentUserDataRequest) -> Result<GetCurrentUserDataResponse, Self::Error>;
async fn get_ibs_history(&mut self, request: GetIbsHistoryRequest) -> Result<GetIbsHistoryResponse, Self::Error>;
async fn get_margin_level(&mut self, request: GetMarginLevelRequest) -> Result<GetMarginLevelResponse, Self::Error>;
async fn get_margin_trade(&mut self, request: GetMarginTradeRequest) -> Result<GetMarginTradeResponse, Self::Error>;
async fn get_news(&mut self, request: GetNewsRequest) -> Result<GetNewsResponse, Self::Error>;
async fn get_profit_calculation(&mut self, request: GetProfitCalculationRequest) -> Result<GetProfitCalculationResponse, Self::Error>;
async fn get_server_time(&mut self, request: GetServerTimeRequest) -> Result<GetServerTimeResponse, Self::Error>;
async fn get_step_rules(&mut self, request: GetStepRulesRequest) -> Result<GetStepRulesResponse, Self::Error>;
async fn get_symbol(&mut self, request: GetSymbolRequest) -> Result<GetSymbolResponse, Self::Error>;
async fn get_tick_prices(&mut self, request: GetTickPricesRequest) -> Result<GetTickPricesResponse, Self::Error>;
async fn get_trade_records(&mut self, request: GetTradeRecordsRequest) -> Result<GetTradeRecordsResponse, Self::Error>;
async fn get_trades(&mut self, request: GetTradesRequest) -> Result<GetTradesResponse, Self::Error>;
async fn get_trades_history(&mut self, request: GetTradesHistoryRequest) -> Result<GetTradesHistoryResponse, Self::Error>;
async fn get_trading_hours(&mut self, request: GetTradingHoursRequest) -> Result<GetTradingHoursResponse, Self::Error>;
async fn get_version(&mut self, request: GetVersionRequest) -> Result<GetVersionResponse, Self::Error>;
async fn trade_transaction(&mut self, request: TradeTransactionRequest) -> Result<TradeTransactionResponse, Self::Error>;
async fn trade_transaction_status(&mut self, request: TradeTransactionStatusRequest) -> Result<TradeTransactionStatusResponse, Self::Error>;
}
#[async_trait]
pub trait StreamApi {
type Error;
type Stream<T: Send + Sync + for<'de> Deserialize<'de>>;
async fn subscribe_balance(&mut self, arguments: StreamGetBalanceSubscribe) -> Result<Self::Stream<StreamGetBalanceData>, Self::Error>;
async fn subscribe_candles(&mut self, arguments: StreamGetCandlesSubscribe) -> Result<Self::Stream<StreamGetCandlesData>, Self::Error>;
async fn subscribe_keep_alive(&mut self, arguments: StreamGetKeepAliveSubscribe) -> Result<Self::Stream<StreamGetKeepAliveData>, Self::Error>;
async fn subscribe_news(&mut self, arguments: StreamGetNewsSubscribe) -> Result<Self::Stream<StreamGetNewsData>, Self::Error>;
async fn subscribe_profits(&mut self, arguments: StreamGetProfitSubscribe) -> Result<Self::Stream<StreamGetProfitData>, Self::Error>;
async fn subscribe_tick_prices(&mut self, arguments: StreamGetTickPricesSubscribe) -> Result<Self::Stream<StreamGetTickPricesData>, Self::Error>;
async fn subscribe_trades(&mut self, arguments: StreamGetTradesSubscribe) -> Result<Self::Stream<StreamGetTradesData>, Self::Error>;
async fn subscribe_trade_status(&mut self, arguments: StreamGetTradeStatusSubscribe) -> Result<Self::Stream<StreamGetTradeStatusData>, Self::Error>;
}
pub struct XtbClient {
connection: Arc<Mutex<BasicXtbConnection>>,
stream_manager: StreamManager,
ping_join_handle: JoinHandle<()>,
stream_ping_join_handle: JoinHandle<()>,
}
impl XtbClient {
pub fn builder() -> XtbClientBuilder {
XtbClientBuilder::default()
}
pub fn new(connection: BasicXtbConnection, stream_connection: BasicXtbStreamConnection, ping_period: u64) -> Self {
let connection = Arc::new(Mutex::new(connection));
let ping_join_handle = spawn_ping(connection.clone(), ping_period);
let stream_manager = StreamManager::new(stream_connection);
let stream_ping_join_handle = spawn_stream_ping(stream_manager.clone(), ping_period);
Self {
connection,
stream_manager,
ping_join_handle,
stream_ping_join_handle,
}
}
async fn send_and_wait_or_default<REQ, RESP>(&mut self, command: &str, request: REQ) -> Result<RESP, XtbClientError>
where
REQ: Serialize,
RESP: for<'de> Deserialize<'de> + Default {
self.send_and_wait(command, request).await.map(|val| val.unwrap_or_default())
}
async fn send_and_wait<REQ, RESP>(&mut self, command: &str, request: REQ) -> Result<Option<RESP>, XtbClientError>
where
REQ: Serialize,
RESP: for<'de> Deserialize<'de>
{
let promise = self.send(command, request).await?;
let response = promise.await.map_err(|err| {
error!("Unexpected error: {:?}", err);
XtbClientError::UnexpectedError
})?;
match response {
ProcessedMessage::Response(response) => {
match response.return_data {
Some(data) => from_value(data).map_err(|err| XtbClientError::DeserializationFailed(err)).map(|v| Some(v)),
None => Ok(None)
}
}
ProcessedMessage::ErrorResponse(err) => Err(XtbClientError::CommandFailed(err)),
}
}
async fn send<A>(&mut self, command: &str, request: A) -> Result<ResponsePromise, XtbClientError>
where
A: Serialize
{
let mut conn = self.connection.lock().await;
let payload = Self::convert_data_to_value(request)?;
conn.send_command(command, Some(payload)).await.map_err(|err| {
match err {
BasicXtbConnectionError::SerializationError(err) => XtbClientError::SerializationFailed(err),
BasicXtbConnectionError::CannotSendRequest(err) => XtbClientError::CannotSendCommand(err),
_ => XtbClientError::UnexpectedError,
}
})
}
fn convert_data_to_value<T: Serialize>(data: T) -> Result<Value, XtbClientError> {
to_value(data).map_err(|err| XtbClientError::SerializationFailed(err))
}
async fn send_simple_stream_command<T, SA, UA>(
&mut self,
subscribe_command: &str,
subscribe_arguments: SA,
unsubscribe_command: &str,
unsubscribe_arguments: UA,
data_command: &str,
) -> Result<DataStream<T>, XtbClientError>
where
T: for<'de> Deserialize<'de> + Send + Sync,
SA: Serialize,
UA: Serialize,
{
let unsubscribe_arguments = Self::convert_data_to_value(unsubscribe_arguments)?;
let filter = DataMessageFilter::Command(data_command.to_owned());
let subscribe_arguments = Self::convert_data_to_value(subscribe_arguments)?;
self.stream_manager.subscribe(subscribe_command, Some(subscribe_arguments), unsubscribe_command, Some(unsubscribe_arguments), data_command, filter).await
}
async fn send_symbol_scoped_stream_command<T, SA, UA>(
&mut self,
subscribe_command: &str,
subscribe_arguments: SA,
unsubscribe_command: &str,
unsubscribe_arguments: UA,
data_command: &str,
symbol: &str,
) -> Result<DataStream<T>, XtbClientError>
where
T: for<'de> Deserialize<'de> + Send + Sync,
SA: Serialize,
UA: Serialize,
{
let unsubscribe_arguments = Self::convert_data_to_value(unsubscribe_arguments)?;
let subscribe_arguments = Self::convert_data_to_value(subscribe_arguments)?;
let subscription_key = format!("{}.{}", data_command, symbol);
let filter = DataMessageFilter::All(vec![
DataMessageFilter::Command(data_command.to_owned()),
DataMessageFilter::FieldValue { name: "symbol".to_owned(), value: Value::String(symbol.to_owned()) },
]);
self.stream_manager.subscribe(subscribe_command, Some(subscribe_arguments), unsubscribe_command, Some(unsubscribe_arguments), &subscription_key, filter).await
}
}
impl Drop for XtbClient {
fn drop(&mut self) {
self.ping_join_handle.abort();
self.stream_ping_join_handle.abort();
}
}
#[async_trait]
impl RequestResponseApi for XtbClient {
type Error = XtbClientError;
async fn get_all_symbols(&mut self, request: GetAllSymbolsRequest) -> Result<GetAllSymbolsResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_GET_ALL_SYMBOLS, request).await
}
async fn get_calendar(&mut self, request: GetCalendarRequest) -> Result<GetCalendarResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_GET_CALENDAR, request).await
}
async fn get_chart_last_request(&mut self, request: GetChartLastRequestRequest) -> Result<GetChartLastRequestResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_GET_CHART_LAST_REQUEST, request).await
}
async fn get_chart_range_request(&mut self, request: GetChartRangeRequestRequest) -> Result<GetChartRangeRequestResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_GET_CHART_RANGE_REQUEST, request).await
}
async fn get_commission_def(&mut self, request: GetCommissionDefRequest) -> Result<GetCommissionDefResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_GET_COMMISSION_DEF, request).await
}
async fn get_current_user_data(&mut self, request: GetCurrentUserDataRequest) -> Result<GetCurrentUserDataResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_GET_CURRENT_USER_DATA, request).await
}
async fn get_ibs_history(&mut self, request: GetIbsHistoryRequest) -> Result<GetIbsHistoryResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_GET_IBS_HISTORY, request).await
}
async fn get_margin_level(&mut self, request: GetMarginLevelRequest) -> Result<GetMarginLevelResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_GET_MARGIN_LEVEL, request).await
}
async fn get_margin_trade(&mut self, request: GetMarginTradeRequest) -> Result<GetMarginTradeResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_GET_MARGIN_TRADE, request).await
}
async fn get_news(&mut self, request: GetNewsRequest) -> Result<GetNewsResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_GET_NEWS, request).await
}
async fn get_profit_calculation(&mut self, request: GetProfitCalculationRequest) -> Result<GetProfitCalculationResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_GET_PROFIT_CALCULATION, request).await
}
async fn get_server_time(&mut self, request: GetServerTimeRequest) -> Result<GetServerTimeResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_GET_SERVER_TIME, request).await
}
async fn get_step_rules(&mut self, request: GetStepRulesRequest) -> Result<GetStepRulesResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_GET_STEP_RULES, request).await
}
async fn get_symbol(&mut self, request: GetSymbolRequest) -> Result<GetSymbolResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_GET_SYMBOL, request).await
}
async fn get_tick_prices(&mut self, request: GetTickPricesRequest) -> Result<GetTickPricesResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_GET_TICK_PRICES, request).await
}
async fn get_trade_records(&mut self, request: GetTradeRecordsRequest) -> Result<GetTradeRecordsResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_GET_TRADE_RECORDS, request).await
}
async fn get_trades(&mut self, request: GetTradesRequest) -> Result<GetTradesResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_GET_TRADES, request).await
}
async fn get_trades_history(&mut self, request: GetTradesHistoryRequest) -> Result<GetTradesHistoryResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_GET_TRADES_HISTORY, request).await
}
async fn get_trading_hours(&mut self, request: GetTradingHoursRequest) -> Result<GetTradingHoursResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_GET_TRADING_HOURS, request).await
}
async fn get_version(&mut self, request: GetVersionRequest) -> Result<GetVersionResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_GET_VERSION, request).await
}
async fn trade_transaction(&mut self, request: TradeTransactionRequest) -> Result<TradeTransactionResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_TRADE_TRANSACTION, request).await
}
async fn trade_transaction_status(&mut self, request: TradeTransactionStatusRequest) -> Result<TradeTransactionStatusResponse, Self::Error> {
self.send_and_wait_or_default(COMMAND_TRADE_TRANSACTION_STATUS, request).await
}
}
#[async_trait]
impl StreamApi for XtbClient {
type Error = XtbClientError;
type Stream<T: Send + Sync + for<'de> Deserialize<'de>> = DataStream<T>;
async fn subscribe_balance(&mut self, arguments: StreamGetBalanceSubscribe) -> Result<Self::Stream<StreamGetBalanceData>, Self::Error> {
let stop_arguments = Self::convert_data_to_value(StreamGetBalanceUnsubscribe::default())?;
self.send_simple_stream_command(STREAM_BALANCE_SUBSCRIBE, arguments, STREAM_BALANCE_UNSUBSCRIBE, stop_arguments, STREAM_BALANCE).await
}
async fn subscribe_candles(&mut self, arguments: StreamGetCandlesSubscribe) -> Result<Self::Stream<StreamGetCandlesData>, Self::Error> {
let stop_arguments = Self::convert_data_to_value(StreamGetCandlesUnsubscribe::default().with_symbol(&arguments.symbol))?;
let symbol = arguments.symbol.clone();
self.send_symbol_scoped_stream_command(STREAM_CANDLES_SUBSCRIBE, arguments, STREAM_CANDLES_UNSUBSCRIBE, stop_arguments, STREAM_CANDLES, &symbol).await
}
async fn subscribe_keep_alive(&mut self, arguments: StreamGetKeepAliveSubscribe) -> Result<Self::Stream<StreamGetKeepAliveData>, Self::Error> {
let stop_arguments = Self::convert_data_to_value(StreamGetKeepAliveUnsubscribe::default())?;
self.send_simple_stream_command(STREAM_KEEP_ALIVE_SUBSCRIBE, arguments, STREAM_KEEP_ALIVE_UNSUBSCRIBE, stop_arguments, STREAM_KEEP_ALIVE).await
}
async fn subscribe_news(&mut self, arguments: StreamGetNewsSubscribe) -> Result<Self::Stream<StreamGetNewsData>, Self::Error> {
let stop_arguments = Self::convert_data_to_value(StreamGetNewsUnsubscribe::default())?;
self.send_simple_stream_command(STREAM_NEWS_SUBSCRIBE, arguments, STREAM_NEWS_UNSUBSCRIBE, stop_arguments, STREAM_NEWS).await
}
async fn subscribe_profits(&mut self, arguments: StreamGetProfitSubscribe) -> Result<Self::Stream<StreamGetProfitData>, Self::Error> {
let stop_arguments = Self::convert_data_to_value(StreamGetProfitUnsubscribe::default())?;
self.send_simple_stream_command(STREAM_PROFITS_SUBSCRIBE, arguments, STREAM_PROFITS_UNSUBSCRIBE, stop_arguments, STREAM_PROFITS).await
}
async fn subscribe_tick_prices(&mut self, arguments: StreamGetTickPricesSubscribe) -> Result<Self::Stream<StreamGetTickPricesData>, Self::Error> {
let stop_arguments = Self::convert_data_to_value(StreamGetTickPricesUnsubscribe::default().with_symbol(&arguments.symbol))?;
let symbol = arguments.symbol.clone();
self.send_symbol_scoped_stream_command(STREAM_TICK_PRICES_SUBSCRIBE, arguments, STREAM_TICK_PRICES_UNSUBSCRIBE, stop_arguments, STREAM_TICK_PRICES, &symbol).await
}
async fn subscribe_trades(&mut self, arguments: StreamGetTradesSubscribe) -> Result<Self::Stream<StreamGetTradesData>, Self::Error> {
let stop_arguments = Self::convert_data_to_value(StreamGetTradesUnsubscribe::default())?;
self.send_simple_stream_command(STREAM_TRADES_SUBSCRIBE, arguments, STREAM_TRADES_UNSUBSCRIBE, stop_arguments, STREAM_TRADES).await
}
async fn subscribe_trade_status(&mut self, arguments: StreamGetTradeStatusSubscribe) -> Result<Self::Stream<StreamGetTradeStatusData>, Self::Error> {
let stop_arguments = Self::convert_data_to_value(StreamGetTradeStatusUnsubscribe::default())?;
self.send_simple_stream_command(STREAM_TRADE_STATUS_SUBSCRIBE, arguments, STREAM_TRADE_STATUS_UNSUBSCRIBE, stop_arguments, STREAM_TRADE_STATUS).await
}
}
#[derive(Debug, Error)]
pub enum XtbClientError {
#[error("Cannot serialize arguments")]
SerializationFailed(serde_json::Error),
#[error("Cannot send command to server")]
CannotSendCommand(tokio_tungstenite::tungstenite::Error),
#[error("Cannot send stream command")]
CannotSendStreamCommand(BasicXtbStreamConnectionError),
#[error("Unexpected error.")]
UnexpectedError,
#[error("Cannot deserialize data")]
DeserializationFailed(serde_json::Error),
#[error("Command failed and an error response was returned")]
CommandFailed(ErrorResponse),
}
#[derive(Debug)]
struct StreamManagerState {
connection: BasicXtbStreamConnection,
subscriptions: HashMap<String, usize>,
}
impl StreamManagerState {
pub fn new(connection: BasicXtbStreamConnection) -> Self {
Self {
connection,
subscriptions: HashMap::new(),
}
}
}
#[derive(Clone, Debug)]
struct StreamManager {
state: Arc<Mutex<StreamManagerState>>,
}
impl StreamManager {
pub fn new(connection: BasicXtbStreamConnection) -> Self {
let state = Arc::new(Mutex::new(StreamManagerState::new(connection)));
Self {
state
}
}
pub async fn subscribe<T: for<'de> Deserialize<'de> + Send + Sync>(
&mut self,
subscribe_command: &str,
subscribe_arguments: Option<Value>,
unsubscribe_command: &str,
unsubscribe_arguments: Option<Value>,
subscription_key: &str,
filter: DataMessageFilter,
) -> Result<DataStream<T>, XtbClientError> {
let mut state = self.state.lock().await;
let stream = state.connection.make_message_stream(filter).await;
state.connection.subscribe(subscribe_command, subscribe_arguments).await.map_err(|err| XtbClientError::CannotSendStreamCommand(err))?;
*state.subscriptions.entry(subscription_key.to_owned()).or_default() += 1;
Ok(DataStream::new(stream, self.clone(), subscription_key.to_owned(), unsubscribe_command.to_owned(), unsubscribe_arguments))
}
pub async fn unsubscribe(&mut self, subscription_key: &str, command: &str, arguments: Option<Value>) -> Result<(), XtbClientError> {
let mut state = self.state.lock().await;
let mut entry = state.subscriptions.entry(subscription_key.to_owned()).or_default();
if *entry > 0 {
*entry -= 1;
}
if *entry == 0 {
state.connection.unsubscribe(command, arguments).await.map_err(|err| XtbClientError::CannotSendStreamCommand(err))?;
}
Ok(())
}
}
pub struct DataStream<T>
where
T: for<'de> Deserialize<'de> + Send + Sync
{
message_stream: BasicMessageStream,
stream_manager: StreamManager,
subscription_key: String,
unsubscribe_command: String,
unsubscribe_arguments: Option<Value>,
type_: PhantomData<T>,
}
impl<T> DataStream<T>
where
T: for<'de> Deserialize<'de> + Send + Sync
{
fn new(message_stream: BasicMessageStream, stream_manager: StreamManager, subscription_key: String, unsubscribe_command: String, unsubscribe_arguments: Option<Value>) -> Self {
Self {
message_stream,
stream_manager,
subscription_key,
unsubscribe_command,
unsubscribe_arguments,
type_: PhantomData::<T>,
}
}
pub async fn next(&mut self) -> Result<Option<T>, DataStreamError> {
let message = self.message_stream.next().await;
match message {
Some(msg) => Self::process_message(msg).map(|r| Some(r)),
None => Ok(None),
}
}
fn process_message(msg: StreamDataMessage) -> Result<T, DataStreamError> {
from_value(msg.data).map_err(|err| DataStreamError::CannotDeserializeValue(err))
}
}
impl<T> Drop for DataStream<T>
where
T: for<'de> Deserialize<'de> + Send + Sync
{
fn drop(&mut self) {
let mut manager = self.stream_manager.clone();
let unsubscribe_command = self.unsubscribe_command.clone();
let unsubscribe_arguments = self.unsubscribe_arguments.take();
let subscription_key = self.subscription_key.clone();
spawn(async move {
let result = manager.unsubscribe(&subscription_key, &unsubscribe_command, unsubscribe_arguments.clone()).await;
match result {
Err(err) => error!("Cannot unsubscribe command '{unsubscribe_command}' ({unsubscribe_arguments:?}). The subscription key was: '{subscription_key}'. The error was: {err:?}"),
_ => (),
};
});
}
}
#[derive(Debug, Error)]
pub enum DataStreamError {
#[error("Cannot deserialize value: {0}")]
CannotDeserializeValue(serde_json::Error)
}
fn spawn_ping(conn: Arc<Mutex<BasicXtbConnection>>, ping_secs: u64) -> JoinHandle<()> {
let ping_value = to_value(PingRequest::default()).expect("Cannot serialize ping message");
spawn(async move {
let mut idx = 1u64;
loop {
let response_promise = {
let mut conn = conn.lock().await;
debug!("Sending ping #{} to connection", idx);
match conn.send_command(COMMAND_PING, Some(ping_value.clone())).await {
Ok(resp) => Some(resp),
Err(err) => {
error!("Cannot send ping #{}: {:?}", idx, err);
None
}
}
};
if let Some(response_promise) = response_promise {
match response_promise.await {
Ok(_) => (),
Err(err) => error!("Cannot await the ping response #{}: {:?}", idx, err)
}
}
idx += 1;
sleep(Duration::from_secs(ping_secs)).await;
}
})
}
fn spawn_stream_ping(stream_manager: StreamManager, ping_secs: u64) -> JoinHandle<()> {
let ping_value = to_value(StreamPingSubscribe::default()).expect("Cannot serialize the stream ping message");
spawn(async move {
let mut idx = 1u64;
loop {
{
debug!("Sending ping #{} to stream connection", idx);
let mut inner_state = stream_manager.state.lock().await;
match inner_state.connection.subscribe(STREAM_PING, Some(ping_value.clone())).await {
Ok(_) => (),
Err(err) => error!("Cannot send ping #{}: {:?}", idx, err)
}
}
idx += 1;
sleep(Duration::from_secs(ping_secs)).await;
}
})
}