use chrono::{DateTime, Utc};
use rust_decimal::Decimal;
use rust_decimal::prelude::ToPrimitive;
use rust_decimal_macros::dec;
use serde::{Deserialize, Serialize};
use uuid::Uuid;
use crate::error::{CoreError, Result};
use crate::models::{
AddLiquidityRequest, CreatePoolRequest, LiquidityPool, LpPosition, PoolStatus,
RemoveLiquidityRequest, SwapRequest,
};
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct OraclePrice {
pub pool_id: Uuid,
pub reference_price: Decimal,
pub max_deviation: Decimal,
pub updated_at: DateTime<Utc>,
pub oracle_source: String,
}
impl OraclePrice {
pub fn is_price_within_bounds(&self, pool_price: Decimal) -> bool {
let lower_bound = self.reference_price * (dec!(1) - self.max_deviation);
let upper_bound = self.reference_price * (dec!(1) + self.max_deviation);
pool_price >= lower_bound && pool_price <= upper_bound
}
pub fn calculate_deviation(&self, pool_price: Decimal) -> Decimal {
if self.reference_price == dec!(0) {
return dec!(0);
}
((pool_price - self.reference_price) / self.reference_price).abs()
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct VolatilityTracker {
pub pool_id: Uuid,
pub recent_changes: Vec<Decimal>,
pub max_samples: usize,
pub current_volatility: Decimal,
pub updated_at: DateTime<Utc>,
}
impl VolatilityTracker {
pub fn new(pool_id: Uuid, max_samples: usize) -> Self {
Self {
pool_id,
recent_changes: Vec::new(),
max_samples,
current_volatility: dec!(0),
updated_at: Utc::now(),
}
}
pub fn add_price_change(&mut self, price_change_pct: Decimal) {
self.recent_changes.push(price_change_pct);
if self.recent_changes.len() > self.max_samples {
self.recent_changes.remove(0);
}
self.current_volatility = self.calculate_volatility();
self.updated_at = Utc::now();
}
fn calculate_volatility(&self) -> Decimal {
if self.recent_changes.is_empty() {
return dec!(0);
}
let n = Decimal::from(self.recent_changes.len());
let mean = self.recent_changes.iter().sum::<Decimal>() / n;
let variance = self
.recent_changes
.iter()
.map(|x| {
let diff = *x - mean;
diff * diff
})
.sum::<Decimal>()
/ n;
if let Some(var_f64) = variance.to_f64() {
if let Some(std_dev) = Decimal::from_f64_retain(var_f64.sqrt()) {
return std_dev;
}
}
dec!(0)
}
pub fn calculate_dynamic_fee(&self, base_fee: Decimal, max_fee: Decimal) -> Decimal {
if self.current_volatility == dec!(0) {
return base_fee;
}
let volatility_multiplier = dec!(1) + (self.current_volatility / dec!(0.1));
let dynamic_fee = base_fee * volatility_multiplier;
dynamic_fee.min(max_fee)
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct FlashSwapRequest {
pub pool_id: Uuid,
pub borrow_amount: Decimal,
pub borrow_token_a: bool,
pub expected_repayment: Decimal,
pub callback_data: Vec<u8>,
}
impl FlashSwapRequest {
pub fn validate(&self) -> Result<()> {
if self.borrow_amount <= dec!(0) {
return Err(CoreError::Validation(
"Borrow amount must be positive".to_string(),
));
}
if self.expected_repayment < self.borrow_amount {
return Err(CoreError::Validation(
"Repayment must be at least borrow amount".to_string(),
));
}
Ok(())
}
}
#[derive(Debug, Serialize)]
pub struct FlashSwapResult {
pub borrowed_amount: Decimal,
pub repaid_amount: Decimal,
pub fee_amount: Decimal,
pub profit: Decimal,
}
pub struct AmmExecutor;
impl AmmExecutor {
pub fn create_pool(request: CreatePoolRequest, _creator_id: Uuid) -> Result<LiquidityPool> {
request.validate().map_err(|e| {
CoreError::Validation(format!("Invalid pool creation request: {}", e.0))
})?;
let now = Utc::now();
let pool_id = Uuid::new_v4();
let fee_percentage = request.fee_percentage.unwrap_or(dec!(0.003));
let product = request.initial_reserve_a * request.initial_reserve_b;
let total_lp_tokens = if let Some(val) = product.to_f64() {
Decimal::from_f64_retain(val.sqrt())
.ok_or_else(|| CoreError::Validation("Invalid initial reserves".to_string()))?
} else {
return Err(CoreError::Validation(
"Invalid initial reserves".to_string(),
));
};
Ok(LiquidityPool {
pool_id,
token_a_id: request.token_a_id,
token_b_id: request.token_b_id,
reserve_a: request.initial_reserve_a,
reserve_b: request.initial_reserve_b,
total_lp_tokens,
fee_percentage,
status: PoolStatus::Active,
cumulative_volume_a: dec!(0),
cumulative_volume_b: dec!(0),
total_fees_a: dec!(0),
total_fees_b: dec!(0),
created_at: now,
updated_at: now,
})
}
pub fn create_initial_position(pool: &LiquidityPool, creator_id: Uuid) -> Result<LpPosition> {
let now = Utc::now();
Ok(LpPosition {
position_id: Uuid::new_v4(),
pool_id: pool.pool_id,
user_id: creator_id,
lp_tokens: pool.total_lp_tokens,
initial_reserve_a: pool.reserve_a,
initial_reserve_b: pool.reserve_b,
pool_share: dec!(1), created_at: now,
updated_at: now,
})
}
pub fn add_liquidity(
pool: &mut LiquidityPool,
request: AddLiquidityRequest,
_user_id: Uuid,
) -> Result<(Decimal, Decimal, Decimal)> {
request.validate().map_err(|e| {
CoreError::Validation(format!("Invalid add liquidity request: {}", e.0))
})?;
if pool.status != PoolStatus::Active {
return Err(CoreError::Validation("Pool is not active".to_string()));
}
let ratio = pool.reserve_b / pool.reserve_a;
let optimal_amount_b = request.amount_a * ratio;
let (actual_amount_a, actual_amount_b) = if optimal_amount_b <= request.amount_b {
(request.amount_a, optimal_amount_b)
} else {
let optimal_amount_a = request.amount_b / ratio;
(optimal_amount_a, request.amount_b)
};
let lp_tokens_a = (actual_amount_a * pool.total_lp_tokens) / pool.reserve_a;
let lp_tokens_b = (actual_amount_b * pool.total_lp_tokens) / pool.reserve_b;
let lp_tokens = lp_tokens_a.min(lp_tokens_b);
if let Some(min_lp) = request.min_lp_tokens {
if lp_tokens < min_lp {
return Err(CoreError::SlippageExceeded {
expected: min_lp,
actual: lp_tokens,
});
}
}
pool.reserve_a += actual_amount_a;
pool.reserve_b += actual_amount_b;
pool.total_lp_tokens += lp_tokens;
pool.updated_at = Utc::now();
Ok((lp_tokens, actual_amount_a, actual_amount_b))
}
pub fn update_or_create_position(
pool: &LiquidityPool,
user_id: Uuid,
lp_tokens_added: Decimal,
amount_a_added: Decimal,
amount_b_added: Decimal,
existing_position: Option<LpPosition>,
) -> Result<LpPosition> {
let now = Utc::now();
match existing_position {
Some(mut pos) => {
pos.lp_tokens += lp_tokens_added;
pos.pool_share = pos.lp_tokens / pool.total_lp_tokens;
pos.updated_at = now;
Ok(pos)
}
None => {
Ok(LpPosition {
position_id: Uuid::new_v4(),
pool_id: pool.pool_id,
user_id,
lp_tokens: lp_tokens_added,
initial_reserve_a: amount_a_added,
initial_reserve_b: amount_b_added,
pool_share: lp_tokens_added / pool.total_lp_tokens,
created_at: now,
updated_at: now,
})
}
}
}
pub fn remove_liquidity(
pool: &mut LiquidityPool,
request: RemoveLiquidityRequest,
position: &LpPosition,
) -> Result<(Decimal, Decimal)> {
request.validate().map_err(|e| {
CoreError::Validation(format!("Invalid remove liquidity request: {}", e.0))
})?;
if pool.status == PoolStatus::Closed {
return Err(CoreError::Validation("Pool is closed".to_string()));
}
if position.lp_tokens < request.lp_tokens {
return Err(CoreError::InsufficientBalance {
required: request.lp_tokens,
available: position.lp_tokens,
});
}
let share = request.lp_tokens / pool.total_lp_tokens;
let amount_a = pool.reserve_a * share;
let amount_b = pool.reserve_b * share;
if let Some(min_a) = request.min_amount_a {
if amount_a < min_a {
return Err(CoreError::SlippageExceeded {
expected: min_a,
actual: amount_a,
});
}
}
if let Some(min_b) = request.min_amount_b {
if amount_b < min_b {
return Err(CoreError::SlippageExceeded {
expected: min_b,
actual: amount_b,
});
}
}
pool.reserve_a -= amount_a;
pool.reserve_b -= amount_b;
pool.total_lp_tokens -= request.lp_tokens;
pool.updated_at = Utc::now();
Ok((amount_a, amount_b))
}
pub fn update_position_after_removal(
position: &mut LpPosition,
pool: &LiquidityPool,
lp_tokens_removed: Decimal,
) {
position.lp_tokens -= lp_tokens_removed;
position.pool_share = if pool.total_lp_tokens > dec!(0) {
position.lp_tokens / pool.total_lp_tokens
} else {
dec!(0)
};
position.updated_at = Utc::now();
}
pub fn swap(pool: &mut LiquidityPool, request: SwapRequest) -> Result<(Decimal, Decimal)> {
request
.validate()
.map_err(|e| CoreError::Validation(format!("Invalid swap request: {}", e.0)))?;
if pool.status != PoolStatus::Active {
return Err(CoreError::Validation("Pool is not active".to_string()));
}
let output_amount = pool.calculate_output(request.input_amount, request.input_is_a);
if output_amount == dec!(0) {
return Err(CoreError::Validation(
"Invalid swap: output is zero".to_string(),
));
}
if let Some(min_output) = request.min_output {
if output_amount < min_output {
return Err(CoreError::SlippageExceeded {
expected: min_output,
actual: output_amount,
});
}
}
if let Some(max_slippage) = request.max_slippage {
let price_impact =
pool.calculate_price_impact(request.input_amount, request.input_is_a);
if price_impact > max_slippage {
return Err(CoreError::SlippageExceeded {
expected: max_slippage,
actual: price_impact,
});
}
}
let fee_amount = request.input_amount * pool.fee_percentage;
if request.input_is_a {
pool.reserve_a += request.input_amount;
pool.reserve_b -= output_amount;
pool.cumulative_volume_a += request.input_amount;
pool.total_fees_a += fee_amount;
} else {
pool.reserve_b += request.input_amount;
pool.reserve_a -= output_amount;
pool.cumulative_volume_b += request.input_amount;
pool.total_fees_b += fee_amount;
}
pool.updated_at = Utc::now();
Ok((output_amount, fee_amount))
}
pub fn calculate_optimal_route(
pools: &[LiquidityPool],
input_token_id: Uuid,
output_token_id: Uuid,
input_amount: Decimal,
) -> Option<(Vec<Uuid>, Decimal)> {
let mut best_route: Option<(Vec<Uuid>, Decimal)> = None;
for pool in pools {
if pool.status != PoolStatus::Active {
continue;
}
if (pool.token_a_id == input_token_id && pool.token_b_id == output_token_id)
|| (pool.token_b_id == input_token_id && pool.token_a_id == output_token_id)
{
let input_is_a = pool.token_a_id == input_token_id;
let output = pool.calculate_output(input_amount, input_is_a);
best_route = Some((vec![pool.pool_id], output));
}
}
for pool1 in pools {
if pool1.status != PoolStatus::Active {
continue;
}
let (intermediate_token, input_is_a1) = if pool1.token_a_id == input_token_id {
(Some(pool1.token_b_id), true)
} else if pool1.token_b_id == input_token_id {
(Some(pool1.token_a_id), false)
} else {
(None, false)
};
if let Some(intermediate) = intermediate_token {
for pool2 in pools {
if pool2.status != PoolStatus::Active {
continue;
}
let input_is_a2 = if pool2.token_a_id == intermediate
&& pool2.token_b_id == output_token_id
{
Some(true)
} else if pool2.token_b_id == intermediate
&& pool2.token_a_id == output_token_id
{
Some(false)
} else {
None
};
if let Some(is_a2) = input_is_a2 {
let intermediate_output = pool1.calculate_output(input_amount, input_is_a1);
let final_output = pool2.calculate_output(intermediate_output, is_a2);
if let Some((_, current_best)) = best_route {
if final_output > current_best {
best_route =
Some((vec![pool1.pool_id, pool2.pool_id], final_output));
}
} else {
best_route = Some((vec![pool1.pool_id, pool2.pool_id], final_output));
}
}
}
}
}
best_route
}
pub fn execute_flash_swap(
pool: &mut LiquidityPool,
request: FlashSwapRequest,
) -> Result<FlashSwapResult> {
request.validate()?;
if pool.status != PoolStatus::Active {
return Err(CoreError::Validation("Pool is not active".to_string()));
}
let min_fee = request.borrow_amount * pool.fee_percentage;
let min_repayment = request.borrow_amount + min_fee;
if request.expected_repayment < min_repayment {
return Err(CoreError::Validation(format!(
"Repayment {} insufficient, minimum required: {}",
request.expected_repayment, min_repayment
)));
}
let fee_amount = request.expected_repayment - request.borrow_amount;
let profit = request.expected_repayment - min_repayment;
if request.borrow_token_a {
pool.total_fees_a += fee_amount;
} else {
pool.total_fees_b += fee_amount;
}
pool.updated_at = Utc::now();
Ok(FlashSwapResult {
borrowed_amount: request.borrow_amount,
repaid_amount: request.expected_repayment,
fee_amount,
profit,
})
}
pub fn update_dynamic_fee(
pool: &mut LiquidityPool,
volatility_tracker: &VolatilityTracker,
base_fee: Decimal,
max_fee: Decimal,
) {
let new_fee = volatility_tracker.calculate_dynamic_fee(base_fee, max_fee);
pool.fee_percentage = new_fee;
pool.updated_at = Utc::now();
}
pub fn validate_swap_with_oracle(
pool: &LiquidityPool,
oracle: &OraclePrice,
input_amount: Decimal,
input_is_a: bool,
) -> Result<()> {
let output_amount = pool.calculate_output(input_amount, input_is_a);
let (new_reserve_a, new_reserve_b) = if input_is_a {
(
pool.reserve_a + input_amount,
pool.reserve_b - output_amount,
)
} else {
(
pool.reserve_a - output_amount,
pool.reserve_b + input_amount,
)
};
let new_price = new_reserve_b / new_reserve_a;
if !oracle.is_price_within_bounds(new_price) {
let deviation = oracle.calculate_deviation(new_price);
return Err(CoreError::Validation(format!(
"Price deviation {} exceeds oracle bounds (max: {})",
deviation, oracle.max_deviation
)));
}
Ok(())
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_create_pool() {
let token_a = Uuid::new_v4();
let token_b = Uuid::new_v4();
let creator = Uuid::new_v4();
let request = CreatePoolRequest {
token_a_id: token_a,
token_b_id: token_b,
initial_reserve_a: dec!(1000),
initial_reserve_b: dec!(2000),
fee_percentage: Some(dec!(0.003)),
};
let pool = AmmExecutor::create_pool(request, creator).unwrap();
assert_eq!(pool.reserve_a, dec!(1000));
assert_eq!(pool.reserve_b, dec!(2000));
assert_eq!(pool.fee_percentage, dec!(0.003));
assert!(pool.total_lp_tokens > dec!(1414) && pool.total_lp_tokens < dec!(1415));
}
#[test]
fn test_swap_calculation() {
let token_a = Uuid::new_v4();
let token_b = Uuid::new_v4();
let creator = Uuid::new_v4();
let request = CreatePoolRequest {
token_a_id: token_a,
token_b_id: token_b,
initial_reserve_a: dec!(1000),
initial_reserve_b: dec!(2000),
fee_percentage: Some(dec!(0.003)),
};
let mut pool = AmmExecutor::create_pool(request, creator).unwrap();
let swap_request = SwapRequest {
pool_id: pool.pool_id,
input_amount: dec!(100),
input_is_a: true,
min_output: None,
max_slippage: None,
};
let (output, fee) = AmmExecutor::swap(&mut pool, swap_request).unwrap();
assert_eq!(fee, dec!(0.3));
assert!(output < dec!(200)); assert!(output > dec!(0));
assert_eq!(pool.reserve_a, dec!(1100));
assert!(pool.reserve_b < dec!(2000));
}
#[test]
fn test_price_impact() {
let pool = LiquidityPool {
pool_id: Uuid::new_v4(),
token_a_id: Uuid::new_v4(),
token_b_id: Uuid::new_v4(),
reserve_a: dec!(1000),
reserve_b: dec!(2000),
total_lp_tokens: dec!(1414.213562373095),
fee_percentage: dec!(0.003),
status: PoolStatus::Active,
cumulative_volume_a: dec!(0),
cumulative_volume_b: dec!(0),
total_fees_a: dec!(0),
total_fees_b: dec!(0),
created_at: Utc::now(),
updated_at: Utc::now(),
};
let small_impact = pool.calculate_price_impact(dec!(10), true);
assert!(small_impact < dec!(0.02));
let large_impact = pool.calculate_price_impact(dec!(500), true);
assert!(large_impact > dec!(0.1)); }
#[test]
fn test_add_remove_liquidity() {
let token_a = Uuid::new_v4();
let token_b = Uuid::new_v4();
let creator = Uuid::new_v4();
let request = CreatePoolRequest {
token_a_id: token_a,
token_b_id: token_b,
initial_reserve_a: dec!(1000),
initial_reserve_b: dec!(2000),
fee_percentage: Some(dec!(0.003)),
};
let mut pool = AmmExecutor::create_pool(request, creator).unwrap();
let initial_lp = pool.total_lp_tokens;
let add_request = AddLiquidityRequest {
pool_id: pool.pool_id,
amount_a: dec!(500),
amount_b: dec!(1000),
min_lp_tokens: None,
};
let user = Uuid::new_v4();
let (lp_tokens, actual_a, actual_b) =
AmmExecutor::add_liquidity(&mut pool, add_request, user).unwrap();
assert_eq!(actual_a, dec!(500));
assert_eq!(actual_b, dec!(1000));
assert_eq!(pool.reserve_a, dec!(1500));
assert_eq!(pool.reserve_b, dec!(3000));
let expected_lp = initial_lp * dec!(0.5);
assert!((lp_tokens - expected_lp).abs() < dec!(0.01));
let position = LpPosition {
position_id: Uuid::new_v4(),
pool_id: pool.pool_id,
user_id: user,
lp_tokens,
initial_reserve_a: actual_a,
initial_reserve_b: actual_b,
pool_share: lp_tokens / pool.total_lp_tokens,
created_at: Utc::now(),
updated_at: Utc::now(),
};
let remove_request = RemoveLiquidityRequest {
pool_id: pool.pool_id,
lp_tokens: lp_tokens / dec!(2), min_amount_a: None,
min_amount_b: None,
};
let (removed_a, removed_b) =
AmmExecutor::remove_liquidity(&mut pool, remove_request, &position).unwrap();
assert!((removed_a - dec!(250)).abs() < dec!(1));
assert!((removed_b - dec!(500)).abs() < dec!(1));
}
#[test]
fn test_volatility_tracker() {
let pool_id = Uuid::new_v4();
let mut tracker = VolatilityTracker::new(pool_id, 10);
tracker.add_price_change(dec!(0.05)); tracker.add_price_change(dec!(-0.03)); tracker.add_price_change(dec!(0.02));
assert!(tracker.current_volatility > dec!(0));
assert_eq!(tracker.recent_changes.len(), 3);
let base_fee = dec!(0.003);
let max_fee = dec!(0.01);
let dynamic_fee = tracker.calculate_dynamic_fee(base_fee, max_fee);
assert!(dynamic_fee >= base_fee);
assert!(dynamic_fee <= max_fee);
}
#[test]
fn test_oracle_price_bounds() {
let pool_id = Uuid::new_v4();
let oracle = OraclePrice {
pool_id,
reference_price: dec!(2), max_deviation: dec!(0.05), updated_at: Utc::now(),
oracle_source: "test_oracle".to_string(),
};
assert!(oracle.is_price_within_bounds(dec!(2.0)));
assert!(oracle.is_price_within_bounds(dec!(2.05)));
assert!(oracle.is_price_within_bounds(dec!(1.95)));
assert!(!oracle.is_price_within_bounds(dec!(2.15)));
assert!(!oracle.is_price_within_bounds(dec!(1.85)));
let deviation = oracle.calculate_deviation(dec!(2.1));
assert_eq!(deviation, dec!(0.05)); }
#[test]
fn test_flash_swap() {
let token_a = Uuid::new_v4();
let token_b = Uuid::new_v4();
let creator = Uuid::new_v4();
let request = CreatePoolRequest {
token_a_id: token_a,
token_b_id: token_b,
initial_reserve_a: dec!(10000),
initial_reserve_b: dec!(20000),
fee_percentage: Some(dec!(0.003)),
};
let mut pool = AmmExecutor::create_pool(request, creator).unwrap();
let flash_request = FlashSwapRequest {
pool_id: pool.pool_id,
borrow_amount: dec!(100),
borrow_token_a: true,
expected_repayment: dec!(100.5), callback_data: vec![],
};
let result = AmmExecutor::execute_flash_swap(&mut pool, flash_request).unwrap();
assert_eq!(result.borrowed_amount, dec!(100));
assert_eq!(result.repaid_amount, dec!(100.5));
assert_eq!(result.fee_amount, dec!(0.5));
assert!(result.profit >= dec!(0));
}
#[test]
fn test_enhanced_routing() {
let token_a = Uuid::new_v4();
let token_b = Uuid::new_v4();
let token_c = Uuid::new_v4();
let pool_ab = LiquidityPool {
pool_id: Uuid::new_v4(),
token_a_id: token_a,
token_b_id: token_b,
reserve_a: dec!(1000),
reserve_b: dec!(2000),
total_lp_tokens: dec!(1414.213562373095),
fee_percentage: dec!(0.003),
status: PoolStatus::Active,
cumulative_volume_a: dec!(0),
cumulative_volume_b: dec!(0),
total_fees_a: dec!(0),
total_fees_b: dec!(0),
created_at: Utc::now(),
updated_at: Utc::now(),
};
let pool_bc = LiquidityPool {
pool_id: Uuid::new_v4(),
token_a_id: token_b,
token_b_id: token_c,
reserve_a: dec!(2000),
reserve_b: dec!(4000),
total_lp_tokens: dec!(2828.427124746190),
fee_percentage: dec!(0.003),
status: PoolStatus::Active,
cumulative_volume_a: dec!(0),
cumulative_volume_b: dec!(0),
total_fees_a: dec!(0),
total_fees_b: dec!(0),
created_at: Utc::now(),
updated_at: Utc::now(),
};
let pools = vec![pool_ab.clone(), pool_bc.clone()];
let route = AmmExecutor::calculate_optimal_route(&pools, token_a, token_c, dec!(100));
assert!(route.is_some());
let (route_pools, output) = route.unwrap();
assert_eq!(route_pools.len(), 2);
assert_eq!(route_pools[0], pool_ab.pool_id);
assert_eq!(route_pools[1], pool_bc.pool_id);
assert!(output > dec!(0));
}
#[test]
fn test_dynamic_fee_update() {
let token_a = Uuid::new_v4();
let token_b = Uuid::new_v4();
let creator = Uuid::new_v4();
let request = CreatePoolRequest {
token_a_id: token_a,
token_b_id: token_b,
initial_reserve_a: dec!(1000),
initial_reserve_b: dec!(2000),
fee_percentage: Some(dec!(0.003)),
};
let mut pool = AmmExecutor::create_pool(request, creator).unwrap();
let initial_fee = pool.fee_percentage;
let mut tracker = VolatilityTracker::new(pool.pool_id, 10);
tracker.add_price_change(dec!(0.08)); tracker.add_price_change(dec!(-0.05)); tracker.add_price_change(dec!(0.06)); tracker.add_price_change(dec!(-0.04));
AmmExecutor::update_dynamic_fee(&mut pool, &tracker, dec!(0.003), dec!(0.01));
assert!(pool.fee_percentage > initial_fee);
assert!(pool.fee_percentage <= dec!(0.01));
}
}