fin_primitives/
async_signals.rs1use crate::ohlcv::OhlcvBar;
20use crate::signals::pipeline::SignalPipeline;
21use crate::signals::SignalValue;
22use chrono::{DateTime, Utc};
23use tokio::sync::mpsc;
24
25#[derive(Debug, Clone)]
29pub struct SignalUpdate {
30 pub signal_name: String,
32 pub value: SignalValue,
34 pub timestamp: DateTime<Utc>,
36}
37
38impl SignalUpdate {
39 pub fn new(signal_name: impl Into<String>, value: SignalValue, timestamp: DateTime<Utc>) -> Self {
41 Self {
42 signal_name: signal_name.into(),
43 value,
44 timestamp,
45 }
46 }
47
48 pub fn is_ready(&self) -> bool {
50 self.value.is_scalar()
51 }
52}
53
54#[derive(Debug, Clone)]
58pub struct StreamingConfig {
59 pub tick_channel_capacity: usize,
61 pub output_channel_capacity: usize,
63}
64
65impl Default for StreamingConfig {
66 fn default() -> Self {
67 Self {
68 tick_channel_capacity: 1_024,
69 output_channel_capacity: 4_096,
70 }
71 }
72}
73
74pub struct StreamingSignalPipeline {
80 pipeline: SignalPipeline,
81 config: StreamingConfig,
82}
83
84impl StreamingSignalPipeline {
85 pub fn new(pipeline: SignalPipeline) -> Self {
87 Self {
88 pipeline,
89 config: StreamingConfig::default(),
90 }
91 }
92
93 pub fn with_config(pipeline: SignalPipeline, config: StreamingConfig) -> Self {
95 Self { pipeline, config }
96 }
97
98 pub fn spawn(
106 self,
107 ) -> (mpsc::Sender<OhlcvBar>, mpsc::Receiver<SignalUpdate>) {
108 let (tick_tx, tick_rx) = mpsc::channel::<OhlcvBar>(self.config.tick_channel_capacity);
109 let (update_tx, update_rx) =
110 mpsc::channel::<SignalUpdate>(self.config.output_channel_capacity);
111
112 tokio::spawn(run_pipeline(self.pipeline, tick_rx, update_tx));
113
114 (tick_tx, update_rx)
115 }
116}
117
118pub fn spawn_signal_stream(
127 pipeline: SignalPipeline,
128 tick_rx: mpsc::Receiver<OhlcvBar>,
129) -> mpsc::Receiver<SignalUpdate> {
130 let capacity = StreamingConfig::default().output_channel_capacity;
131 let (update_tx, update_rx) = mpsc::channel::<SignalUpdate>(capacity);
132 tokio::spawn(run_pipeline(pipeline, tick_rx, update_tx));
133 update_rx
134}
135
136async fn run_pipeline(
143 mut pipeline: SignalPipeline,
144 mut tick_rx: mpsc::Receiver<OhlcvBar>,
145 update_tx: mpsc::Sender<SignalUpdate>,
146) {
147 let mut known_names: Vec<String> = Vec::new();
152
153 while let Some(bar) = tick_rx.recv().await {
154 let ts = Utc::now();
155
156 let map = pipeline.update(&bar);
158
159 if known_names.is_empty() {
161 known_names = map.names().iter().map(|s| (*s).to_owned()).collect();
162 }
163
164 for name in &known_names {
166 let value = map
167 .get(name)
168 .cloned()
169 .unwrap_or(SignalValue::Unavailable);
170
171 let update = SignalUpdate::new(name.clone(), value, ts);
172
173 if update_tx.send(update).await.is_err() {
175 return;
176 }
177 }
178 }
179}
180
181#[cfg(test)]
184mod tests {
185 use super::*;
186 use crate::ohlcv::OhlcvBar;
187 use crate::signals::indicators::Sma;
188 use crate::signals::pipeline::SignalPipeline;
189 use crate::types::{NanoTimestamp, Price, Quantity, Symbol};
190 use rust_decimal_macros::dec;
191
192 fn make_bar(close: rust_decimal::Decimal, ts: i64) -> OhlcvBar {
193 let sym = Symbol::new("X").unwrap();
194 let p = Price::new(close).unwrap();
195 OhlcvBar {
196 symbol: sym,
197 open: p,
198 high: p,
199 low: p,
200 close: p,
201 volume: Quantity::new(dec!(1)).unwrap(),
202 ts_open: NanoTimestamp::new(ts),
203 ts_close: NanoTimestamp::new(ts + 1),
204 tick_count: 1,
205 }
206 }
207
208 #[tokio::test]
209 async fn test_streaming_pipeline_receives_updates() {
210 let sma = Sma::new("sma3", 3).unwrap();
211 let pipeline = SignalPipeline::new().add(sma);
212
213 let (tick_tx, mut update_rx) = StreamingSignalPipeline::new(pipeline).spawn();
214
215 for i in 1u32..=5 {
217 let bar = make_bar(rust_decimal::Decimal::from(i) * dec!(10), i64::from(i));
218 tick_tx.send(bar).await.unwrap();
219 }
220
221 drop(tick_tx);
223
224 let mut updates: Vec<SignalUpdate> = Vec::new();
225 while let Some(u) = update_rx.recv().await {
226 updates.push(u);
227 }
228
229 assert_eq!(updates.len(), 5);
231 assert!(updates[0].value.is_unavailable());
233 assert!(updates[1].value.is_unavailable());
234 assert!(updates[2].value.is_scalar());
236 }
237
238 #[tokio::test]
239 async fn test_spawn_signal_stream_convenience() {
240 let sma = Sma::new("sma2", 2).unwrap();
241 let pipeline = SignalPipeline::new().add(sma);
242
243 let (tick_tx, tick_rx) = mpsc::channel::<OhlcvBar>(16);
244 let mut update_rx = spawn_signal_stream(pipeline, tick_rx);
245
246 for i in 1u32..=3 {
247 let bar = make_bar(rust_decimal::Decimal::from(i) * dec!(5), i64::from(i));
248 tick_tx.send(bar).await.unwrap();
249 }
250 drop(tick_tx);
251
252 let mut count = 0usize;
253 while update_rx.recv().await.is_some() {
254 count += 1;
255 }
256 assert_eq!(count, 3);
257 }
258
259 #[tokio::test]
260 async fn test_pipeline_closes_when_sender_dropped() {
261 let sma = Sma::new("sma5", 5).unwrap();
262 let pipeline = SignalPipeline::new().add(sma);
263 let (tick_tx, mut update_rx) = StreamingSignalPipeline::new(pipeline).spawn();
264
265 drop(tick_tx);
267
268 let result = update_rx.recv().await;
270 assert!(result.is_none());
271 }
272}