use anchor_lang::prelude::*;
use crate::errors::ReflectErrorCodes;
use crate::reflect::{AutoCompound, deserialise_autocompound};
use crate::spl_mint::get_mint_supply;
use crate::drift::get_usdc_amount_drift_data;
use crate::ids;
pub fn compute_usdc_from_tokens(
token_amount: u64,
deposited_vault_value: u64,
effective_supply: u64,
) -> Result<u64> {
if effective_supply == 0 {
return Err(ReflectErrorCodes::MathError.into());
}
let numerator = (token_amount as u128) * (deposited_vault_value as u128);
let result = numerator / (effective_supply as u128);
if result > u64::MAX as u128 {
return Err(ReflectErrorCodes::MathError.into());
}
Ok(result as u64)
}
pub fn compute_tokens_from_usdc(
usdc_amount: u64,
deposited_vault_value: u64,
effective_supply: u64,
) -> Result<u64> {
if effective_supply == 0 || deposited_vault_value == 0 {
return Ok(usdc_amount);
}
let numerator = (usdc_amount as u128) * (effective_supply as u128);
let result = numerator / (deposited_vault_value as u128);
if result > u64::MAX as u128 {
return Err(ReflectErrorCodes::MathError.into());
}
Ok(result as u64)
}
pub fn exchange_rate_usdc_input(
data_usdc_controller: &[u8],
data_spot_market_usdc: &[u8],
data_user_account: &[u8],
data_usdc_plus_mint: &[u8],
usdc_amount: u64,
) -> Result<u64> {
let mut auto_compound: AutoCompound = deserialise_autocompound(data_usdc_controller)?;
let usdc_plus_supply = get_mint_supply(data_usdc_plus_mint)?;
let usdc_amount_drift = get_usdc_amount_drift_data(data_user_account, data_spot_market_usdc)?;
auto_compound.update_pool(usdc_amount_drift, &vec![10_000], usdc_plus_supply)?;
let deposited_vault_value: u64 = auto_compound.deposited_vault_value;
compute_tokens_from_usdc(usdc_amount, deposited_vault_value, usdc_plus_supply)
}
pub fn exchange_rate_usdc_accounts(
usdc_controller_account: &AccountInfo,
spot_market_usdc_account: &AccountInfo,
reflect_user_account: &AccountInfo,
usdc_plus_mint_account: &AccountInfo,
usdc_amount: u64,
) -> Result<u64> {
require!(usdc_controller_account.key == &ids::usdc_controller::ID, ReflectErrorCodes::InvalidUsdcControllerAccount);
require!(spot_market_usdc_account.key == &ids::usdc_spot_market::ID, ReflectErrorCodes::InvalidUsdcSpotMarketAccount);
require!(reflect_user_account.key == &ids::reflect_user_account_strategy_0::ID, ReflectErrorCodes::InvalidReflectUserAccount);
exchange_rate_usdc_input(
&usdc_controller_account.data.borrow(),
&spot_market_usdc_account.data.borrow(),
&reflect_user_account.data.borrow(),
&usdc_plus_mint_account.data.borrow(),
usdc_amount,
)
}
pub fn exchange_rate_receipt_input(
data_usdc_controller: &[u8],
data_spot_market_usdc: &[u8],
data_user_account: &[u8],
data_usdc_plus_mint: &[u8],
receipt_amount: u64,
) -> Result<u64> {
let mut auto_compound: AutoCompound = deserialise_autocompound(data_usdc_controller)?;
let usdc_plus_supply = get_mint_supply(data_usdc_plus_mint)?;
let usdc_amount_drift = get_usdc_amount_drift_data(data_user_account, data_spot_market_usdc)?;
auto_compound.update_pool(usdc_amount_drift, &vec![10_000], usdc_plus_supply)?;
let deposited_vault_value: u64 = auto_compound.deposited_vault_value;
compute_usdc_from_tokens(receipt_amount, deposited_vault_value, usdc_plus_supply)
}
pub fn exchange_rate_receipt_accounts(
usdc_controller_account: &AccountInfo,
spot_market_usdc_account: &AccountInfo,
reflect_user_account: &AccountInfo,
usdc_plus_mint_account: &AccountInfo,
receipt_amount: u64,
) -> Result<u64> {
require!(usdc_controller_account.key == &ids::usdc_controller::ID, ReflectErrorCodes::InvalidUsdcControllerAccount);
require!(spot_market_usdc_account.key == &ids::usdc_spot_market::ID, ReflectErrorCodes::InvalidUsdcSpotMarketAccount);
require!(reflect_user_account.key == &ids::reflect_user_account_strategy_0::ID, ReflectErrorCodes::InvalidReflectUserAccount);
exchange_rate_receipt_input(
&usdc_controller_account.data.borrow(),
&spot_market_usdc_account.data.borrow(),
&reflect_user_account.data.borrow(),
&usdc_plus_mint_account.data.borrow(),
receipt_amount,
)
}
#[cfg(test)]
mod tests {
use super::*;
use std::fs;
use std::path::PathBuf;
const LOCAL_CONTROLLER: &str = "usdc_controller_account.bin";
const LOCAL_SPOT_MARKET: &str = "spot_market_usdc.bin";
const LOCAL_USER_ACCOUNT: &str = "user_account.bin";
const LOCAL_USDC_PLUS_MINT: &str = "usdc_plus_mint.bin";
fn get_test_assets_dir() -> PathBuf {
let mut path = PathBuf::from(env!("CARGO_MANIFEST_DIR"));
path.push("./test_assets/local/");
path
}
#[test]
fn test_compute() {
let result = compute_usdc_from_tokens(1000, 1_100_000, 1_000_000).unwrap();
assert_eq!(result, 1100);
let result = compute_usdc_from_tokens(1000, 1_000_000, 0);
assert!(result.is_err());
let result = compute_usdc_from_tokens(u64::MAX, u64::MAX, 1);
assert!(result.is_err());
let result = compute_tokens_from_usdc(u64::MAX, 1, u64::MAX);
assert!(result.is_err());
}
#[test]
fn test_exchange_rate_functions() {
let assets_dir = get_test_assets_dir();
let controller_data = fs::read(assets_dir.join(LOCAL_CONTROLLER)).expect(&format!("Missing: {}", LOCAL_CONTROLLER));
let spot_data = fs::read(assets_dir.join(LOCAL_SPOT_MARKET)).expect(&format!("Missing: {}", LOCAL_SPOT_MARKET));
let user_data: Vec<u8> = fs::read(assets_dir.join(LOCAL_USER_ACCOUNT)).expect(&format!("Missing: {}", LOCAL_USER_ACCOUNT));
let usdc_plus_mint_data = fs::read(assets_dir.join(LOCAL_USDC_PLUS_MINT)).expect(&format!("Missing: {}", LOCAL_USDC_PLUS_MINT));
println!("\n=== Deposit Tests ===");
for amount in [1_000_000, 100_000_000, 1_000_000_000] {
let tokens = exchange_rate_usdc_input(
&controller_data,
&spot_data,
&user_data,
&usdc_plus_mint_data,
amount,
).expect("Failed to calculate deposit");
let rate = tokens as f64 / amount as f64;
println!(
"Deposit {:.2} USDC → {:.6} USDC+ (1 USDC = {:.6} USDC+)",
amount as f64 / 1_000_000.0,
tokens as f64 / 1_000_000.0,
rate
);
}
println!("\n=== Redemption Tests ===");
for amount in [1_000_000, 100_000_000, 1_000_000_000] {
let usdc = exchange_rate_receipt_input(
&controller_data,
&spot_data,
&user_data,
&usdc_plus_mint_data,
amount,
).expect("Failed to calculate redemption");
let rate = usdc as f64 / amount as f64;
println!(
"Redeem {:.2} USDC+ → {:.6} USDC (1 USDC+ = {:.6} USDC)",
amount as f64 / 1_000_000.0,
usdc as f64 / 1_000_000.0,
rate
);
}
}
#[test]
fn test_round_trip_conversion() {
let assets_dir = get_test_assets_dir();
let controller_data = fs::read(assets_dir.join(LOCAL_CONTROLLER)).expect(&format!("Missing: {}", LOCAL_CONTROLLER));
let spot_data = fs::read(assets_dir.join(LOCAL_SPOT_MARKET)).expect(&format!("Missing: {}", LOCAL_SPOT_MARKET));
let user_data = fs::read(assets_dir.join(LOCAL_USER_ACCOUNT)).expect(&format!("Missing: {}", LOCAL_USER_ACCOUNT));
let mint_data = fs::read(assets_dir.join(LOCAL_USDC_PLUS_MINT)).expect(&format!("Missing: {}", LOCAL_USDC_PLUS_MINT));
let initial_usdc = 1_000_000_000;
let tokens = exchange_rate_usdc_input(
&controller_data, &spot_data, &user_data, &mint_data, initial_usdc
).expect("Failed to calculate deposit");
let final_usdc = exchange_rate_receipt_input(
&controller_data, &spot_data, &user_data, &mint_data, tokens
).expect("Failed to calculate redemption");
println!(
"\nRound trip: {} USDC → {} USDC+ → {} USDC",
initial_usdc / 1_000_000,
tokens / 1_000_000,
final_usdc / 1_000_000
);
let difference = (final_usdc as i64 - initial_usdc as i64).abs();
assert!(difference <= 1, "Round trip loss exceeds tolerance: {} units", difference);
}
#[test]
fn test_yield_impact_on_exchange_rate() {
let assets_dir = get_test_assets_dir();
let controller_data = fs::read(assets_dir.join(LOCAL_CONTROLLER)).expect(&format!("Missing: {}", LOCAL_CONTROLLER));
let spot_data = fs::read(assets_dir.join(LOCAL_SPOT_MARKET)).expect(&format!("Missing: {}", LOCAL_SPOT_MARKET));
let user_data = fs::read(assets_dir.join(LOCAL_USER_ACCOUNT)).expect(&format!("Missing: {}", LOCAL_USER_ACCOUNT));
let mint_data = vec![0u8; 82];
let auto_compound = deserialise_autocompound(&controller_data).expect("Failed to deserialize autocompound");
let usdc_plus_supply = get_mint_supply(&mint_data).unwrap_or(0);
let stored_rate = if usdc_plus_supply > 0 {
auto_compound.deposited_vault_value as f64 / usdc_plus_supply as f64
} else {
1.0
};
let test_amount = 1_000_000_000;
let tokens = exchange_rate_usdc_input(
&controller_data, &spot_data, &user_data, &mint_data, test_amount
).expect("Failed to calculate with yield");
let updated_rate = test_amount as f64 / tokens as f64;
println!("\nStored rate (no pending yield): 1 USDC+ = {:.6} USDC", stored_rate);
println!("Updated rate (with pending yield): 1 USDC+ = {:.6} USDC", updated_rate);
println!("Yield impact: {:.4}%", ((updated_rate / stored_rate) - 1.0) * 100.0);
}
#[test]
fn test_error_handling() {
let bad_controller = vec![0u8; 100]; let valid_spot = vec![0u8; 10000];
let valid_user = vec![0u8; 10000];
let mint_data = vec![0u8; 82];
assert!(exchange_rate_usdc_input(&bad_controller, &valid_spot, &valid_user, &mint_data, 1_000_000).is_err());
assert!(exchange_rate_receipt_input(&bad_controller, &valid_spot, &valid_user, &mint_data, 1_000_000).is_err());
}
#[test]
fn test_exchange_rate_progression() {
let assets_dir = get_test_assets_dir();
let controller_data = fs::read(assets_dir.join(LOCAL_CONTROLLER)).expect(&format!("Missing: {}", LOCAL_CONTROLLER));
let spot_data = fs::read(assets_dir.join(LOCAL_SPOT_MARKET)).expect(&format!("Missing: {}", LOCAL_SPOT_MARKET));
let user_data = fs::read(assets_dir.join(LOCAL_USER_ACCOUNT)).expect(&format!("Missing: {}", LOCAL_USER_ACCOUNT));
let usdc_plus_mint_data = fs::read(assets_dir.join(LOCAL_USDC_PLUS_MINT)).expect(&format!("Missing: {}", LOCAL_USDC_PLUS_MINT));
let mut auto_compound = deserialise_autocompound(&controller_data).expect("Failed to deserialize");
let initial_drift_value = get_usdc_amount_drift_data(&user_data, &spot_data).expect("Failed to get drift value");
let usdc_plus_supply = get_mint_supply(&usdc_plus_mint_data).unwrap_or(0);
println!("\n=== Exchange Rate Progression Test ===");
println!("Initial drift value: {:.2}", initial_drift_value as f64 / 1_000_000.0);
let test_usdc = 1_000_000_000; let test_receipt = 1_000_000_000;
let mut last_deposit_rate = 0.0;
let mut last_redeem_rate = 0.0;
for round in 0..5 {
println!("\n--- Round {} ---", round);
let simulated_yield = (initial_drift_value as f64 * 0.02 * (round + 1) as f64) as u64;
let current_drift_value = initial_drift_value + simulated_yield;
println!(
"Simulated drift value: {:.2} (yield: {:.2})",
current_drift_value as f64 / 1_000_000.0,
simulated_yield as f64 / 1_000_000.0
);
let token_supply = usdc_plus_supply;
auto_compound.update_pool(current_drift_value, &[10_000], token_supply).unwrap();
let deposit_tokens = compute_tokens_from_usdc(
test_usdc,
auto_compound.deposited_vault_value,
token_supply,
).unwrap();
let redeem_usdc = compute_usdc_from_tokens(
test_receipt,
auto_compound.deposited_vault_value,
token_supply,
).unwrap();
let deposit_rate = deposit_tokens as f64 / test_usdc as f64;
let redeem_rate = redeem_usdc as f64 / test_receipt as f64;
println!("After capture:");
println!(" 1 USDC = {:.6} USDC+ (deposit rate)", deposit_rate);
println!(" 1 USDC+ = {:.6} USDC (redeem rate)", redeem_rate);
println!(" deposited_vault_value (V): {}", auto_compound.deposited_vault_value);
println!(" token_supply (S from mint): {}", token_supply);
if round > 0 {
assert!(
deposit_rate <= last_deposit_rate + 0.000001,
"Deposit rate improved after capture: {} > {} at round {}",
deposit_rate, last_deposit_rate, round
);
assert!(
redeem_rate >= last_redeem_rate - 0.000001,
"Redeem rate worsened after capture: {} < {} at round {}",
redeem_rate, last_redeem_rate, round
);
}
last_deposit_rate = deposit_rate;
last_redeem_rate = redeem_rate;
for amount in [1_000_000, 100_000_000, 10_000_000_000] {
let tokens = compute_tokens_from_usdc(
amount,
auto_compound.deposited_vault_value,
token_supply,
).unwrap();
let rate = tokens as f64 / amount as f64;
assert!(
(rate - deposit_rate).abs() < 0.000001,
"Rate inconsistent for different amounts"
);
}
}
println!("\n=== Summary ===");
println!("Deposit rate progression (USDC+ per USDC should decrease):");
println!(" Started at: ~1.0");
println!(" Ended at: {:.6}", last_deposit_rate);
println!("Redeem rate progression (USDC per USDC+ should increase):");
println!(" Started at: ~1.0");
println!(" Ended at: {:.6}", last_redeem_rate);
}
}