Files
pulse-trader/src/engine/engine/plugin.rs
T

241 lines
7.7 KiB
Rust

use hypersdk::hypercore::{self, CandleInterval, Subscription, WebSocket};
use pulse_wire::prelude::*;
use std::{
collections::HashSet,
path::PathBuf,
process::Stdio,
sync::{Arc, Weak},
time::{SystemTime, UNIX_EPOCH},
};
use tokio::{
fs,
process::{Child, Command},
sync::Mutex,
};
use crate::{
engine::Engine,
store::{plugin::Plugin, pulse_plugin},
};
pub struct StrategyEngine {
pub strategy: Arc<Plugin<StrategyEngineMessage, StrategyMessage, StrategyManifest>>,
pub risk: Arc<Plugin<RiskEngineMessage, RiskMessage, RiskManifest>>,
pub engine: Weak<Engine>,
pub ws: WebSocket,
pub subscriptions: Mutex<HashSet<Subscription>>,
}
impl StrategyEngine {
pub async fn new(strategy_id: &str, risk_id: &str) -> tokio::io::Result<Self> {
let strategy = pulse_plugin(strategy_id)?;
let risk = pulse_plugin(risk_id)?;
let (strategy, strategy_manifest) = get_manifest_plugin_pair(
&strategy,
&strategy.join("strategy.bash"),
&fs::read(strategy.join("strategy.toml")).await?,
)?;
let (risk, risk_manifest) = get_manifest_plugin_pair(
&risk,
&risk.join("risk.bash"),
&fs::read(risk.join("risk.toml")).await?,
)?;
Ok(Self {
strategy: Arc::new(Plugin::new(strategy, strategy_manifest)),
risk: Arc::new(Plugin::new(risk, risk_manifest)),
engine: Weak::new(),
ws: hypercore::mainnet_ws(),
subscriptions: Mutex::new(HashSet::new()),
})
}
pub fn initialize(mut self, engine: Weak<Engine>) -> Arc<Self> {
self.engine = engine;
Arc::new(self)
}
pub async fn run_strategy(&self) -> anyhow::Result<()> {
let engine = self
.engine
.upgrade()
.expect("Failed to upgrade engine (StrategyEngine)");
self.strategy
.send(&StrategyEngineMessage::Initialize)
.await?;
loop {
match self.strategy.recv().await? {
None => {}
Some(StrategyMessage::GetWatchList) => {
let mut v = vec![1];
v.extend(engine.config.lock().await.watchlist.to_com());
self.strategy.send_raw(&v).await?;
}
Some(StrategyMessage::Log(mut log)) => {
log.name.insert_str(0, "strategy::");
engine.terminal_server.log_raw(log).await?;
}
Some(StrategyMessage::Signal(signal)) => {
self.risk.send(&RiskEngineMessage::Signal(signal)).await?;
}
Some(StrategyMessage::Subscribe(subscription)) => {
self.ws.subscribe(subscription.clone());
self.subscriptions.lock().await.insert(subscription);
}
Some(StrategyMessage::Unsubscribe(subscription)) => {
self.subscriptions.lock().await.remove(&subscription);
self.ws.unsubscribe(subscription);
}
Some(StrategyMessage::UnsubscribeAll) => {
for sub in self.subscriptions.lock().await.drain() {
self.ws.unsubscribe(sub);
}
}
Some(StrategyMessage::RequestCandlestick {
symbol,
interval,
count,
}) => {
let client = hypercore::mainnet();
let now = SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap()
.as_millis() as u64;
let interval_ms = match interval {
CandleInterval::OneMinute => 60_000,
CandleInterval::ThreeMinutes => 3 * 60_000,
CandleInterval::FiveMinutes => 5 * 60_000,
CandleInterval::FifteenMinutes => 15 * 60_000,
CandleInterval::ThirtyMinutes => 30 * 60_000,
CandleInterval::OneHour => 60 * 60_000,
CandleInterval::TwoHours => 2 * 60 * 60_000,
CandleInterval::FourHours => 4 * 60 * 60_000,
CandleInterval::EightHours => 8 * 60 * 60_000,
CandleInterval::TwelveHours => 12 * 60 * 60_000,
CandleInterval::OneDay => 24 * 60 * 60_000,
CandleInterval::ThreeDays => 3 * 24 * 60 * 60_000,
CandleInterval::OneWeek => 7 * 24 * 60 * 60_000,
CandleInterval::OneMonth => 30 * 24 * 60 * 60_000,
};
let start_time = now.saturating_sub(interval_ms * count as u64);
client
.candle_snapshot(symbol, interval, start_time, now)
.await?;
}
}
}
}
pub async fn run_risk(&self) -> tokio::io::Result<()> {
let engine = self
.engine
.upgrade()
.expect("Failed to upgrade engine (StrategyEngine)");
self.risk.send(&RiskEngineMessage::Initialize).await?;
loop {
match self.risk.recv().await? {
None => {}
Some(RiskMessage::GetWatchList) => {
let mut v = vec![1];
v.extend(engine.config.lock().await.watchlist.to_com());
self.risk.send_raw(&v).await?;
}
Some(RiskMessage::Log(mut log)) => {
log.name.insert_str(0, "risk::");
engine.terminal_server.log_raw(log).await?;
}
Some(RiskMessage::Approve(signal)) => {}
Some(RiskMessage::Reject { reason }) => {}
}
}
}
pub async fn spawn(self: &Arc<Self>) {
let engine = self.clone();
tokio::spawn(async move { engine.run_strategy().await });
let engine = self.clone();
tokio::spawn(async move { engine.run_risk().await });
}
pub async fn reload_strategy(self: &Arc<Self>, id: &str) -> tokio::io::Result<()> {
let plugin = pulse_plugin(id)?;
let (child, manifest) = get_manifest_plugin_pair(
&plugin,
&plugin.join("strategy.bash"),
&fs::read(plugin.join("strategy.toml")).await?,
)?;
self.strategy.reload(child, manifest).await?;
let engine = self.clone();
tokio::spawn(async move { engine.run_strategy().await });
Ok(())
}
pub async fn reload_risk(self: &Arc<Self>, id: &str) -> tokio::io::Result<()> {
let plugin = pulse_plugin(id)?;
let (child, manifest) = get_manifest_plugin_pair(
&plugin,
&plugin.join("risk.bash"),
&fs::read(plugin.join("risk.toml")).await?,
)?;
self.risk.reload(child, manifest).await?;
let engine = self.clone();
tokio::spawn(async move { engine.run_risk().await });
Ok(())
}
}
pub fn get_manifest_plugin_pair<'de, M: serde::Deserialize<'de>>(
plugin_dir: &PathBuf,
plugin_path: &PathBuf,
manifest: &'de [u8],
) -> tokio::io::Result<(Child, M)> {
Ok((
Command::new("bash")
.arg(plugin_path)
.current_dir(plugin_dir)
.stdin(Stdio::piped())
.stdout(Stdio::piped())
.stderr(Stdio::inherit())
.spawn()?,
toml::from_slice(manifest)
.map_err(|v| tokio::io::Error::new(std::io::ErrorKind::InvalidInput, v.to_string()))?,
))
}