diff --git a/pulse-sdk/src/general.rs b/pulse-sdk/src/general.rs index 56e6253..f7491aa 100644 --- a/pulse-sdk/src/general.rs +++ b/pulse-sdk/src/general.rs @@ -30,7 +30,7 @@ pub enum Mode { Manual, } -#[derive(Debug, Clone)] +#[derive(Debug, Clone, Copy)] pub enum Allocation { Fixed(Decimal), Percent(Decimal), diff --git a/src/engine/engine/risk.rs b/src/engine/engine/risk.rs index 923a9a0..e49f295 100644 --- a/src/engine/engine/risk.rs +++ b/src/engine/engine/risk.rs @@ -11,7 +11,7 @@ use hypersdk::{ use pulse_sdk::general::Signal; use tokio::sync::Mutex; -use crate::engine::Engine; +use crate::{engine::Engine, store::config::RiskConfig}; #[derive(Debug, Clone)] pub struct OrderIds { @@ -74,6 +74,8 @@ impl RiskEngine { .sum(); } + self.state.lock().await.pnl = 0.into(); + Ok(()) } @@ -140,20 +142,25 @@ impl RiskEngine { Ok(()) } - pub async fn validate_signal(&self, _signal: &mut Signal) -> bool { - true + pub async fn validate_signal(&self, _signal: &mut Signal) -> tokio::io::Result { + let state = self.state.lock().await; + + let under_max_losses = state.pnl + < self + .get_risk_config() + .await? + .max_daily_loss + .get(state.starting_equity); + + Ok(under_max_losses) } pub async fn create_order(&self) -> tokio::io::Result { let state = self.state.lock().await; Ok(OrderIds::new( - self.get_engine() - .config - .lock() - .await - .get_ref()? - .risk + self.get_risk_config() + .await? .risk_per_trade .get(state.starting_equity), )) @@ -163,6 +170,10 @@ impl RiskEngine { self.orders.lock().await.push(order); } + pub async fn get_risk_config(&self) -> tokio::io::Result { + Ok(self.get_engine().config.lock().await.get_ref()?.risk) + } + pub fn get_engine(&self) -> Arc { self.engine .upgrade() diff --git a/src/engine/engine/strategy.rs b/src/engine/engine/strategy.rs index 56621fc..88e02e1 100644 --- a/src/engine/engine/strategy.rs +++ b/src/engine/engine/strategy.rs @@ -151,7 +151,7 @@ impl StrategyEngine { } Some(StrategyMessage::Signal(mut signal)) => { - if !engine.risk_engine.validate_signal(&mut signal).await { + if !engine.risk_engine.validate_signal(&mut signal).await? { engine.signals.lock().await.push(Err(signal)); engine diff --git a/src/engine/store/config.rs b/src/engine/store/config.rs index 7e321fe..a39bd9a 100644 --- a/src/engine/store/config.rs +++ b/src/engine/store/config.rs @@ -2,7 +2,7 @@ use std::collections::HashMap; use pulse_sdk::general::Allocation; -#[derive(Debug, serde::Serialize, serde::Deserialize)] +#[derive(Debug, Clone, Copy, serde::Serialize, serde::Deserialize)] pub struct RiskConfig { pub risk_per_trade: Allocation, pub max_open_positions: u32,