diff --git a/engine-core/src/protocol_fee.rs b/engine-core/src/protocol_fee.rs index ac18b1c..e434141 100644 --- a/engine-core/src/protocol_fee.rs +++ b/engine-core/src/protocol_fee.rs @@ -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"; @@ -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) { @@ -72,16 +76,36 @@ pub fn get_fee_config(env: &Env) -> (u32, Option
) { (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); @@ -105,7 +129,11 @@ 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; @@ -113,17 +141,21 @@ mod tests { #[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); @@ -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); @@ -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); @@ -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); @@ -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() { @@ -183,4 +254,4 @@ mod tests { let recipient = Address::generate(&env); env.as_contract(&contract_id, || init(&env, MAX_BPS + 1, &recipient)); } -} + }