Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
99 changes: 85 additions & 14 deletions engine-core/src/protocol_fee.rs
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,8 @@ use soroban_sdk::{

const KEY_FEE_BPS: Symbol = symbol_short!("FEE_BPS");
const KEY_FEE_RECIPIENT: Symbol = symbol_short!("FEE_RCP");
/// Shared with control_plane — admin that may change fee parameters.
const KEY_ADMIN: Symbol = symbol_short!("ADMIN");
const MAX_BPS: u32 = 10_000;
const ZERO_ADDRESS: &str = "GAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAWHF";

Expand All @@ -19,6 +21,8 @@ pub enum FeeError {
InvalidRecipient = 2,
FeeCalculationOverflow = 3,
InvalidAmount = 4,
Unauthorized = 5,
NotInitialized = 6,
}

pub fn init(env: &Env, fee_bps: u32, recipient: &Address) {
Expand Down Expand Up @@ -72,16 +76,36 @@ pub fn get_fee_config(env: &Env) -> (u32, Option<Address>) {
(fee_bps, recipient)
}

pub fn set_fee_bps(env: &Env, fee_bps: u32) {
/// Update fee rate. Caller must authorize and must be the stored admin.
pub fn set_fee_bps(env: &Env, caller: &Address, fee_bps: u32) {
require_admin(env, caller);
validate_bps(env, fee_bps);
env.storage().instance().set(&KEY_FEE_BPS, &fee_bps);
}

pub fn set_fee_recipient(env: &Env, recipient: &Address) {
/// Update fee recipient. Caller must authorize and must be the stored admin.
pub fn set_fee_recipient(env: &Env, caller: &Address, recipient: &Address) {
require_admin(env, caller);
validate_address(env, recipient);
env.storage().instance().set(&KEY_FEE_RECIPIENT, recipient);
}

/// Mirror of control_plane::require_admin — auth + role check inside the module
/// so callers cannot forget it when wiring #[contractimpl] entrypoints.
fn require_admin(env: &Env, caller: &Address) {
caller.require_auth();

let admin: Address = env
.storage()
.instance()
.get(&KEY_ADMIN)
.unwrap_or_else(|| panic_with_error!(env, FeeError::NotInitialized));

if caller != &admin {
panic_with_error!(env, FeeError::Unauthorized);
}
}

fn validate_bps(env: &Env, fee_bps: u32) {
if fee_bps > MAX_BPS {
panic_with_error!(env, FeeError::InvalidBasisPoints);
Expand All @@ -105,25 +129,33 @@ fn fee_as_u64(env: &Env, fee: i128) -> u64 {
#[cfg(test)]
mod tests {
use super::*;
use soroban_sdk::{contract, contractimpl, testutils::Address as _, Env};
use soroban_sdk::{
contract, contractimpl,
testutils::{Address as _, MockAuth, MockAuthInvoke},
Env, IntoVal,
};

#[contract]
pub struct TestContract;

#[contractimpl]
impl TestContract {}

fn setup(env: &Env, fee_bps: u32) -> (Address, Address) {
fn setup(env: &Env, fee_bps: u32) -> (Address, Address, Address) {
let contract_id = env.register_contract(None, TestContract);
let recipient = Address::generate(env);
env.as_contract(&contract_id, || init(env, fee_bps, &recipient));
(contract_id, recipient)
let admin = Address::generate(env);
env.as_contract(&contract_id, || {
init(env, fee_bps, &recipient);
env.storage().instance().set(&KEY_ADMIN, &admin);
});
(contract_id, recipient, admin)
}

#[test]
fn fee_calculation_standard() {
let env = Env::default();
let (contract_id, _) = setup(&env, 500);
let (contract_id, _, _) = setup(&env, 500);
env.as_contract(&contract_id, || {
let (fee, net) = calculate_fee(&env, 1000);
assert_eq!(fee, 50);
Expand All @@ -134,7 +166,7 @@ mod tests {
#[test]
fn partial_bps_fee_rounds_down() {
let env = Env::default();
let (contract_id, _) = setup(&env, 333);
let (contract_id, _, _) = setup(&env, 333);
env.as_contract(&contract_id, || {
let (fee, net) = calculate_fee(&env, 100);
assert_eq!(fee, 3);
Expand All @@ -145,7 +177,7 @@ mod tests {
#[test]
fn fee_at_full_bps_is_full_amount() {
let env = Env::default();
let (contract_id, _) = setup(&env, MAX_BPS);
let (contract_id, _, _) = setup(&env, MAX_BPS);
env.as_contract(&contract_id, || {
let (fee, net) = calculate_fee(&env, 1000);
assert_eq!(fee, 1000);
Expand All @@ -156,7 +188,7 @@ mod tests {
#[test]
fn get_fee_config_returns_stored_values() {
let env = Env::default();
let (contract_id, recipient) = setup(&env, 250);
let (contract_id, recipient, _) = setup(&env, 250);
env.as_contract(&contract_id, || {
let (bps, rec) = get_fee_config(&env);
assert_eq!(bps, 250);
Expand All @@ -165,16 +197,55 @@ mod tests {
}

#[test]
fn set_fee_bps_updates_rate() {
fn set_fee_bps_updates_rate_when_admin() {
let env = Env::default();
let (contract_id, _) = setup(&env, 100);
env.mock_all_auths();
let (contract_id, _, admin) = setup(&env, 100);
env.as_contract(&contract_id, || {
set_fee_bps(&env, 750);
set_fee_bps(&env, &admin, 750);
let (bps, _) = get_fee_config(&env);
assert_eq!(bps, 750);
});
}

#[test]
fn set_fee_recipient_updates_when_admin() {
let env = Env::default();
env.mock_all_auths();
let (contract_id, _, admin) = setup(&env, 100);
let new_recipient = Address::generate(&env);
env.as_contract(&contract_id, || {
set_fee_recipient(&env, &admin, &new_recipient);
let (_, rec) = get_fee_config(&env);
assert_eq!(rec.unwrap(), new_recipient);
});
}

#[test]
#[should_panic]
fn set_fee_bps_rejects_non_admin() {
let env = Env::default();
env.mock_all_auths();
let (contract_id, _, _) = setup(&env, 100);
let stranger = Address::generate(&env);
env.as_contract(&contract_id, || {
set_fee_bps(&env, &stranger, 750);
});
}

#[test]
#[should_panic]
fn set_fee_recipient_rejects_non_admin() {
let env = Env::default();
env.mock_all_auths();
let (contract_id, _, _) = setup(&env, 100);
let stranger = Address::generate(&env);
let new_recipient = Address::generate(&env);
env.as_contract(&contract_id, || {
set_fee_recipient(&env, &stranger, &new_recipient);
});
}

#[test]
#[should_panic]
fn rejects_bps_over_max() {
Expand All @@ -183,4 +254,4 @@ mod tests {
let recipient = Address::generate(&env);
env.as_contract(&contract_id, || init(&env, MAX_BPS + 1, &recipient));
}
}
}
Loading