quantoxide 0.6.2

Rust framework for developing, backtesting, and deploying Bitcoin futures trading strategies.
Documentation
use std::{collections::HashMap, sync::Arc};

use async_trait::async_trait;
use lnm_sdk::rest::v3::models::PercentageCapped;
use sqlx::{QueryBuilder, Sqlite, SqlitePool};
use uuid::Uuid;

use crate::{
    db::{
        error::{DbError, Result},
        models::RunningTrade,
        repositories::{RunningTradesRepository, RunningTradesRepositoryRead},
    },
    trade::TradeTrailingStoploss,
};

pub(crate) struct SqliteRunningTradesRepo {
    pool: Arc<SqlitePool>,
}

impl SqliteRunningTradesRepo {
    pub(crate) fn new(pool: Arc<SqlitePool>) -> Self {
        Self { pool }
    }

    fn pool(&self) -> &SqlitePool {
        self.pool.as_ref()
    }
}

#[async_trait]
impl RunningTradesRepositoryRead for SqliteRunningTradesRepo {
    async fn get_running_trades_map(
        &self,
        account_id: Uuid,
    ) -> Result<HashMap<Uuid, Option<TradeTrailingStoploss>>> {
        let running_trades = sqlx::query_as!(
            RunningTrade,
            r#"
                SELECT
                    trade_id as "trade_id!: Uuid",
                    trailing_stoploss as "trailing_stoploss?: f64"
                FROM running_trades
                WHERE account_id = ?1
                ORDER BY created_at ASC
            "#,
            account_id,
        )
        .fetch_all(self.pool())
        .await
        .map_err(DbError::Query)?;

        let mut running_trades_map = HashMap::new();

        for trade in running_trades.into_iter() {
            let trailing_stoploss = trade
                .trailing_stoploss
                .map(|tsl| {
                    PercentageCapped::try_from(tsl)
                        .map_err(|e| {
                            DbError::UnexpectedQueryResult(format!(
                                "`trailing_stoploss` ({tsl}) cannot be casted as `PercentageCapped`: {e}"
                            ))
                        })
                        .map(TradeTrailingStoploss::prev_validated)
                })
                .transpose()?;

            running_trades_map.insert(trade.trade_id, trailing_stoploss);
        }

        Ok(running_trades_map)
    }
}

#[async_trait]
impl RunningTradesRepository for SqliteRunningTradesRepo {
    async fn add_running_trade(
        &self,
        account_id: Uuid,
        trade_id: Uuid,
        trailing_stoploss: Option<TradeTrailingStoploss>,
    ) -> Result<()> {
        sqlx::query!(
            r#"
                INSERT INTO running_trades (account_id, trade_id, trailing_stoploss)
                VALUES (?1, ?2, ?3)
            "#,
            account_id,
            trade_id,
            trailing_stoploss.map(|tsl| tsl.as_f64()),
        )
        .execute(self.pool())
        .await
        .map_err(DbError::Query)?;

        Ok(())
    }

    async fn remove_running_trades(&self, account_id: Uuid, trade_ids: &[Uuid]) -> Result<()> {
        if trade_ids.is_empty() {
            return Ok(());
        }

        let mut query_builder =
            QueryBuilder::<Sqlite>::new("DELETE FROM running_trades WHERE account_id = ");
        query_builder.push_bind(account_id);
        query_builder.push(" AND trade_id IN (");

        let mut separated = query_builder.separated(", ");
        for trade_id in trade_ids {
            separated.push_bind(trade_id);
        }
        separated.push_unseparated(")");

        query_builder
            .build()
            .execute(self.pool())
            .await
            .map_err(DbError::Query)?;

        Ok(())
    }
}

#[cfg(test)]
mod tests {
    use std::sync::Arc;

    use lnm_sdk::rest::v3::models::PercentageCapped;
    use sqlx::sqlite::SqlitePoolOptions;

    use super::*;

    async fn repo() -> SqliteRunningTradesRepo {
        let pool = SqlitePoolOptions::new()
            .max_connections(1)
            .connect("sqlite::memory:")
            .await
            .unwrap();

        sqlx::migrate!("./migrations/sqlite")
            .run(&pool)
            .await
            .unwrap();

        SqliteRunningTradesRepo::new(Arc::new(pool))
    }

    fn trailing_stoploss(value: f64) -> TradeTrailingStoploss {
        TradeTrailingStoploss::prev_validated(PercentageCapped::try_from(value).unwrap())
    }

    #[tokio::test]
    async fn running_trades_are_scoped_by_account() {
        let repo = repo().await;
        let account_a = Uuid::new_v4();
        let account_b = Uuid::new_v4();
        let shared_trade_id = Uuid::new_v4();

        repo.add_running_trade(account_a, shared_trade_id, Some(trailing_stoploss(5.0)))
            .await
            .unwrap();
        repo.add_running_trade(account_b, shared_trade_id, None)
            .await
            .unwrap();

        let account_a_trades = repo.get_running_trades_map(account_a).await.unwrap();
        assert_eq!(account_a_trades.len(), 1);
        assert_eq!(
            account_a_trades.get(&shared_trade_id).copied().flatten(),
            Some(trailing_stoploss(5.0))
        );

        let account_b_trades = repo.get_running_trades_map(account_b).await.unwrap();
        assert_eq!(account_b_trades.len(), 1);
        assert_eq!(
            account_b_trades.get(&shared_trade_id).copied().flatten(),
            None
        );
    }

    #[tokio::test]
    async fn remove_running_trades_deletes_only_requested_trades_for_account() {
        let repo = repo().await;
        let account_a = Uuid::new_v4();
        let account_b = Uuid::new_v4();
        let removed_trade_id = Uuid::new_v4();
        let kept_trade_id = Uuid::new_v4();

        repo.add_running_trade(account_a, removed_trade_id, None)
            .await
            .unwrap();
        repo.add_running_trade(account_a, kept_trade_id, None)
            .await
            .unwrap();
        repo.add_running_trade(account_b, removed_trade_id, None)
            .await
            .unwrap();

        repo.remove_running_trades(account_a, &[removed_trade_id])
            .await
            .unwrap();

        let account_a_trades = repo.get_running_trades_map(account_a).await.unwrap();
        assert!(!account_a_trades.contains_key(&removed_trade_id));
        assert!(account_a_trades.contains_key(&kept_trade_id));

        let account_b_trades = repo.get_running_trades_map(account_b).await.unwrap();
        assert!(account_b_trades.contains_key(&removed_trade_id));
    }

    #[tokio::test]
    async fn remove_running_trades_accepts_empty_trade_ids() {
        let repo = repo().await;
        let account_id = Uuid::new_v4();
        let trade_id = Uuid::new_v4();

        repo.add_running_trade(account_id, trade_id, None)
            .await
            .unwrap();

        repo.remove_running_trades(account_id, &[]).await.unwrap();

        let trades = repo.get_running_trades_map(account_id).await.unwrap();
        assert!(trades.contains_key(&trade_id));
    }
}