use crate::czsc_signals::CzscSignals;
use crate::sig_parse::SignalConfig;
use czsc_core::analyze::CZSC;
use czsc_core::objects::bar::RawBar;
use czsc_core::objects::position::{LiteBar, Position};
use czsc_core::objects::state::TraderState;
use czsc_signals::types::TraderSignalFn;
use czsc_utils::bar_generator::BarGenerator;
use polars::prelude::*;
use serde_json::Value;
use std::collections::HashMap;
use std::fs::File;
use std::path::Path;
use std::time::Instant;
#[derive(Debug, Clone, Copy, Default)]
pub struct UpdateProfile {
pub signals_update_ns: u128,
pub trader_signals_ns: u128,
pub position_update_ns: u128,
pub pos_event_match_ns: u128,
pub pos_fsm_ns: u128,
pub pos_risk_ns: u128,
pub pos_holds_ns: u128,
}
#[derive(Clone)]
struct CompiledTraderSignalOp {
func: TraderSignalFn,
params: HashMap<String, Value>,
}
pub struct CzscTrader {
pub name: String,
pub signals: CzscSignals,
pub positions: Vec<Position>,
compiled_trader_ops: Vec<CompiledTraderSignalOp>,
compiled_cfg_ptr: usize,
compiled_cfg_len: usize,
}
impl CzscTrader {
pub fn new(symbol: String, bg: BarGenerator, positions: Vec<Position>) -> Self {
Self {
name: "CzscTrader".to_string(),
signals: CzscSignals::new(symbol, bg),
positions,
compiled_trader_ops: Vec::new(),
compiled_cfg_ptr: 0,
compiled_cfg_len: 0,
}
}
fn ensure_compiled_trader_ops(&mut self, signals_config: &[SignalConfig]) {
let ptr = signals_config.as_ptr() as usize;
let len = signals_config.len();
if self.compiled_cfg_ptr == ptr && self.compiled_cfg_len == len {
return;
}
self.compiled_trader_ops.clear();
self.compiled_trader_ops.reserve(signals_config.len());
for config in signals_config {
if config.freq.is_none()
&& let Some(meta) =
czsc_signals::registry::TRADER_SIGNAL_REGISTRY.get(config.name.as_str())
{
self.compiled_trader_ops.push(CompiledTraderSignalOp {
func: meta.func,
params: config.params.clone(),
});
}
}
self.compiled_cfg_ptr = ptr;
self.compiled_cfg_len = len;
}
pub fn update(&mut self, bar: &RawBar, signals_config: &[SignalConfig]) {
let _ = self.update_profiled(bar, signals_config);
}
pub fn update_profiled(
&mut self,
bar: &RawBar,
signals_config: &[SignalConfig],
) -> UpdateProfile {
self.ensure_compiled_trader_ops(signals_config);
let t_signals = Instant::now();
self.signals.update_signals(bar, signals_config);
let signals_update_ns = t_signals.elapsed().as_nanos();
let t_trader_sig = Instant::now();
let mut trader_sigs = Vec::new();
for op in &self.compiled_trader_ops {
let sigs = (op.func)(self, &op.params);
trader_sigs.extend(sigs);
}
let trader_signals_ns = t_trader_sig.elapsed().as_nanos();
let t_pos = Instant::now();
for sig in trader_sigs {
let (k, v) = (sig.key(), sig.value());
self.signals.s.insert(k.clone(), v.clone());
self.signals.signal_map.insert(k, v);
self.signals.sigs.insert(sig);
}
let lite_bar = LiteBar {
id: bar.id,
dt: bar.dt.into(),
price: bar.close,
};
let mut pos_event_match_ns = 0u128;
let mut pos_fsm_ns = 0u128;
let mut pos_risk_ns = 0u128;
let mut pos_holds_ns = 0u128;
for pos in &mut self.positions {
let p =
pos.update_profiled_with_signal_map(lite_bar, None, Some(&self.signals.signal_map));
pos_event_match_ns += p.event_match_ns;
pos_fsm_ns += p.fsm_ns;
pos_risk_ns += p.risk_ns;
pos_holds_ns += p.holds_ns;
}
let position_update_ns = t_pos.elapsed().as_nanos();
UpdateProfile {
signals_update_ns,
trader_signals_ns,
position_update_ns,
pos_event_match_ns,
pos_fsm_ns,
pos_risk_ns,
pos_holds_ns,
}
}
pub fn dump_results(&self, out_dir: &str) -> anyhow::Result<()> {
let path = Path::new(out_dir);
if !path.exists() {
std::fs::create_dir_all(path)?;
}
let mut all_pairs = Vec::new();
let mut all_holds = Vec::new();
for pos in &self.positions {
if let Ok(df) = pos.pairs()
&& df.height() > 0
{
all_pairs.push(df.lazy());
}
if let Ok(df) = pos.holds()
&& df.height() > 0
{
all_holds.push(df.lazy());
}
}
if !all_pairs.is_empty() {
let mut combined_pairs = concat(all_pairs, UnionArgs::default())?.collect()?;
let mut file = File::create(path.join("pairs.parquet"))?;
ParquetWriter::new(&mut file).finish(&mut combined_pairs)?;
}
if !all_holds.is_empty() {
let mut combined_holds = concat(all_holds, UnionArgs::default())?.collect()?;
let mut file = File::create(path.join("holds.parquet"))?;
ParquetWriter::new(&mut file).finish(&mut combined_holds)?;
}
Ok(())
}
}
impl TraderState for CzscTrader {
#[inline]
fn get_position(&self, name: &str) -> Option<&Position> {
self.positions.iter().find(|p| p.name == name)
}
#[inline]
fn get_czsc(&self, freq: &str) -> Option<&CZSC> {
self.signals.kas.get(freq)
}
#[inline]
fn latest_price(&self) -> Option<f64> {
self.signals
.s
.get("close")
.and_then(|x| x.parse::<f64>().ok())
}
}