use crate::ohlcv::OhlcvBar;
use crate::signals::pipeline::SignalPipeline;
use crate::signals::SignalValue;
use chrono::{DateTime, Utc};
use tokio::sync::mpsc;
#[derive(Debug, Clone)]
pub struct SignalUpdate {
pub signal_name: String,
pub value: SignalValue,
pub timestamp: DateTime<Utc>,
}
impl SignalUpdate {
pub fn new(signal_name: impl Into<String>, value: SignalValue, timestamp: DateTime<Utc>) -> Self {
Self {
signal_name: signal_name.into(),
value,
timestamp,
}
}
pub fn is_ready(&self) -> bool {
self.value.is_scalar()
}
}
#[derive(Debug, Clone)]
pub struct StreamingConfig {
pub tick_channel_capacity: usize,
pub output_channel_capacity: usize,
}
impl Default for StreamingConfig {
fn default() -> Self {
Self {
tick_channel_capacity: 1_024,
output_channel_capacity: 4_096,
}
}
}
pub struct StreamingSignalPipeline {
pipeline: SignalPipeline,
config: StreamingConfig,
}
impl StreamingSignalPipeline {
pub fn new(pipeline: SignalPipeline) -> Self {
Self {
pipeline,
config: StreamingConfig::default(),
}
}
pub fn with_config(pipeline: SignalPipeline, config: StreamingConfig) -> Self {
Self { pipeline, config }
}
pub fn spawn(
self,
) -> (mpsc::Sender<OhlcvBar>, mpsc::Receiver<SignalUpdate>) {
let (tick_tx, tick_rx) = mpsc::channel::<OhlcvBar>(self.config.tick_channel_capacity);
let (update_tx, update_rx) =
mpsc::channel::<SignalUpdate>(self.config.output_channel_capacity);
tokio::spawn(run_pipeline(self.pipeline, tick_rx, update_tx));
(tick_tx, update_rx)
}
}
pub fn spawn_signal_stream(
pipeline: SignalPipeline,
tick_rx: mpsc::Receiver<OhlcvBar>,
) -> mpsc::Receiver<SignalUpdate> {
let capacity = StreamingConfig::default().output_channel_capacity;
let (update_tx, update_rx) = mpsc::channel::<SignalUpdate>(capacity);
tokio::spawn(run_pipeline(pipeline, tick_rx, update_tx));
update_rx
}
async fn run_pipeline(
mut pipeline: SignalPipeline,
mut tick_rx: mpsc::Receiver<OhlcvBar>,
update_tx: mpsc::Sender<SignalUpdate>,
) {
let mut known_names: Vec<String> = Vec::new();
while let Some(bar) = tick_rx.recv().await {
let ts = Utc::now();
let map = pipeline.update(&bar);
if known_names.is_empty() {
known_names = map.names().iter().map(|s| (*s).to_owned()).collect();
}
for name in &known_names {
let value = map
.get(name)
.cloned()
.unwrap_or(SignalValue::Unavailable);
let update = SignalUpdate::new(name.clone(), value, ts);
if update_tx.send(update).await.is_err() {
return;
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::ohlcv::OhlcvBar;
use crate::signals::indicators::Sma;
use crate::signals::pipeline::SignalPipeline;
use crate::types::{NanoTimestamp, Price, Quantity, Symbol};
use rust_decimal_macros::dec;
fn make_bar(close: rust_decimal::Decimal, ts: i64) -> OhlcvBar {
let sym = Symbol::new("X").unwrap();
let p = Price::new(close).unwrap();
OhlcvBar {
symbol: sym,
open: p,
high: p,
low: p,
close: p,
volume: Quantity::new(dec!(1)).unwrap(),
ts_open: NanoTimestamp::new(ts),
ts_close: NanoTimestamp::new(ts + 1),
tick_count: 1,
}
}
#[tokio::test]
async fn test_streaming_pipeline_receives_updates() {
let sma = Sma::new("sma3", 3).unwrap();
let pipeline = SignalPipeline::new().add(sma);
let (tick_tx, mut update_rx) = StreamingSignalPipeline::new(pipeline).spawn();
for i in 1u32..=5 {
let bar = make_bar(rust_decimal::Decimal::from(i) * dec!(10), i64::from(i));
tick_tx.send(bar).await.unwrap();
}
drop(tick_tx);
let mut updates: Vec<SignalUpdate> = Vec::new();
while let Some(u) = update_rx.recv().await {
updates.push(u);
}
assert_eq!(updates.len(), 5);
assert!(updates[0].value.is_unavailable());
assert!(updates[1].value.is_unavailable());
assert!(updates[2].value.is_scalar());
}
#[tokio::test]
async fn test_spawn_signal_stream_convenience() {
let sma = Sma::new("sma2", 2).unwrap();
let pipeline = SignalPipeline::new().add(sma);
let (tick_tx, tick_rx) = mpsc::channel::<OhlcvBar>(16);
let mut update_rx = spawn_signal_stream(pipeline, tick_rx);
for i in 1u32..=3 {
let bar = make_bar(rust_decimal::Decimal::from(i) * dec!(5), i64::from(i));
tick_tx.send(bar).await.unwrap();
}
drop(tick_tx);
let mut count = 0usize;
while update_rx.recv().await.is_some() {
count += 1;
}
assert_eq!(count, 3);
}
#[tokio::test]
async fn test_pipeline_closes_when_sender_dropped() {
let sma = Sma::new("sma5", 5).unwrap();
let pipeline = SignalPipeline::new().add(sma);
let (tick_tx, mut update_rx) = StreamingSignalPipeline::new(pipeline).spawn();
drop(tick_tx);
let result = update_rx.recv().await;
assert!(result.is_none());
}
}