scuriolus 0.3.0

Scuriolus is a modular trading bot platform.
Documentation
use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use std::marker::PhantomData;
use surrealdb::Surreal;
use surrealdb::engine::local::Db;

use crate::core::{CoreError, CoreResult};
use crate::provider::data::{Data, DataBasics, DataQuery, DataSpecifier};
use crate::provider::interval::Interval;

#[derive(Debug, Serialize, Deserialize)]
struct DataSection {
    pub begin: DateTime<Utc>,
    pub end: DateTime<Utc>,
}

#[derive(Debug, Serialize, Deserialize)]
#[serde(bound = "D: Data")]
struct DataWrapper<D: Data> {
    begin: DateTime<Utc>,
    end: DateTime<Utc>,
    data: D,
}

impl<D: Data> From<D> for DataWrapper<D> {
    fn from(data: D) -> Self {
        Self {
            begin: data.basics().begin,
            end: data.basics().end,
            data,
        }
    }
}

impl<D: Data> DataWrapper<D> {
    fn into_data(self) -> D {
        self.data
    }
}

/// A [`Store`] is an object responsible for storing and retrieving data based on dates.
#[derive(Debug)]
pub struct Store<D: Data> {
    db: Surreal<Db>,
    _marker: PhantomData<D>,
}

impl<D: Data> Clone for Store<D> {
    fn clone(&self) -> Self {
        Self {
            db: self.db.clone(),
            _marker: PhantomData,
        }
    }
}

impl<D: Data> Store<D> {
    /// Create a new [`Store`]. Should only called once per [`Market`].
    pub fn new(db: Surreal<Db>) -> Self {
        Self {
            db,
            _marker: PhantomData,
        }
    }

    #[allow(dead_code)]
    fn check_basics(&self, interval: Interval, basics: &DataBasics) -> CoreResult<()> {
        let (begin, end) = interval.get_time_bounds(basics.begin)?;

        if begin != basics.begin || end != basics.end {
            Err(CoreError::ParamError(format!(
                "Wrong time bounds in {basics:?} for interval {interval} (begin: {begin}, end: {end})"
            )))
        } else {
            Ok(())
        }
    }

    pub async fn inject(
        &self,
        specifier: &D::Specifier,
        interval: Interval,
        elems: Vec<D>,
    ) -> CoreResult<()> {
        if elems.is_empty() {
            tracing::warn!("No elements to inject");
            return Ok(());
        }

        let first = elems.first().unwrap();

        #[cfg(test)]
        {
            let mut begin = first.basics().begin;
            for elem in &elems {
                let basics = elem.basics();
                self.check_basics(interval, basics)?;
                if begin != basics.begin {
                    return Err(CoreError::ParamError(format!(
                        "Hole between {begin} and {}",
                        basics.begin
                    )));
                }
                begin = basics.end;
            }
        }

        tracing::debug!(
            "Injecting {} in {} ({})",
            elems.len(),
            specifier.name(),
            interval
        );
        tracing::trace!("elems: {elems:#?}");

        let begin = first.basics().begin;
        let end = elems.last().unwrap().basics().end;

        let mut response = self.db
            .query("SELECT * FROM type::table($register) WHERE begin <= $begin AND end >= $begin ORDER BY begin ASC")
            .query("SELECT * FROM type::table($register) WHERE begin <= $end AND end >= $end ORDER BY end DESC")
            .query("DELETE type::table($register) WHERE begin >= $begin AND begin <= $end OR end >= $begin AND end <= $end OR begin <= $begin AND end >= $end")
            .bind(("register", register_from(specifier, interval)))
            .bind(("begin", begin))
            .bind(("end", end))
            .await?;

        let begin_section: Vec<DataSection> = response.take(0)?;
        let end_section: Vec<DataSection> = response.take(1)?;

        if begin_section.len() > 1 || end_section.len() > 1 {
            return Err(CoreError::db_error(
                "Db corrupted, multiple sections found for a single date",
            ));
        }

        tracing::trace!(
            "begin_section found locally: {:#?}, \nend_section found locally: {:#?}",
            begin_section,
            end_section
        );

        let new_begin = match begin_section.first() {
            Some(section) => section.begin,
            None => begin,
        };
        let new_end = match end_section.first() {
            Some(section) => section.end,
            None => end,
        };

        let section = DataSection {
            begin: new_begin,
            end: new_end,
        };

        let _: Option<DataSection> = self
            .db
            .upsert((
                register_from(specifier, interval),
                section.begin.timestamp(),
            ))
            .content(section)
            .await?;

        for elem in elems {
            let _: Option<DataWrapper<D>> = self
                .db
                .upsert((
                    table_from(specifier, interval),
                    elem.basics().begin.timestamp(),
                ))
                .content(DataWrapper::from(elem))
                .await?;
        }

        Ok(())
    }

    pub async fn try_get_data(&self, params: &DataQuery<D>) -> CoreResult<(Vec<D>, bool)> {
        let mut response = self
            .db
            .query("SELECT * FROM type::table($register) WHERE begin <= $begin AND end >= $begin")
            .bind((
                "register",
                register_from(params.specifier(), *params.interval()),
            ))
            .bind(("begin", *params.begin()))
            .await?;

        let content: Vec<DataSection> = response.take(0)?;

        if content.is_empty() {
            return Ok((Vec::new(), false));
        }

        let section = content.first().unwrap();

        let (last_elem_should_begin, last_elem_might_end) =
            params.interval().get_time_bounds(section.end)?;

        let complete = section.end == *params.end()
            || (last_elem_should_begin > *params.end() && *params.end() <= last_elem_might_end);

        tracing::debug!(
            "Asking DB for {} ({})",
            params.specifier().name(),
            params.interval()
        );
        tracing::trace!(
            "table: {}, begin: {}, end: {}",
            table_from(params.specifier(), *params.interval()),
            params.begin(),
            params.end()
        );

        let mut data_response = self
                .db
                .query("SELECT * FROM type::table($table) WHERE begin >= $begin AND end <= $end_section AND end <= $end ORDER BY begin ASC")
                .bind(("begin", *params.begin()))
                .bind(("end_section", section.end))
                .bind(("end", *params.end()))
                .bind(("table", table_from(params.specifier(), *params.interval())))
                .await?;
        let wrapped_elems: Vec<DataWrapper<D>> = data_response.take(0)?;
        let elems: Vec<D> = wrapped_elems
            .into_iter()
            .map(DataWrapper::into_data)
            .collect();
        Ok((elems, complete))
    }
}

