scuriolus 0.3.0

Scuriolus is a modular trading bot platform.
Documentation
use chrono::{DateTime, TimeDelta, Utc};
use std::{
    cmp::min,
    sync::{
        Arc,
        atomic::{AtomicI64, Ordering},
    },
};
use tokio::{select, sync::Notify, time::sleep};

use super::{Clock, ClockBase, ClockFactory, ClockRemote, SAFETY_REFRESH_DELAY};
use crate::core::{CoreError, CoreResult};

trait RunningCheatClockBase {
    fn offset_ms(&self) -> &Arc<AtomicI64>;
    fn running_now(&self) -> DateTime<Utc> {
        let now = Utc::now();
        let offset = self.offset_ms().load(Ordering::Relaxed);
        now - TimeDelta::milliseconds(offset)
    }
}

/// The [`RunningCheatClock`] is a [`Clock`] that mimic a real clock but starting in the past and that can be advanced.
#[derive(Clone, Debug, Default)]
pub struct RunningCheatClock {
    offset_ms: Arc<AtomicI64>,
    remote_offset_ms: Arc<AtomicI64>,
    notify: Arc<Notify>,
}

impl RunningCheatClockBase for RunningCheatClock {
    fn offset_ms(&self) -> &Arc<AtomicI64> {
        &self.offset_ms
    }
}

impl ClockBase for RunningCheatClock {
    const NAME: &'static str = "RunningCheatClock";
    fn now(&self) -> DateTime<Utc> {
        self.running_now()
    }
}

impl Clock for RunningCheatClock {
    type Remote = RunningCheatClockRemote;

    async fn sleep(&self, delay: TimeDelta) {
        self.sleep_until(self.now() + delay).await
    }

    async fn sleep_until(&self, date: DateTime<Utc>) {
        loop {
            let offset = self.offset_ms.load(Ordering::Relaxed);
            let remote_offset = self.remote_offset_ms.load(Ordering::Relaxed);
            let potential_new_offset = (Utc::now() - date).num_milliseconds();

            #[cfg(any(test, debug_assertions))]
            assert!(offset >= remote_offset);

            if offset != remote_offset && offset != potential_new_offset {
                if potential_new_offset <= remote_offset {
                    self.offset_ms.store(remote_offset, Ordering::Relaxed);
                    self.notify.notify_waiters();
                } else {
                    self.offset_ms
                        .store(potential_new_offset, Ordering::Relaxed);
                }
            }

            let now = self.now();
            if now >= date {
                break;
            }

            select! {
                _ = sleep(min(TimeDelta::to_std(&(date-now)).unwrap_or_else(|_| {
                    tracing::error!("Failed to sleep, conversion error from TimeDelta to Duration");
                    SAFETY_REFRESH_DELAY
                }), SAFETY_REFRESH_DELAY)) => {},
                _ = self.notify.notified() => {}
            }
        }
    }

    fn synchronize(&self) {
        self.offset_ms.store(
            self.remote_offset_ms.load(Ordering::Relaxed),
            Ordering::Relaxed,
        );
    }
}

/// A [`ClockRemote`] for [`RunningCheatClock`].
#[derive(Clone, Debug)]
pub struct RunningCheatClockRemote {
    offset_ms: Arc<AtomicI64>,
    remote_offset_ms: Arc<AtomicI64>,
    notify: Arc<Notify>,
}

impl RunningCheatClockBase for RunningCheatClockRemote {
    fn offset_ms(&self) -> &Arc<AtomicI64> {
        &self.offset_ms
    }
}

impl ClockBase for RunningCheatClockRemote {
    const NAME: &'static str = "RunningCheatClockRemote";
    fn now(&self) -> DateTime<Utc> {
        self.running_now()
    }
}

impl ClockRemote for RunningCheatClockRemote {
    #[cfg(test)]
    fn not_blocking(&self) {
        self.remote_offset_ms.store(0, Ordering::Relaxed);
    }

    fn advance_toward(&self, date: DateTime<Utc>) {
        self.remote_offset_ms
            .store((Utc::now() - date).num_milliseconds(), Ordering::Relaxed);
        self.notify.notify_waiters();
    }

    fn remote_now(&self) -> DateTime<Utc> {
        let now = Utc::now();
        let offset = self.remote_offset_ms.load(Ordering::Relaxed);
        now - TimeDelta::milliseconds(offset)
    }

    fn notify(&self) -> &Notify {
        &self.notify
    }

    fn is_sync(&self) -> bool {
        self.offset_ms.load(Ordering::Relaxed) == self.remote_offset_ms.load(Ordering::Relaxed)
    }
}

