diff --git a/Cargo.lock b/Cargo.lock index 1aad095..5b95774 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -2264,6 +2264,18 @@ dependencies = [ "zeroize", ] +[[package]] +name = "embedded-io" +version = "0.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ef1a6892d9eef45c8fa6b9e0086428a2cca8491aca8f787c534a3d6d0bcb3ced" + +[[package]] +name = "embedded-io" +version = "0.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "edd0f118536f44f5ccd48bcb8b111bdc3de888b58c74639dfb034a357d0f206d" + [[package]] name = "encoding_rs" version = "0.8.35" @@ -3740,6 +3752,8 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6764c3b5dd454e283a30e6dfe78e9b31096d9e32036b5d1eaac7a6119ccb9a24" dependencies = [ "cobs", + "embedded-io 0.4.0", + "embedded-io 0.6.1", "heapless", "serde", ] @@ -3958,7 +3972,9 @@ name = "pulse-wire" version = "0.1.0-alpha.0" dependencies = [ "hypersdk", + "postcard", "serde", + "tokio", ] [[package]] diff --git a/Cargo.toml b/Cargo.toml index 7f050ea..33d3eee 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -29,7 +29,7 @@ toml = "1.1.3" members = ["pulse-ui", "pulse-wire"] [workspace.dependencies] -postcard = "1.1.3" +postcard = { version = "1.1.3", features = ["alloc"]} pulse-ui = { path = "pulse-ui", version = "0.1.0-alpha.0" } pulse-wire = { path = "pulse-wire", version = "0.1.0-alpha.0" } hypersdk = "0.2.14" diff --git a/pulse-wire/Cargo.toml b/pulse-wire/Cargo.toml index 54d4acc..bd40410 100644 --- a/pulse-wire/Cargo.toml +++ b/pulse-wire/Cargo.toml @@ -6,3 +6,5 @@ edition = "2024" [dependencies] serde = { workspace = true } hypersdk = { workspace = true } +postcard = { workspace = true } +tokio = { workspace = true } diff --git a/pulse-wire/src/general.rs b/pulse-wire/src/general.rs index 97f9dfc..a2fffad 100644 --- a/pulse-wire/src/general.rs +++ b/pulse-wire/src/general.rs @@ -1,13 +1,13 @@ use crate::units::{Direction, Symbol, USD}; -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub enum MarketTrend { Bullish, Bearish, Neutral, } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub enum LogKind { Info, Warn, @@ -15,7 +15,7 @@ pub enum LogKind { Debug, } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct Signal { pub symbol: String, pub kind: Direction, @@ -26,14 +26,14 @@ pub struct Signal { pub stop_loss: USD, } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct EventLog { pub kind: LogKind, pub name: String, pub message: String, } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct Position { pub symbol: Symbol, pub size: f64, diff --git a/pulse-wire/src/lib.rs b/pulse-wire/src/lib.rs index cca42fc..dd2211e 100644 --- a/pulse-wire/src/lib.rs +++ b/pulse-wire/src/lib.rs @@ -19,3 +19,7 @@ pub mod prelude { pub fn server_path() -> PathBuf { PathBuf::from("/tmp/pulse-engine.sock") } + +pub fn map_postcard_err(res: postcard::Result) -> tokio::io::Result { + res.map_err(|e| tokio::io::Error::new(std::io::ErrorKind::Other, e)) +} diff --git a/pulse-wire/src/plugin.rs b/pulse-wire/src/plugin.rs index 875413c..821a226 100644 --- a/pulse-wire/src/plugin.rs +++ b/pulse-wire/src/plugin.rs @@ -1,34 +1,30 @@ use crate::{ - PulseWire, general::{EventLog, Signal}, terminal::MarketItem, units::Direction, }; use hypersdk::hypercore::{Candle, CandleInterval, Subscription}; -use pulse_macros::pwp; -#[pwp] -#[derive(serde::Deserialize, serde::Serialize)] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct StrategyManifest { - name: String, - description: String, - author: String, - version: String, + pub name: String, + pub description: String, + pub author: String, + pub version: String, } -#[pwp] -#[derive(serde::Deserialize, serde::Serialize)] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct RiskManifest { - name: String, - description: String, - author: String, - version: String, + pub name: String, + pub description: String, + pub author: String, + pub version: String, - max_loss: u8, - cooldown: CandleInterval, + pub max_loss: u8, + pub cooldown: CandleInterval, } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub enum StrategyMessage { Log(EventLog), @@ -49,7 +45,7 @@ pub enum StrategyMessage { Signal(StrategySignal), } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub enum RiskMessage { Log(EventLog), @@ -60,7 +56,7 @@ pub enum RiskMessage { Reject { reason: String }, } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub enum StrategyEngineMessage { Initialize, @@ -84,7 +80,7 @@ pub enum StrategyEngineMessage { }, } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub enum RiskEngineMessage { Initialize, @@ -95,7 +91,7 @@ pub enum RiskEngineMessage { Signal(StrategySignal), } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct StrategySignal { pub symbol: String, pub side: Direction, diff --git a/pulse-wire/src/terminal.rs b/pulse-wire/src/terminal.rs index 072001d..a07e9f1 100644 --- a/pulse-wire/src/terminal.rs +++ b/pulse-wire/src/terminal.rs @@ -1,13 +1,11 @@ use crate::{ - PulseWire, general::{EventLog, MarketTrend, Position, Signal}, plugin::{RiskManifest, StrategyManifest}, units::{Symbol, USD, Volatility}, }; use hypersdk::hypercore::CandleInterval; -use pulse_macros::pwp; -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub enum TerminalServerMessage { // WatchList WatchListUpdated(Vec), @@ -32,12 +30,12 @@ pub enum TerminalServerMessage { AddLog(EventLog), } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub enum TerminalClientMessage { ExecuteCommand(String), } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct MarketItem { pub symbol: Symbol, pub price: USD, @@ -45,7 +43,7 @@ pub struct MarketItem { pub volume_24h: USD, } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct MarketOverview { pub trend: MarketTrend, pub volatility: Volatility, @@ -53,32 +51,32 @@ pub struct MarketOverview { pub alerts: Vec, } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub enum AlertLevel { High, Medium, Low, } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct Alert { - level: AlertLevel, - message: String, + pub level: AlertLevel, + pub message: String, } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct Balance { pub asset: String, pub amount: f64, pub value: f64, } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub enum InspectTarget { None, Some(Vec), } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub enum InspectItem { String(String), Symbol(String), @@ -86,35 +84,35 @@ pub enum InspectItem { F64(f64), } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct Status { - feed: Mode, - exchange: String, - dex: String, - latency: u16, + pub feed: Mode, + pub exchange: String, + pub dex: String, + pub latency: u16, } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub enum Mode { Auto, Manual, } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub enum ItemState { Running, Stopped, Error, } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct Strategy { - strategy: StrategyManifest, - risk: RiskManifest, + pub strategy: StrategyManifest, + pub risk: RiskManifest, - mode: Mode, - state: ItemState, - cooldown: CandleInterval, + pub mode: Mode, + pub state: ItemState, + pub cooldown: CandleInterval, } impl std::fmt::Display for AlertLevel { diff --git a/pulse-wire/src/units.rs b/pulse-wire/src/units.rs index c7db0e6..e2758e9 100644 --- a/pulse-wire/src/units.rs +++ b/pulse-wire/src/units.rs @@ -1,20 +1,16 @@ -use pulse_macros::pwp; - -use crate::PulseWire; - -#[derive(Debug, Clone)] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub struct Symbol(pub String); -#[derive(Debug, Clone, Copy)] +#[derive(Debug, Clone, Copy, serde::Serialize, serde::Deserialize)] pub struct USD(pub f64); -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub enum Direction { Buy, Sell, } -#[pwp] +#[derive(Debug, Clone, serde::Serialize, serde::Deserialize)] pub enum Volatility { Low, Medium, @@ -37,26 +33,6 @@ impl std::fmt::Display for USD { } } -impl PulseWire for Symbol { - fn from_com(com: &mut Vec) -> Self { - Self(String::from_com(com)) - } - - fn to_com(&self) -> Vec { - self.0.to_com() - } -} - -impl PulseWire for USD { - fn from_com(com: &mut Vec) -> Self { - Self(f64::from_com(com)) - } - - fn to_com(&self) -> Vec { - self.0.to_com() - } -} - pub fn format_f64(value: f64) -> String { let abs = value.abs(); diff --git a/src/engine/engine/plugin.rs b/src/engine/engine/plugin.rs index 223e59b..4109b13 100644 --- a/src/engine/engine/plugin.rs +++ b/src/engine/engine/plugin.rs @@ -166,11 +166,11 @@ impl StrategyEngine { None => {} Some(RiskMessage::GetWatchList) => { - let mut v = vec![1]; - - v.extend(engine.config.lock().await.watchlist.to_com()); - - self.risk.send_raw(&v).await?; + self.risk + .send(&&RiskEngineMessage::WatchList( + engine.watch_list.lock().await.clone(), + )) + .await?; } Some(RiskMessage::Log(mut log)) => { diff --git a/src/engine/store/plugin.rs b/src/engine/store/plugin.rs index 761fa1f..9763dc9 100644 --- a/src/engine/store/plugin.rs +++ b/src/engine/store/plugin.rs @@ -1,16 +1,14 @@ use std::marker::PhantomData; -use pulse_wire::PulseWire; - -use serde::Deserialize; +use pulse_wire::map_postcard_err; +use serde::{Deserialize, Serialize}; use tokio::{ io::{AsyncReadExt, AsyncWriteExt}, process::{Child, ChildStdout}, sync::Mutex, }; - #[derive(Debug)] -pub struct Plugin Deserialize<'de>> { +pub struct Plugin Deserialize<'de>, M: for<'de> Deserialize<'de>> { pub manifest: Mutex, pub stdout: Mutex, pub process: Mutex, @@ -18,7 +16,7 @@ pub struct Plugin Deserialize<'de>> { pub _p: (PhantomData, PhantomData), } -impl Deserialize<'de>> Plugin { +impl Deserialize<'de>, M: for<'de> Deserialize<'de>> Plugin { pub fn new(mut child: Child, manifest: M) -> Self { Self { stdout: Mutex::new(child.stdout.take().expect("Failed to obtain child stdout")), @@ -45,11 +43,12 @@ impl Deserialize<'de>> Plugin { stdout.read_exact(&mut buffer).await?; - Ok(Some(R::from_com(&mut buffer))) + Ok(Some(map_postcard_err(postcard::from_bytes(&buffer))?)) } pub async fn send(&self, msg: &S) -> tokio::io::Result<()> { - self.send_raw(&msg.to_com()).await + self.send_raw(&map_postcard_err(postcard::to_allocvec(msg))?) + .await } pub async fn send_raw(&self, msg: &[u8]) -> tokio::io::Result<()> { diff --git a/src/engine/terminal.rs b/src/engine/terminal.rs index 6623b6a..df601a7 100644 --- a/src/engine/terminal.rs +++ b/src/engine/terminal.rs @@ -1,5 +1,5 @@ use crate::engine::Engine; -use pulse_wire::prelude::*; +use pulse_wire::{map_postcard_err, prelude::*}; use std::{ collections::HashMap, sync::{Arc, Weak}, @@ -104,7 +104,7 @@ impl TerminalServer { reader.read_exact(&mut buffer).await?; - match TerminalClientMessage::from_com(&mut buffer) { + match map_postcard_err(postcard::from_bytes(&buffer))? { TerminalClientMessage::ExecuteCommand(command) => { let command = command.as_str(); @@ -126,7 +126,7 @@ impl TerminalServer { self: &Arc, message: pulse_wire::terminal::TerminalServerMessage, ) -> tokio::io::Result<()> { - let msg = message.to_com(); + let msg = map_postcard_err(postcard::to_allocvec(&message))?; let mut clients = self.clients.lock().await; @@ -158,7 +158,7 @@ impl TerminalServer { format!("Client({id}) does not exist"), ) })?, - &message.to_com(), + &map_postcard_err(postcard::to_allocvec(&message))?, ) .await } diff --git a/src/terminal/terminal.rs b/src/terminal/terminal.rs index 468a20d..5519562 100644 --- a/src/terminal/terminal.rs +++ b/src/terminal/terminal.rs @@ -1,4 +1,5 @@ -use pulse_wire::prelude::*; +use pulse_ui::state::State; +use pulse_wire::{map_postcard_err, prelude::*}; use tokio::{ io::{AsyncReadExt, AsyncWriteExt}, @@ -29,7 +30,7 @@ impl TerminalClient { &mut self, message: pulse_wire::terminal::TerminalClientMessage, ) -> tokio::io::Result<()> { - let msg = message.to_com(); + let msg = map_postcard_err(postcard::to_allocvec(&message))?; self.writer.write(&msg.len().to_le_bytes()).await?; self.writer.write(&msg).await?; self.writer.flush().await?; @@ -42,7 +43,7 @@ impl TerminalClient { std::mem::swap(&mut self.reader, &mut reader); - let mut reader = reader.expect("Reader failed to swap"); + let reader = reader.expect("Reader failed to swap"); let watch_list = app.watch_list.clone(); let active_positions = app.active_positions.clone(); @@ -52,61 +53,82 @@ impl TerminalClient { let status = app.status.clone(); let inspect = app.inspect.clone(); - tokio::spawn(async move { - loop { - let mut len_buf = [0u8; size_of::()]; - reader - .read_exact(&mut len_buf) - .await - .expect("Failed to get header length"); - - let len = usize::from_le_bytes(len_buf); - - let mut buffer = vec![0u8; len]; - - reader - .read_exact(&mut buffer) - .await - .expect("Failed to read socket"); - - match TerminalServerMessage::from_com(&mut buffer) { - TerminalServerMessage::WatchListUpdated(v) => { - *watch_list.lock().await = v; - } - - TerminalServerMessage::PositionsUpdated(v) => { - *active_positions.lock().await = v; - } - - TerminalServerMessage::StrategyUpdated(v) => { - *market_overview.lock().await = Some(v); - } - - TerminalServerMessage::SignalsUpdated(v) => { - *signals.lock().await = v; - } - - TerminalServerMessage::Inspect(v) => { - *inspect.lock().await = v; - } - - TerminalServerMessage::StatusUpdated(v) => { - *status.lock().await = Some(v); - } - - TerminalServerMessage::SetLogs(v) => { - *logs.lock().await = v; - } - - TerminalServerMessage::AddLog(v) => { - logs.lock().await.push(v); - } - } - } - }); + tokio::spawn(Self::run_client( + reader, + watch_list, + active_positions, + logs, + signals, + market_overview, + status, + inspect, + )); app.sock = Some(self); app } + + pub async fn run_client( + mut reader: OwnedReadHalf, + + watch_list: State>, + active_positions: State>, + logs: State>, + signals: State>, + market_overview: State>, + status: State>, + inspect: State, + ) -> tokio::io::Result<()> { + let mut len_buf = [0u8; size_of::()]; + reader + .read_exact(&mut len_buf) + .await + .expect("Failed to get header length"); + + let len = usize::from_le_bytes(len_buf); + + let mut buffer = vec![0u8; len]; + + reader + .read_exact(&mut buffer) + .await + .expect("Failed to read socket"); + + match map_postcard_err(postcard::from_bytes(&buffer))? { + TerminalServerMessage::WatchListUpdated(v) => { + *watch_list.lock().await = v; + } + + TerminalServerMessage::PositionsUpdated(v) => { + *active_positions.lock().await = v; + } + + TerminalServerMessage::StrategyUpdated(v) => { + *market_overview.lock().await = Some(v); + } + + TerminalServerMessage::SignalsUpdated(v) => { + *signals.lock().await = v; + } + + TerminalServerMessage::Inspect(v) => { + *inspect.lock().await = v; + } + + TerminalServerMessage::StatusUpdated(v) => { + *status.lock().await = Some(v); + } + + TerminalServerMessage::SetLogs(v) => { + *logs.lock().await = v; + } + + TerminalServerMessage::AddLog(v) => { + logs.lock().await.push(v); + } + } + + Ok(()) + } }