//TODO utiliser hash au lieu de name
//TODO reunir en une seule table mais enjoutant le hash dans l'objet
fn table_from<S: DataSpecifier>(specifier: &S, interval: Interval) -> String {
    format!("{}_{}", specifier.name(), interval)
}

fn register_from<S: DataSpecifier>(specifier: &S, interval: Interval) -> String {
    format!("{}_{}_reg", specifier.name(), interval)
}

#[cfg(test)]
mod tests {
    use chrono::{TimeZone as _, Utc};
    use surrealdb::{Surreal, engine::local::Mem};

    use crate::provider::{
        Store,
        data::test::{TestData, TestDataSpecifier},
        interval::Interval,
        store::{DataBasics, DataQuery},
    };

    fn default_wrapper() -> DataBasics {
        DataBasics {
            source: String::from("test"),
            ..Default::default()
        }
    }

    #[tokio::test]
    async fn simple() {
        let db = Surreal::new::<Mem>(()).await.unwrap();
        db.use_ns("test").use_db("test").await.unwrap();

        let store = Store::<TestData>::new(db);
        let specifier = TestDataSpecifier::default();

        let jan = TestData {
            content: String::from("test"),
            basics: DataBasics {
                begin: Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap(),
                end: Utc.with_ymd_and_hms(2024, 2, 1, 0, 0, 0).unwrap(),
                ..default_wrapper()
            },
        };
        let feb = TestData {
            content: String::from("test"),
            basics: DataBasics {
                begin: Utc.with_ymd_and_hms(2024, 2, 1, 0, 0, 0).unwrap(),
                end: Utc.with_ymd_and_hms(2024, 3, 1, 0, 0, 0).unwrap(),
                ..default_wrapper()
            },
        };
        let mar = TestData {
            content: String::from("test"),
            basics: DataBasics {
                begin: Utc.with_ymd_and_hms(2024, 3, 1, 0, 0, 0).unwrap(),
                end: Utc.with_ymd_and_hms(2024, 4, 1, 0, 0, 0).unwrap(),
                ..default_wrapper()
            },
        };

        let data1 = vec![feb.clone(), mar.clone()];

        store
            .inject(&specifier, Interval::OneMonth, data1)
            .await
            .unwrap();

        let data2 = vec![jan.clone(), feb.clone()];

        store
            .inject(&specifier, Interval::OneMonth, data2)
            .await
            .unwrap();

        let (elems, complete) = store
            .try_get_data(&DataQuery::new(
                TestDataSpecifier::default(),
                Interval::OneMonth,
                Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap(),
                Utc.with_ymd_and_hms(2024, 4, 1, 0, 0, 0).unwrap(),
            ))
            .await
            .unwrap();

        assert!(complete);
        assert_eq!(elems, vec![jan, feb, mar]);
    }

    #[tokio::test]
    async fn uncomplete() {
        let db = Surreal::new::<Mem>(()).await.unwrap();
        db.use_ns("test").use_db("test").await.unwrap();

        let store = Store::<TestData>::new(db);
        let specifier = TestDataSpecifier::default();

        let feb = TestData {
            content: String::from("test"),
            basics: DataBasics {
                begin: Utc.with_ymd_and_hms(2024, 2, 1, 0, 0, 0).unwrap(),
                end: Utc.with_ymd_and_hms(2024, 3, 1, 0, 0, 0).unwrap(),
                ..default_wrapper()
            },
        };
        let mar = TestData {
            content: String::from("test"),
            basics: DataBasics {
                begin: Utc.with_ymd_and_hms(2024, 3, 1, 0, 0, 0).unwrap(),
                end: Utc.with_ymd_and_hms(2024, 4, 1, 0, 0, 0).unwrap(),
                ..default_wrapper()
            },
        };

        let elems = vec![feb.clone(), mar.clone()];

        store
            .inject(&specifier, Interval::OneMonth, elems)
            .await
            .unwrap();

        let (elems, complete) = store
            .try_get_data(&DataQuery::new(
                TestDataSpecifier::default(),
                Interval::OneMonth,
                Utc.with_ymd_and_hms(2024, 1, 1, 0, 0, 0).unwrap(),
                Utc.with_ymd_and_hms(2024, 4, 1, 0, 0, 0).unwrap(),
            ))
            .await
            .unwrap();

        assert!(!complete);
        assert_eq!(elems, vec![]);

        let (elems, complete) = store
            .try_get_data(&DataQuery::new(
                TestDataSpecifier::default(),
                Interval::OneMonth,
                Utc.with_ymd_and_hms(2024, 2, 1, 0, 0, 0).unwrap(),
                Utc.with_ymd_and_hms(2024, 5, 1, 0, 0, 0).unwrap(),
            ))
            .await
            .unwrap();

        assert!(!complete);
        assert_eq!(elems, vec![feb, mar]);
    }
}