/// A [`ClockFactory`] for [`RunningCheatClock`].
#[derive(Clone, Debug)]
pub struct RunningCheatClockFactory {
    pub date: DateTime<Utc>,
}

impl ClockFactory for RunningCheatClockFactory {
    type _Clock = RunningCheatClock;

    fn build(self) -> CoreResult<(Option<RunningCheatClockRemote>, RunningCheatClock)> {
        let now = Utc::now();
        if now < self.date {
            return Err(CoreError::param_error("Can't set time in the future"));
        }

        let offset_base = (now - self.date).num_milliseconds();

        let offset_ms = Arc::new(AtomicI64::new(offset_base));
        let remote_offset_ms = Arc::new(AtomicI64::new(offset_base));
        let notify = Arc::new(Notify::new());
        Ok((
            Some(RunningCheatClockRemote {
                offset_ms: offset_ms.clone(),
                remote_offset_ms: remote_offset_ms.clone(),
                notify: notify.clone(),
            }),
            RunningCheatClock {
                offset_ms,
                remote_offset_ms,
                notify,
            },
        ))
    }
}

#[cfg(test)]
mod tests {
    use std::{
        panic,
        str::FromStr as _,
        sync::{
            Arc,
            atomic::{AtomicI16, Ordering},
        },
    };

    use chrono::{DateTime, TimeDelta, Utc};
    use tokio::time::sleep as tokio_sleep;

    use crate::clock::ClockBase as _;

    use super::{
        super::REMOTE_REFRESH_RATE, Clock as _, ClockFactory as _, ClockRemote as _,
        RunningCheatClockFactory,
    };

    const TOLERANCE: TimeDelta = TimeDelta::milliseconds(10);
    const DELAY: TimeDelta = TimeDelta::milliseconds(100);

    #[tokio::test]
    async fn sleep() {
        let date = DateTime::from_str("2012-12-12 00:00:00Z").unwrap();
        let (remote, clock) = RunningCheatClockFactory { date }.build().unwrap();

        remote.unwrap().not_blocking();

        assert!(clock.now() - date < TOLERANCE);

        let before = Utc::now();
        clock.sleep(DELAY).await;
        let after = Utc::now();

        assert!(after - before < TOLERANCE);

        assert!(clock.now() - date > DELAY - TOLERANCE);
        assert!(clock.now() - date < DELAY + TOLERANCE);
    }

    #[tokio::test]
    async fn sleep_until() {
        let date = DateTime::from_str("2012-12-12 00:00:00Z").unwrap();
        let (remote, clock) = RunningCheatClockFactory { date }.build().unwrap();

        remote.unwrap().not_blocking();

        assert!(clock.now() - date < TOLERANCE);

        let before = Utc::now();
        clock.sleep_until(clock.now() + DELAY).await;
        let after = Utc::now();

        assert!(after - before < TOLERANCE);

        assert!(clock.now() - date > DELAY - TOLERANCE);
        assert!(clock.now() - date < DELAY + TOLERANCE);
    }

    #[tokio::test]
    async fn remote() {
        let date = DateTime::from_str("2012-12-12 00:00:00Z").unwrap();
        let (_remote, clock) = RunningCheatClockFactory { date }.build().unwrap();
        let remote = _remote.unwrap();

        assert!(clock.now() - date < TOLERANCE);
        assert!(remote.now() - date < TOLERANCE);

        let marker = Arc::new(AtomicI16::new(0));
        let marker_clone = marker.clone();

        tokio::spawn(async move {
            clock.sleep(DELAY).await;
            marker_clone.store(1, Ordering::Relaxed);
            clock.sleep(DELAY).await;
            marker_clone.store(2, Ordering::Relaxed);
            clock
                .sleep_until(DateTime::from_str("2020-01-01 00:00:00Z").unwrap())
                .await;
            panic!("should not reach here");
        });

        assert_eq!(marker.load(Ordering::Relaxed), 0);
        tokio_sleep(DELAY.to_std().unwrap()).await;
        //now wait for the clock to auto-update
        tokio_sleep(REMOTE_REFRESH_RATE + TOLERANCE.to_std().unwrap()).await;
        assert_eq!(marker.load(Ordering::Relaxed), 1);

        remote.advance_toward(date + DELAY + DELAY);

        //now wait for the arc to update
        tokio_sleep(TOLERANCE.to_std().unwrap()).await;

        assert_eq!(marker.load(Ordering::Relaxed), 2);
    }
}