use anyhow::Error;
use chrono::{DateTime, Utc};
use futures::{FutureExt as _, select};
use linregress::{FormulaRegressionBuilder, RegressionDataBuilder};
use rust_decimal::{Decimal, prelude::FromPrimitive};
use std::{
future::Future,
pin::Pin,
sync::{Arc, atomic::Ordering},
};
use tokio::sync::{
Notify,
watch::{Receiver, Sender, channel},
};
use crate::{
clock::Clock,
core::{Core, CoreError, CoreResult, market::Market},
generics::order::{Crypto, OrderDraft, OrderSide, OrderType},
provider::{
DataQuery, Interval, ProviderTrait, Source,
kline::{BaseKline, Kline, KlineSpecifier},
},
};
use super::{AtomicState, ErrorInStrat, State, StratStat, Strategy, StrategyFactory};
pub struct TrendTunnelFactory<
C: Clock,
M: Market<_Clock = C>,
S: Source<_Data = Kline>,
P: ProviderTrait<_Source = S>,
> {
quote: Crypto,
amount: Decimal,
asset: Crypto,
core: Option<Core<C, M>>,
provider: Option<P>,
run_date: DateTime<Utc>,
interval: Interval,
sensibility_percentage: u16,
}
impl<C: Clock, M: Market<_Clock = C>, S: Source<_Data = Kline>, P: ProviderTrait<_Source = S>>
TrendTunnelFactory<C, M, S, P>
{
pub fn new(
quote: Crypto,
amount: Decimal,
asset: Crypto,
run_date: DateTime<Utc>,
interval: Interval,
sensibility_percentage: u16,
) -> Self {
Self {
quote,
amount,
asset,
core: None,
provider: None,
run_date,
interval,
sensibility_percentage,
}
}
pub fn with_provider(self, provider: P) -> Self {
Self {
provider: Some(provider),
..self
}
}
}
impl<C: Clock, M: Market<_Clock = C>, S: Source<_Data = Kline>, P: ProviderTrait<_Source = S>>
StrategyFactory for TrendTunnelFactory<C, M, S, P>
{
type _Strategy = TrendTunnel<C, M, S, P>;
fn build(mut self) -> TrendTunnel<C, M, S, P> {
let core = self.core.take().expect("Core not set");
let provider = self.provider.take().expect("Provider not set");
let future_state = Arc::new(AtomicState::new(START_STATE));
let notify = Arc::new(Notify::new());
let (tx, rx) = channel(Err(ErrorInStrat::new(
"Not initialized".to_string(),
core.clock().now(),
)));
TrendTunnel {
core,
provider,
future_state,
state: START_STATE,
notify,
asset: self.asset,
quote: self.quote,
run_date: self.run_date,
interval: self.interval,
invested_status: false,
sensibility_percentage: self.sensibility_percentage,
tx,
rx,
avail_asset: Decimal::ZERO,
avail_quote: self.amount,
}
}
fn with_core(self, core: Core<C, M>) -> Self {
Self {
core: Some(core),
..self
}
}
}
pub struct TrendTunnel<
C: Clock,
M: Market<_Clock = C>,
S: Source<_Data = Kline>,
P: ProviderTrait<_Source = S>,
> {
core: Core<C, M>,
provider: P,
future_state: Arc<AtomicState>,
state: State,
notify: Arc<Notify>,
tx: Sender<Result<StratStat, ErrorInStrat>>,
rx: Receiver<Result<StratStat, ErrorInStrat>>,
avail_asset: Decimal,
avail_quote: Decimal,
run_date: DateTime<Utc>,
invested_status: bool,
quote: Crypto,
asset: Crypto,
interval: Interval,
sensibility_percentage: u16,
}
const START_STATE: State = State::Paused;
const SIZE_KLINES_VEC: i32 = 100;
impl<C: Clock, M: Market<_Clock = C>, S: Source<_Data = Kline>, P: ProviderTrait<_Source = S>>
TrendTunnel<C, M, S, P>
{
async fn get_stats(&self) -> Result<StratStat, Error> {
let asset = self.core.get_balance(self.asset).await?;
let quote = self.core.get_balance(self.quote).await?;
Ok(StratStat {
asset,
quote,
time: self.core.clock().now(),
})
}
async fn send_stats(&self) {
let stat = self
.get_stats()
.await
.map_err(|err| ErrorInStrat::new(err.to_string(), self.core.clock().now()));
if let Err(err) = self.tx.send(stat) {
tracing::error!("Failed to send stat: {}", err);
}
}
async fn await_date_or_state_change(&mut self, optional_date: Option<DateTime<Utc>>) {
let actual_future_state = self.future_state.load(Ordering::Relaxed);
if actual_future_state != self.state {
self.state = actual_future_state;
return;
}
if let Some(date) = optional_date {
select! {
_ = self.core.clock().sleep_until(date).fuse() => {},
_ = self.notify.notified().fuse() => {
self.core.clock().synchronize();
}
}
} else {
self.notify.notified().await;
self.core.clock().synchronize();
}
self.state = self.future_state.load(Ordering::Relaxed);
}
async fn latest_basekline(&self, now: DateTime<Utc>) -> CoreResult<BaseKline> {
let mut interval = Interval::OneHour;
loop {
tracing::trace!("Getting smallest kline with interval: {}", interval);
let maybe_kline = self
.provider
.provide_or_empty(&DataQuery::new(
KlineSpecifier::new(self.asset, self.quote),
interval,
now - interval.time_delta(),
now,
))
.await?
.pop();
if let Some(kline) = maybe_kline
&& let Some(base) = kline.content.clone()
{
tracing::trace!(
"Latest kline is : {:#?} with interval {}",
kline.basics.begin,
interval
);
return Ok(base);
}
interval = interval
.next_interval_to_zoom_out()
.ok_or(CoreError::comput_error("No interval found"))?;
}
}
async fn decide(&mut self) -> CoreResult<()> {
let now = self.core.clock().now();
tracing::trace!(
"Getting klines from {} to {}",
now - (self
.interval
.time_delta()
.checked_mul(SIZE_KLINES_VEC)
.ok_or(CoreError::comput_error("todo"))?),
now
);
let klines = self
.provider
.provide(&DataQuery::new(
KlineSpecifier::new(self.asset, self.quote),
self.interval,
now - (self
.interval
.time_delta()
.checked_mul(SIZE_KLINES_VEC)
.ok_or(CoreError::comput_error("todo"))?),
now,
))
.await?;
let len = klines.len();
tracing::trace!("Got {} klines", len);
let (highs, lows) = klines
.iter()
.fold((vec![], vec![]), |(mut highs, mut lows), kline| {
if let Some(kline) = &kline.content {
highs.push(kline.high.try_into().unwrap());
lows.push(kline.low.try_into().unwrap());
} else {
tracing::error!("Missing kline {}", kline.basics.begin);
}
(highs, lows)
});
let x: Vec<f64> = (1..=len).map(|n| f64::from_usize(n).unwrap()).collect();
let data_highs = vec![("x", x.clone()), ("y", highs)];
let data_lows = vec![("x", x), ("y", lows)];
let reg_data_highs = RegressionDataBuilder::new()
.build_from(data_highs)
.map_err(|_| CoreError::comput_error("todo"))?;
let reg_data_lows = RegressionDataBuilder::new()
.build_from(data_lows)
.map_err(|_| CoreError::comput_error("todo"))?;
let formula = "y ~ x";
let model_highs = FormulaRegressionBuilder::new()
.data(®_data_highs)
.formula(formula)
.fit()
.unwrap();
let model_lows = FormulaRegressionBuilder::new()
.data(®_data_lows)
.formula(formula)
.fit()
.unwrap();
let next_high_vec = model_highs
.predict(vec![("x", vec![f64::from_usize(len).unwrap() + 1.])])
.unwrap();
let next_low_vec = model_lows
.predict(vec![("x", vec![f64::from_usize(len).unwrap() + 1.])])
.unwrap();
let next_high = next_high_vec[0];
let next_low = next_low_vec[0];
let level_high =
next_high - (next_high - next_low) * self.sensibility_percentage as f64 / 100.0;
let level_low =
next_low + (next_high - next_low) * self.sensibility_percentage as f64 / 100.0;
tracing::trace!("Getting kline to evaluate now ({}) cost", now);
let now_kline = self.latest_basekline(now).await?;
let double_price: f64 = (now_kline.open + now_kline.close).try_into().unwrap();
let avg_price: f64 = double_price / 2.0;
match self.invested_status {
true => {
if avg_price > level_high {
tracing::debug!("About to buy");
self.core
.send_order(OrderDraft {
asset: self.asset,
quote: self.quote,
side: OrderSide::Sell,
order_type: OrderType::Market,
qty_asset: Some(self.avail_asset),
qty_quote: None,
price: None,
})
.await?;
self.avail_quote = self.core.get_balance(self.quote).await?.free;
self.avail_asset = Decimal::ZERO;
self.invested_status = false;
}
}
false => {
if avg_price < level_low {
tracing::debug!("About to sell");
self.core
.send_order(OrderDraft {
asset: self.asset,
quote: self.quote,
side: OrderSide::Buy,
order_type: OrderType::Market,
qty_asset: None,
qty_quote: Some(self.avail_quote),
price: None,
})
.await?;
self.avail_asset = self.core.get_balance(self.asset).await?.free;
self.avail_quote = Decimal::ZERO;
self.invested_status = true;
}
}
}
Ok(())
}
}
impl<C: Clock, M: Market<_Clock = C>, S: Source<_Data = Kline>, P: ProviderTrait<_Source = S>>
Strategy for TrendTunnel<C, M, S, P>
{
type _Clock = C;
type _Market = M;
fn state(&self) -> &Arc<AtomicState> {
&self.future_state
}
fn notify(&self) -> &Arc<Notify> {
&self.notify
}
fn stat_channel(&self) -> Receiver<Result<StratStat, ErrorInStrat>> {
self.rx.clone()
}
fn run(&mut self) -> Pin<Box<dyn Future<Output = ()> + Send + '_>> {
Box::pin(async move {
tracing::info!("Trend tunnel in place");
let mut stop_date = None;
loop {
let last_state = self.state;
self.await_date_or_state_change(stop_date).await;
tracing::trace!("State : {}", self.state);
match self.state {
State::Starting => {
self.future_state.store(State::Running, Ordering::Relaxed);
stop_date = Some(self.run_date);
}
State::Running => {
if let Some(last_date) = stop_date {
tracing::trace!("Running date : {}", last_date);
stop_date = Some(last_date + self.interval.time_delta());
} else {
tracing::error!("Running without stop date");
return;
}
if let Err(e) = self.decide().await {
tracing::error!("Error in decide: {}", e);
return;
};
}
State::Stats => {
self.send_stats().await;
self.state = last_state;
}
State::Stopped => unimplemented!(),
State::Paused => {
stop_date = None;
tracing::info!("Trend tunnel is paused");
}
State::Terminated => {
tracing::info!("Trend tunnel terminating");
let asset_balance = self.core.get_balance(self.asset).await.unwrap();
self.core
.send_order(OrderDraft {
asset: self.asset,
quote: self.quote,
side: OrderSide::Sell,
order_type: OrderType::Market,
qty_asset: Some(asset_balance.free),
qty_quote: None,
price: None,
})
.await
.unwrap();
self.send_stats().await;
break;
}
State::Killed => {
break;
}
};
}
self.send_stats().await;
})
}
}