use std::sync::Arc;
use std::time::Duration;
use serde::{Deserialize, Serialize};
use super::client::StreamResult;
use super::handle::{RECONNECT_BACKOFF, SourceStream, stream_builder, stream_handle};
use super::polygon::{AssetClass, PolygonTradeSource};
use super::source::ReconnectConfig;
const CHANNEL_CAPACITY: usize = 4096;
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
#[serde(rename_all = "camelCase")]
#[non_exhaustive]
pub struct TradeTick {
pub symbol: String,
pub price: f64,
pub size: f64,
pub exchange: Option<i32>,
pub conditions: Vec<i32>,
pub trade_id: Option<String>,
pub time: i64,
}
impl TradeTick {
pub fn notional(&self) -> f64 {
self.price * self.size
}
}
stream_handle! {
TradeStream(TradeTick);
add: add_symbols = "Add symbols to the subscription.",
remove: remove_symbols = "Remove symbols from the subscription.",
}
impl TradeStream {
pub async fn subscribe<S, I>(symbols: I) -> StreamResult<Self>
where
S: Into<String>,
I: IntoIterator<Item = S>,
{
TradeStreamBuilder::new().symbols(symbols).build().await
}
}
pub struct TradeStreamBuilder {
symbols: Vec<String>,
asset_class: AssetClass,
retry_delay: Duration,
max_reconnect_attempts: Option<u32>,
}
impl TradeStreamBuilder {
pub fn new() -> Self {
Self {
symbols: Vec::new(),
asset_class: AssetClass::Stocks,
retry_delay: RECONNECT_BACKOFF,
max_reconnect_attempts: None,
}
}
pub fn asset_class(mut self, class: AssetClass) -> Self {
self.asset_class = class;
self
}
pub async fn build(self) -> StreamResult<TradeStream> {
let source = Arc::new(PolygonTradeSource::new(self.asset_class)?);
let reconnect =
ReconnectConfig::new(self.retry_delay).max_attempts(self.max_reconnect_attempts);
Ok(TradeStream {
inner: SourceStream::start(source, self.symbols, reconnect, CHANNEL_CAPACITY),
})
}
}
stream_builder!(TradeStreamBuilder, symbols = "Add symbols to subscribe to.");
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn classes_without_trade_prints_are_rejected() {
for class in [AssetClass::Forex, AssetClass::Indices] {
assert!(
TradeStreamBuilder::new()
.symbols(["X"])
.asset_class(class)
.build()
.await
.is_err()
);
}
}
#[test]
fn notional_multiplies_price_by_size() {
let tick = TradeTick {
price: 10.0,
size: 25.0,
..Default::default()
};
assert!((tick.notional() - 250.0).abs() < 1e-9);
}
}