use crate::error::{CoreError, Result};
use crate::models::{Balance, Order, OrderType, Trade};
use crate::pricing::{FeeBreakdown, FeeSchedule};
use chrono::Utc;
use rust_decimal::Decimal;
use rust_decimal_macros::dec;
use sqlx::PgPool;
use uuid::Uuid;
pub struct TradeExecutor {
pool: PgPool,
fee_schedule: FeeSchedule,
}
impl TradeExecutor {
pub fn new(pool: PgPool) -> Self {
Self {
pool,
fee_schedule: FeeSchedule::default(),
}
}
pub fn with_fee_schedule(pool: PgPool, fee_schedule: FeeSchedule) -> Self {
Self { pool, fee_schedule }
}
pub fn calculate_fees(&self, amount_btc: Decimal) -> FeeBreakdown {
self.fee_schedule.calculate(amount_btc)
}
pub async fn execute_buy_order(&self, order_id: Uuid) -> Result<Trade> {
let mut tx = self.pool.begin().await?;
let order: Order = sqlx::query_as(
r#"
SELECT * FROM orders
WHERE order_id = $1 AND status = 'pending'
FOR UPDATE
"#,
)
.bind(order_id)
.fetch_optional(&mut *tx)
.await?
.ok_or_else(|| CoreError::NotFound("Order not found".to_string()))?;
if order.order_type != OrderType::Buy {
return Err(CoreError::Validation(
"Order is not a buy order".to_string(),
));
}
let (circulating_supply, total_supply, issuer_user_id): (Decimal, Decimal, Uuid) =
sqlx::query_as(
r#"
SELECT circulating_supply, total_supply, issuer_user_id FROM tokens
WHERE token_id = $1
FOR UPDATE
"#,
)
.bind(order.token_id)
.fetch_one(&mut *tx)
.await?;
if circulating_supply + order.amount > total_supply {
return Err(CoreError::Validation(
"Would exceed total supply".to_string(),
));
}
let fees = self.calculate_fees(order.total_btc);
sqlx::query(
r#"
UPDATE tokens
SET circulating_supply = circulating_supply + $1
WHERE token_id = $2
"#,
)
.bind(order.amount)
.bind(order.token_id)
.execute(&mut *tx)
.await?;
sqlx::query(
r#"
INSERT INTO balances (balance_id, user_id, token_id, amount, locked_amount, updated_at)
VALUES ($1, $2, $3, $4, 0, NOW())
ON CONFLICT (user_id, token_id)
DO UPDATE SET amount = balances.amount + $4, updated_at = NOW()
"#,
)
.bind(Uuid::new_v4())
.bind(order.user_id)
.bind(order.token_id)
.bind(order.amount)
.execute(&mut *tx)
.await?;
let trade_id = Uuid::new_v4();
let trade = Trade {
trade_id,
buyer_user_id: order.user_id,
seller_user_id: None, token_id: order.token_id,
amount: order.amount,
price_btc: order.price_btc,
total_btc: order.total_btc,
platform_fee_btc: fees.platform_fee_btc,
issuer_royalty_btc: fees.issuer_royalty_btc,
executed_at: Utc::now(),
};
sqlx::query(
r#"
INSERT INTO trades (trade_id, buyer_user_id, seller_user_id, token_id, amount, price_btc, total_btc, platform_fee_btc, issuer_royalty_btc, executed_at)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10)
"#,
)
.bind(trade.trade_id)
.bind(trade.buyer_user_id)
.bind(trade.seller_user_id)
.bind(trade.token_id)
.bind(trade.amount)
.bind(trade.price_btc)
.bind(trade.total_btc)
.bind(trade.platform_fee_btc)
.bind(trade.issuer_royalty_btc)
.bind(trade.executed_at)
.execute(&mut *tx)
.await?;
sqlx::query(
r#"
UPDATE orders
SET status = 'completed', completed_at = NOW()
WHERE order_id = $1
"#,
)
.bind(order_id)
.execute(&mut *tx)
.await?;
if fees.issuer_royalty_btc > dec!(0) {
sqlx::query(
r#"
INSERT INTO issuer_earnings (earning_id, user_id, token_id, trade_id, amount_btc, created_at)
VALUES ($1, $2, $3, $4, $5, NOW())
"#,
)
.bind(Uuid::new_v4())
.bind(issuer_user_id)
.bind(order.token_id)
.bind(trade.trade_id)
.bind(fees.issuer_royalty_btc)
.execute(&mut *tx)
.await
.ok(); }
tx.commit().await?;
tracing::info!(
order_id = %order_id,
trade_id = %trade_id,
amount = %order.amount,
fees = %fees.total_fees_btc,
"Buy order executed successfully"
);
Ok(trade)
}
pub async fn execute_sell_order(&self, order_id: Uuid) -> Result<Trade> {
let mut tx = self.pool.begin().await?;
let order: Order = sqlx::query_as(
r#"
SELECT * FROM orders
WHERE order_id = $1 AND status = 'pending'
FOR UPDATE
"#,
)
.bind(order_id)
.fetch_optional(&mut *tx)
.await?
.ok_or_else(|| CoreError::NotFound("Order not found".to_string()))?;
if order.order_type != OrderType::Sell {
return Err(CoreError::Validation(
"Order is not a sell order".to_string(),
));
}
let balance: Balance = sqlx::query_as(
r#"
SELECT * FROM balances
WHERE user_id = $1 AND token_id = $2
FOR UPDATE
"#,
)
.bind(order.user_id)
.bind(order.token_id)
.fetch_optional(&mut *tx)
.await?
.ok_or_else(|| CoreError::InsufficientBalance {
required: order.amount,
available: dec!(0),
})?;
if balance.available() < order.amount {
return Err(CoreError::InsufficientBalance {
required: order.amount,
available: balance.available(),
});
}
let (circulating_supply, issuer_user_id): (Decimal, Uuid) = sqlx::query_as(
r#"
SELECT circulating_supply, issuer_user_id FROM tokens
WHERE token_id = $1
FOR UPDATE
"#,
)
.bind(order.token_id)
.fetch_one(&mut *tx)
.await?;
if order.amount > circulating_supply {
return Err(CoreError::Validation(
"Cannot sell more than circulating supply".to_string(),
));
}
let fees = self.calculate_fees(order.total_btc);
sqlx::query(
r#"
UPDATE balances
SET amount = amount - $1, updated_at = NOW()
WHERE user_id = $2 AND token_id = $3
"#,
)
.bind(order.amount)
.bind(order.user_id)
.bind(order.token_id)
.execute(&mut *tx)
.await?;
sqlx::query(
r#"
UPDATE tokens
SET circulating_supply = circulating_supply - $1
WHERE token_id = $2
"#,
)
.bind(order.amount)
.bind(order.token_id)
.execute(&mut *tx)
.await?;
let trade_id = Uuid::new_v4();
let trade = Trade {
trade_id,
buyer_user_id: Uuid::nil(), seller_user_id: Some(order.user_id),
token_id: order.token_id,
amount: order.amount,
price_btc: order.price_btc,
total_btc: order.total_btc,
platform_fee_btc: fees.platform_fee_btc,
issuer_royalty_btc: fees.issuer_royalty_btc,
executed_at: Utc::now(),
};
sqlx::query(
r#"
INSERT INTO trades (trade_id, buyer_user_id, seller_user_id, token_id, amount, price_btc, total_btc, platform_fee_btc, issuer_royalty_btc, executed_at)
VALUES ($1, $2, $3, $4, $5, $6, $7, $8, $9, $10)
"#,
)
.bind(trade.trade_id)
.bind(trade.buyer_user_id)
.bind(trade.seller_user_id)
.bind(trade.token_id)
.bind(trade.amount)
.bind(trade.price_btc)
.bind(trade.total_btc)
.bind(trade.platform_fee_btc)
.bind(trade.issuer_royalty_btc)
.bind(trade.executed_at)
.execute(&mut *tx)
.await?;
sqlx::query(
r#"
UPDATE orders
SET status = 'completed', completed_at = NOW()
WHERE order_id = $1
"#,
)
.bind(order_id)
.execute(&mut *tx)
.await?;
let net_proceeds = fees.net_amount_btc;
sqlx::query(
r#"
INSERT INTO pending_payouts (payout_id, user_id, trade_id, amount_btc, status, created_at)
VALUES ($1, $2, $3, $4, 'pending', NOW())
"#,
)
.bind(Uuid::new_v4())
.bind(order.user_id)
.bind(trade.trade_id)
.bind(net_proceeds)
.execute(&mut *tx)
.await
.ok();
if fees.issuer_royalty_btc > dec!(0) {
sqlx::query(
r#"
INSERT INTO issuer_earnings (earning_id, user_id, token_id, trade_id, amount_btc, created_at)
VALUES ($1, $2, $3, $4, $5, NOW())
"#,
)
.bind(Uuid::new_v4())
.bind(issuer_user_id)
.bind(order.token_id)
.bind(trade.trade_id)
.bind(fees.issuer_royalty_btc)
.execute(&mut *tx)
.await
.ok();
}
tx.commit().await?;
tracing::info!(
order_id = %order_id,
trade_id = %trade_id,
amount = %order.amount,
net_proceeds = %net_proceeds,
"Sell order executed successfully"
);
Ok(trade)
}
pub async fn cancel_order(&self, order_id: Uuid, user_id: Uuid) -> Result<()> {
let result = sqlx::query(
r#"
UPDATE orders
SET status = 'cancelled', completed_at = NOW()
WHERE order_id = $1 AND user_id = $2 AND status = 'pending'
"#,
)
.bind(order_id)
.bind(user_id)
.execute(&self.pool)
.await?;
if result.rows_affected() == 0 {
return Err(CoreError::NotFound(
"Order not found or already processed".to_string(),
));
}
tracing::info!(order_id = %order_id, "Order cancelled");
Ok(())
}
pub async fn expire_stale_orders(&self, max_age_hours: i64) -> Result<u64> {
let result = sqlx::query(
r#"
UPDATE orders
SET status = 'expired', completed_at = NOW()
WHERE status = 'pending'
AND created_at < NOW() - INTERVAL '1 hour' * $1
"#,
)
.bind(max_age_hours)
.execute(&self.pool)
.await?;
let count = result.rows_affected();
if count > 0 {
tracing::info!(count = count, "Expired stale orders");
}
Ok(count)
}
pub async fn execute_buy_orders_batch(
&self,
order_ids: Vec<Uuid>,
) -> Result<BatchExecutionResult> {
let mut successful_trades = Vec::new();
let mut failed_orders = Vec::new();
for order_id in order_ids {
match self.execute_buy_order(order_id).await {
Ok(trade) => successful_trades.push(trade),
Err(e) => {
tracing::warn!(order_id = %order_id, error = %e, "Failed to execute buy order in batch");
failed_orders.push((order_id, e.to_string()));
}
}
}
tracing::info!(
successful = successful_trades.len(),
failed = failed_orders.len(),
"Batch buy order execution completed"
);
Ok(BatchExecutionResult {
successful_trades,
failed_orders,
})
}
pub async fn execute_sell_orders_batch(
&self,
order_ids: Vec<Uuid>,
) -> Result<BatchExecutionResult> {
let mut successful_trades = Vec::new();
let mut failed_orders = Vec::new();
for order_id in order_ids {
match self.execute_sell_order(order_id).await {
Ok(trade) => successful_trades.push(trade),
Err(e) => {
tracing::warn!(order_id = %order_id, error = %e, "Failed to execute sell order in batch");
failed_orders.push((order_id, e.to_string()));
}
}
}
tracing::info!(
successful = successful_trades.len(),
failed = failed_orders.len(),
"Batch sell order execution completed"
);
Ok(BatchExecutionResult {
successful_trades,
failed_orders,
})
}
pub async fn get_pending_orders_by_token(&self, token_id: Uuid) -> Result<Vec<Order>> {
let orders = sqlx::query_as::<_, Order>(
r#"
SELECT * FROM orders
WHERE token_id = $1 AND status = 'pending'
ORDER BY created_at ASC
"#,
)
.bind(token_id)
.fetch_all(&self.pool)
.await?;
Ok(orders)
}
pub async fn get_pending_orders_by_user(&self, user_id: Uuid) -> Result<Vec<Order>> {
let orders = sqlx::query_as::<_, Order>(
r#"
SELECT * FROM orders
WHERE user_id = $1 AND status = 'pending'
ORDER BY created_at ASC
"#,
)
.bind(user_id)
.fetch_all(&self.pool)
.await?;
Ok(orders)
}
}
#[derive(Debug)]
pub struct BatchExecutionResult {
pub successful_trades: Vec<Trade>,
pub failed_orders: Vec<(Uuid, String)>,
}
impl BatchExecutionResult {
pub fn total_processed(&self) -> usize {
self.successful_trades.len() + self.failed_orders.len()
}
pub fn success_rate(&self) -> Decimal {
if self.total_processed() == 0 {
return dec!(0);
}
let successful = Decimal::from(self.successful_trades.len());
let total = Decimal::from(self.total_processed());
(successful / total) * dec!(100)
}
pub fn all_successful(&self) -> bool {
self.failed_orders.is_empty()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::pricing::FeeSchedule;
use chrono::Utc;
use rust_decimal_macros::dec;
fn make_trade() -> Trade {
Trade {
trade_id: Uuid::new_v4(),
buyer_user_id: Uuid::new_v4(),
seller_user_id: None,
token_id: Uuid::new_v4(),
amount: dec!(10),
price_btc: dec!(0.001),
total_btc: dec!(0.01),
platform_fee_btc: dec!(0.00025),
issuer_royalty_btc: dec!(0.00005),
executed_at: Utc::now(),
}
}
#[test]
fn test_total_processed_empty_batch() {
let result = BatchExecutionResult {
successful_trades: vec![],
failed_orders: vec![],
};
assert_eq!(
result.total_processed(),
0,
"Empty batch must report zero total processed"
);
}
#[test]
fn test_total_processed_all_successful() {
let result = BatchExecutionResult {
successful_trades: vec![make_trade(), make_trade(), make_trade()],
failed_orders: vec![],
};
assert_eq!(
result.total_processed(),
3,
"All-successful batch must count only the successful trades"
);
}
#[test]
fn test_total_processed_all_failed() {
let result = BatchExecutionResult {
successful_trades: vec![],
failed_orders: vec![
(Uuid::new_v4(), "supply exceeded".to_string()),
(Uuid::new_v4(), "order not found".to_string()),
],
};
assert_eq!(
result.total_processed(),
2,
"All-failed batch must count only the failed orders"
);
}
#[test]
fn test_total_processed_mixed_batch() {
let result = BatchExecutionResult {
successful_trades: vec![make_trade(), make_trade()],
failed_orders: vec![(Uuid::new_v4(), "validation error".to_string())],
};
assert_eq!(
result.total_processed(),
3,
"Mixed batch must sum successful and failed counts"
);
}
#[test]
fn test_success_rate_empty_batch_returns_zero() {
let result = BatchExecutionResult {
successful_trades: vec![],
failed_orders: vec![],
};
assert_eq!(
result.success_rate(),
dec!(0),
"Empty batch must return success_rate of 0"
);
}
#[test]
fn test_success_rate_all_successful_is_100() {
let result = BatchExecutionResult {
successful_trades: vec![make_trade(), make_trade()],
failed_orders: vec![],
};
assert_eq!(
result.success_rate(),
dec!(100),
"All-successful batch must return success_rate of 100"
);
}
#[test]
fn test_success_rate_all_failed_is_zero() {
let result = BatchExecutionResult {
successful_trades: vec![],
failed_orders: vec![
(Uuid::new_v4(), "err".to_string()),
(Uuid::new_v4(), "err".to_string()),
],
};
assert_eq!(
result.success_rate(),
dec!(0),
"All-failed batch must return success_rate of 0"
);
}
#[test]
fn test_success_rate_fifty_percent() {
let result = BatchExecutionResult {
successful_trades: vec![make_trade()],
failed_orders: vec![(Uuid::new_v4(), "err".to_string())],
};
assert_eq!(
result.success_rate(),
dec!(50),
"One success out of two total must return success_rate of 50"
);
}
#[test]
fn test_success_rate_one_in_four_is_25() {
let result = BatchExecutionResult {
successful_trades: vec![make_trade()],
failed_orders: vec![
(Uuid::new_v4(), "e".to_string()),
(Uuid::new_v4(), "e".to_string()),
(Uuid::new_v4(), "e".to_string()),
],
};
assert_eq!(
result.success_rate(),
dec!(25),
"One success out of four total must return success_rate of 25"
);
}
#[test]
fn test_all_successful_true_when_no_failures() {
let result = BatchExecutionResult {
successful_trades: vec![make_trade()],
failed_orders: vec![],
};
assert!(
result.all_successful(),
"all_successful must be true when failed_orders is empty"
);
}
#[test]
fn test_all_successful_false_when_any_failure() {
let result = BatchExecutionResult {
successful_trades: vec![make_trade()],
failed_orders: vec![(Uuid::new_v4(), "one failure".to_string())],
};
assert!(
!result.all_successful(),
"all_successful must be false when there is at least one failure"
);
}
#[test]
fn test_all_successful_true_for_empty_batch() {
let result = BatchExecutionResult {
successful_trades: vec![],
failed_orders: vec![],
};
assert!(
result.all_successful(),
"all_successful must be true for an empty batch (no failures)"
);
}
#[test]
fn test_fee_calculation_default_schedule() {
let schedule = FeeSchedule::default();
let breakdown = schedule.calculate(dec!(1.0));
assert_eq!(
breakdown.platform_fee_btc,
dec!(0.025),
"Default platform fee must be 2.5% of trade amount"
);
assert_eq!(
breakdown.issuer_royalty_btc,
dec!(0.005),
"Default issuer royalty must be 0.5% of trade amount"
);
assert_eq!(
breakdown.net_amount_btc,
dec!(0.97),
"Net amount must be trade amount minus all fees"
);
assert_eq!(
breakdown.total_fees_btc,
dec!(0.03),
"Total fees must be 3% of trade amount"
);
}
#[test]
fn test_fee_calculation_zero_amount() {
let schedule = FeeSchedule::default();
let breakdown = schedule.calculate(dec!(0));
assert_eq!(breakdown.platform_fee_btc, dec!(0));
assert_eq!(breakdown.issuer_royalty_btc, dec!(0));
assert_eq!(breakdown.total_fees_btc, dec!(0));
assert_eq!(breakdown.net_amount_btc, dec!(0));
}
#[test]
fn test_fee_schedule_with_discount() {
let schedule = FeeSchedule::default();
let discounted = schedule.with_discount(dec!(50));
assert_eq!(discounted.platform_fee_rate, dec!(0.0125));
assert_eq!(
discounted.issuer_royalty_rate, schedule.issuer_royalty_rate,
"Royalty rate must not change with discount"
);
}
#[test]
fn test_fee_schedule_discount_capped_at_50pct() {
let schedule = FeeSchedule::default();
let discounted = schedule.with_discount(dec!(100));
let capped = schedule.with_discount(dec!(50));
assert_eq!(
discounted.platform_fee_rate, capped.platform_fee_rate,
"Discount must be capped at 50% of the platform fee rate"
);
}
}