diff --git a/risk_score/Cargo.toml b/risk_score/Cargo.toml
index 13824bca..d1732921 100644
--- a/risk_score/Cargo.toml
+++ b/risk_score/Cargo.toml
@@ -9,6 +9,9 @@ crate-type = ["cdylib"]
[dependencies]
soroban-sdk = "22.0.8"
+[dev-dependencies]
+soroban-sdk = { version = "22.0.8", features = ["testutils"] }
+
[profile.release]
opt-level = "z"
overflow-checks = true
diff --git a/risk_score/src/lib.rs b/risk_score/src/lib.rs
index c5f14569..95eeabce 100644
--- a/risk_score/src/lib.rs
+++ b/risk_score/src/lib.rs
@@ -6,6 +6,12 @@ use soroban_sdk::{contract, contractimpl, contracttype, Address, Env, Map, Symbo
#[contract]
pub struct RiskTierContract;
+/// Storage keys
+#[contracttype]
+pub enum DataKey {
+ Admin,
+}
+
/// Risk and tier data structure - using contracttype for Soroban serialization
#[contracttype]
#[derive(Clone, Debug, Eq, PartialEq)]
@@ -18,9 +24,46 @@ pub struct RiskTierData {
#[contractimpl]
impl RiskTierContract {
- /// Set risk score with tier classification and timestamp
- /// Following Soroban persistent storage best practices with tuple keys
- pub fn set_risk_tier(env: Env, user: Address, score: u32, tier: Symbol, chosen_tier: Symbol) {
+ /// Initialize the contract with a trusted admin address.
+ /// Must be called once before set_risk_tier can be used.
+ /// Panics if already initialized.
+ pub fn initialize(env: Env, admin: Address) {
+ assert!(
+ !env.storage().instance().has(&DataKey::Admin),
+ "already initialized"
+ );
+ admin.require_auth();
+ env.storage().instance().set(&DataKey::Admin, &admin);
+ }
+
+ /// Returns the current admin address, or None if not yet initialized.
+ pub fn get_admin(env: Env) -> Option
{
+ env.storage().instance().get(&DataKey::Admin)
+ }
+
+ /// Set risk score with tier classification and timestamp.
+ ///
+ /// Access control (fixes issue #50):
+ /// - The contract must be initialized first.
+ /// - Only the admin OR the user themselves can set/update the score.
+ /// - Unauthorized callers are rejected by require_auth().
+ pub fn set_risk_tier(env: Env, caller: Address, user: Address, score: u32, tier: Symbol, chosen_tier: Symbol) {
+ // Enforce initialization
+ let admin: Address = env
+ .storage()
+ .instance()
+ .get(&DataKey::Admin)
+ .expect("contract not initialized — call initialize() first");
+
+ // Caller must be admin OR the user themselves
+ if caller == admin {
+ admin.require_auth();
+ } else if caller == user {
+ user.require_auth();
+ } else {
+ panic!("unauthorized: caller must be admin or the user themselves");
+ }
+
// Validate inputs
assert!(score <= 100, "Score must be 0-100");
assert!(
@@ -100,7 +143,7 @@ impl RiskTierContract {
env.storage()
.persistent()
.get(&chosen_key)
- .unwrap_or(Symbol::new(&env, "TIER_3")) // Default to most conservative
+ .unwrap_or(Symbol::new(&env, "TIER_3"))
}
/// Get all users in a specific tier
@@ -112,8 +155,11 @@ impl RiskTierContract {
.unwrap_or(Vec::new(&env))
}
- /// Update user's chosen tier (risk-based validation)
+ /// Update user's chosen tier (risk-based validation).
+ /// Only the user themselves can update their chosen tier.
pub fn update_chosen_tier(env: Env, user: Address, new_chosen_tier: Symbol) {
+ user.require_auth();
+
let tuple_key = (user.clone(), Symbol::new(&env, "risk_tier"));
if let Some(mut risk_data) = env
@@ -122,7 +168,6 @@ impl RiskTierContract {
.get::<_, RiskTierData>(&tuple_key)
{
// Risk-based tier access control
- // High risk users (>70) can only choose TIER_3 for "opportunity" access
if risk_data.score > 70 {
assert!(
new_chosen_tier == Symbol::new(&env, "TIER_3"),
@@ -131,17 +176,13 @@ impl RiskTierContract {
}
risk_data.chosen_tier = new_chosen_tier.clone();
- risk_data.timestamp = env.ledger().timestamp(); // Update timestamp
+ risk_data.timestamp = env.ledger().timestamp();
env.storage().persistent().set(&tuple_key, &risk_data);
- // Update chosen tier cache
let chosen_key = (user.clone(), Symbol::new(&env, "chosen_tier"));
- env.storage()
- .persistent()
- .set(&chosen_key, &new_chosen_tier);
+ env.storage().persistent().set(&chosen_key, &new_chosen_tier);
- // Emit Event for Indexers
env.events()
.publish((Symbol::new(&env, "tier_updated"), user), new_chosen_tier);
}
@@ -150,26 +191,21 @@ impl RiskTierContract {
/// Get tier statistics
pub fn get_tier_stats(env: Env) -> Map {
let mut stats = Map::new(&env);
-
let tiers = [
Symbol::new(&env, "TIER_1"),
Symbol::new(&env, "TIER_2"),
Symbol::new(&env, "TIER_3"),
];
-
for tier in tiers {
let tier_users = Self::get_tier_users(env.clone(), tier.clone());
stats.set(tier, tier_users.len());
}
-
stats
}
/// Check if user can access specific tier based on risk score
- /// Following Goldfinch/Maple risk-liquidity mapping methodology
pub fn can_access_tier(env: Env, user: Address, target_tier: Symbol) -> bool {
let tuple_key = (user, Symbol::new(&env, "risk_tier"));
-
if let Some(risk_data) = env
.storage()
.persistent()
@@ -178,15 +214,14 @@ impl RiskTierContract {
let tier_1 = Symbol::new(&env, "TIER_1");
let tier_2 = Symbol::new(&env, "TIER_2");
let tier_3 = Symbol::new(&env, "TIER_3");
-
match target_tier {
- t if t == tier_1 => risk_data.score <= 30, // Low risk only
- t if t == tier_2 => risk_data.score <= 70, // Low to medium risk
- t if t == tier_3 => true, // All users (with opportunity badge for high risk)
+ t if t == tier_1 => risk_data.score <= 30,
+ t if t == tier_2 => risk_data.score <= 70,
+ t if t == tier_3 => true,
_ => false,
}
} else {
- false // No risk data means no access
+ false
}
}
}
@@ -196,202 +231,196 @@ mod tests {
use super::*;
use soroban_sdk::{testutils::Address as _, Env};
- #[test]
- fn test_set_and_get_risk_tier() {
+ fn setup() -> (Env, soroban_sdk::Address, soroban_sdk::Address) {
let env = Env::default();
+ env.mock_all_auths();
let contract_id = env.register_contract(None, RiskTierContract);
let client = RiskTierContractClient::new(&env, &contract_id);
+ let admin = Address::generate(&env);
+ client.initialize(&admin);
+ (env, contract_id, admin)
+ }
- let user = Address::generate(&env);
- let tier_1 = Symbol::new(&env, "TIER_1");
-
- client.set_risk_tier(&user, &25, &tier_1, &tier_1);
-
- let risk_data = client.get_risk_tier(&user).unwrap();
- assert_eq!(risk_data.score, 25);
- assert_eq!(risk_data.tier, tier_1);
- assert_eq!(risk_data.chosen_tier, tier_1);
+ // ── initialize ───────────────────────────────────────────────────────────
+
+ #[test]
+ fn test_initialize_sets_admin() {
+ let (env, contract_id, admin) = setup();
+ let client = RiskTierContractClient::new(&env, &contract_id);
+ assert_eq!(client.get_admin(), Some(admin));
}
#[test]
- fn test_score_validation_upper_bound() {
- let env = Env::default();
- let contract_id = env.register_contract(None, RiskTierContract);
+ #[should_panic(expected = "already initialized")]
+ fn test_initialize_twice_panics() {
+ let (env, contract_id, admin) = setup();
let client = RiskTierContractClient::new(&env, &contract_id);
+ let admin2 = Address::generate(&env);
+ client.initialize(&admin2);
+ }
+ // ── access control ───────────────────────────────────────────────────────
+
+ #[test]
+ fn test_admin_can_set_risk_tier_for_any_user() {
+ let (env, contract_id, admin) = setup();
+ let client = RiskTierContractClient::new(&env, &contract_id);
let user = Address::generate(&env);
- let tier_3 = Symbol::new(&env, "TIER_3");
-
- client.set_risk_tier(&user, &100, &tier_3, &tier_3);
-
- let score = client.get_score(&user);
- assert_eq!(score, 100);
+ let tier_1 = Symbol::new(&env, "TIER_1");
+ client.set_risk_tier(&admin, &user, &25, &tier_1, &tier_1);
+ let data = client.get_risk_tier(&user).unwrap();
+ assert_eq!(data.score, 25);
}
#[test]
- #[should_panic(expected = "Score must be 0-100")]
- fn test_score_validation_exceeds_limit() {
- let env = Env::default();
- let contract_id = env.register_contract(None, RiskTierContract);
+ fn test_user_can_set_own_risk_tier() {
+ let (env, contract_id, _admin) = setup();
let client = RiskTierContractClient::new(&env, &contract_id);
-
let user = Address::generate(&env);
- let tier_3 = Symbol::new(&env, "TIER_3");
-
- client.set_risk_tier(&user, &101, &tier_3, &tier_3);
+ let tier_2 = Symbol::new(&env, "TIER_2");
+ client.set_risk_tier(&user, &user, &50, &tier_2, &tier_2);
+ assert_eq!(client.get_score(&user), 50);
}
#[test]
- #[should_panic(expected = "Invalid tier")]
- fn test_invalid_tier_validation() {
- let env = Env::default();
- let contract_id = env.register_contract(None, RiskTierContract);
+ #[should_panic(expected = "unauthorized")]
+ fn test_third_party_cannot_set_risk_tier() {
+ let (env, contract_id, _admin) = setup();
let client = RiskTierContractClient::new(&env, &contract_id);
-
let user = Address::generate(&env);
- let invalid_tier = Symbol::new(&env, "TIER_4");
-
- client.set_risk_tier(&user, &50, &invalid_tier, &invalid_tier);
+ let attacker = Address::generate(&env);
+ let tier_1 = Symbol::new(&env, "TIER_1");
+ // attacker tries to set user's score
+ client.set_risk_tier(&attacker, &user, &5, &tier_1, &tier_1);
}
#[test]
- fn test_tier_access_tier1_low_risk() {
+ #[should_panic(expected = "contract not initialized")]
+ fn test_set_risk_tier_requires_initialization() {
let env = Env::default();
+ env.mock_all_auths();
let contract_id = env.register_contract(None, RiskTierContract);
let client = RiskTierContractClient::new(&env, &contract_id);
-
let user = Address::generate(&env);
let tier_1 = Symbol::new(&env, "TIER_1");
-
- client.set_risk_tier(&user, &25, &tier_1, &tier_1);
-
- assert!(client.can_access_tier(&user, &tier_1));
+ client.set_risk_tier(&user, &user, &25, &tier_1, &tier_1);
}
+ // ── existing tests (adapted for new signature) ───────────────────────────
+
#[test]
- fn test_tier_access_tier1_boundary() {
- let env = Env::default();
- let contract_id = env.register_contract(None, RiskTierContract);
+ fn test_set_and_get_risk_tier() {
+ let (env, contract_id, admin) = setup();
let client = RiskTierContractClient::new(&env, &contract_id);
-
let user = Address::generate(&env);
let tier_1 = Symbol::new(&env, "TIER_1");
-
- client.set_risk_tier(&user, &30, &tier_1, &tier_1);
-
- assert!(client.can_access_tier(&user, &tier_1));
+ client.set_risk_tier(&admin, &user, &25, &tier_1, &tier_1);
+ let risk_data = client.get_risk_tier(&user).unwrap();
+ assert_eq!(risk_data.score, 25);
+ assert_eq!(risk_data.tier, tier_1);
+ assert_eq!(risk_data.chosen_tier, tier_1);
}
#[test]
- fn test_tier_access_tier1_denied() {
- let env = Env::default();
- let contract_id = env.register_contract(None, RiskTierContract);
+ fn test_score_validation_upper_bound() {
+ let (env, contract_id, admin) = setup();
let client = RiskTierContractClient::new(&env, &contract_id);
-
let user = Address::generate(&env);
- let tier_2 = Symbol::new(&env, "TIER_2");
- let tier_1 = Symbol::new(&env, "TIER_1");
-
- client.set_risk_tier(&user, &50, &tier_2, &tier_2);
-
- assert!(!client.can_access_tier(&user, &tier_1));
+ let tier_3 = Symbol::new(&env, "TIER_3");
+ client.set_risk_tier(&admin, &user, &100, &tier_3, &tier_3);
+ assert_eq!(client.get_score(&user), 100);
}
#[test]
- fn test_tier_access_tier2_medium_risk() {
- let env = Env::default();
- let contract_id = env.register_contract(None, RiskTierContract);
+ #[should_panic(expected = "Score must be 0-100")]
+ fn test_score_validation_exceeds_limit() {
+ let (env, contract_id, admin) = setup();
let client = RiskTierContractClient::new(&env, &contract_id);
+ let user = Address::generate(&env);
+ let tier_3 = Symbol::new(&env, "TIER_3");
+ client.set_risk_tier(&admin, &user, &101, &tier_3, &tier_3);
+ }
+ #[test]
+ #[should_panic(expected = "Invalid tier")]
+ fn test_invalid_tier_validation() {
+ let (env, contract_id, admin) = setup();
+ let client = RiskTierContractClient::new(&env, &contract_id);
let user = Address::generate(&env);
- let tier_2 = Symbol::new(&env, "TIER_2");
-
- client.set_risk_tier(&user, &50, &tier_2, &tier_2);
-
- assert!(client.can_access_tier(&user, &tier_2));
+ let invalid_tier = Symbol::new(&env, "TIER_4");
+ client.set_risk_tier(&admin, &user, &50, &invalid_tier, &invalid_tier);
}
#[test]
- fn test_tier_access_tier3_always_accessible() {
- let env = Env::default();
- let contract_id = env.register_contract(None, RiskTierContract);
+ fn test_tier_access_tier1_low_risk() {
+ let (env, contract_id, admin) = setup();
let client = RiskTierContractClient::new(&env, &contract_id);
+ let user = Address::generate(&env);
+ let tier_1 = Symbol::new(&env, "TIER_1");
+ client.set_risk_tier(&admin, &user, &25, &tier_1, &tier_1);
+ assert!(client.can_access_tier(&user, &tier_1));
+ }
+ #[test]
+ fn test_tier_access_tier1_denied() {
+ let (env, contract_id, admin) = setup();
+ let client = RiskTierContractClient::new(&env, &contract_id);
let user = Address::generate(&env);
- let tier_3 = Symbol::new(&env, "TIER_3");
-
- client.set_risk_tier(&user, &85, &tier_3, &tier_3);
-
- assert!(client.can_access_tier(&user, &tier_3));
+ let tier_2 = Symbol::new(&env, "TIER_2");
+ let tier_1 = Symbol::new(&env, "TIER_1");
+ client.set_risk_tier(&admin, &user, &50, &tier_2, &tier_2);
+ assert!(!client.can_access_tier(&user, &tier_1));
}
#[test]
fn test_get_tier_users() {
- let env = Env::default();
- let contract_id = env.register_contract(None, RiskTierContract);
+ let (env, contract_id, admin) = setup();
let client = RiskTierContractClient::new(&env, &contract_id);
-
let user1 = Address::generate(&env);
let user2 = Address::generate(&env);
let tier_1 = Symbol::new(&env, "TIER_1");
-
- client.set_risk_tier(&user1, &20, &tier_1, &tier_1);
- client.set_risk_tier(&user2, &25, &tier_1, &tier_1);
-
- let tier_users = client.get_tier_users(&tier_1);
- assert_eq!(tier_users.len(), 2);
+ client.set_risk_tier(&admin, &user1, &20, &tier_1, &tier_1);
+ client.set_risk_tier(&admin, &user2, &25, &tier_1, &tier_1);
+ assert_eq!(client.get_tier_users(&tier_1).len(), 2);
}
#[test]
fn test_update_chosen_tier_valid() {
- let env = Env::default();
- let contract_id = env.register_contract(None, RiskTierContract);
+ let (env, contract_id, admin) = setup();
let client = RiskTierContractClient::new(&env, &contract_id);
-
let user = Address::generate(&env);
let tier_1 = Symbol::new(&env, "TIER_1");
let tier_2 = Symbol::new(&env, "TIER_2");
-
- client.set_risk_tier(&user, &25, &tier_1, &tier_1);
+ client.set_risk_tier(&admin, &user, &25, &tier_1, &tier_1);
client.update_chosen_tier(&user, &tier_2);
-
- let chosen = client.get_chosen_tier(&user);
- assert_eq!(chosen, tier_2);
+ assert_eq!(client.get_chosen_tier(&user), tier_2);
}
#[test]
#[should_panic(expected = "High risk users can only access TIER_3")]
fn test_update_chosen_tier_high_risk_restriction() {
- let env = Env::default();
- let contract_id = env.register_contract(None, RiskTierContract);
+ let (env, contract_id, admin) = setup();
let client = RiskTierContractClient::new(&env, &contract_id);
-
let user = Address::generate(&env);
let tier_3 = Symbol::new(&env, "TIER_3");
let tier_1 = Symbol::new(&env, "TIER_1");
-
- client.set_risk_tier(&user, &85, &tier_3, &tier_3);
+ client.set_risk_tier(&admin, &user, &85, &tier_3, &tier_3);
client.update_chosen_tier(&user, &tier_1);
}
#[test]
fn test_get_tier_stats() {
- let env = Env::default();
- let contract_id = env.register_contract(None, RiskTierContract);
+ let (env, contract_id, admin) = setup();
let client = RiskTierContractClient::new(&env, &contract_id);
-
let user1 = Address::generate(&env);
let user2 = Address::generate(&env);
let user3 = Address::generate(&env);
-
let tier_1 = Symbol::new(&env, "TIER_1");
let tier_2 = Symbol::new(&env, "TIER_2");
let tier_3 = Symbol::new(&env, "TIER_3");
-
- client.set_risk_tier(&user1, &20, &tier_1, &tier_1);
- client.set_risk_tier(&user2, &50, &tier_2, &tier_2);
- client.set_risk_tier(&user3, &80, &tier_3, &tier_3);
-
+ client.set_risk_tier(&admin, &user1, &20, &tier_1, &tier_1);
+ client.set_risk_tier(&admin, &user2, &50, &tier_2, &tier_2);
+ client.set_risk_tier(&admin, &user3, &80, &tier_3, &tier_3);
let stats = client.get_tier_stats();
assert_eq!(stats.get(tier_1).unwrap(), 1);
assert_eq!(stats.get(tier_2).unwrap(), 1);
@@ -400,64 +429,48 @@ mod tests {
#[test]
fn test_score_update_overwrites_previous() {
- let env = Env::default();
- let contract_id = env.register_contract(None, RiskTierContract);
+ let (env, contract_id, admin) = setup();
let client = RiskTierContractClient::new(&env, &contract_id);
-
let user = Address::generate(&env);
let tier_2 = Symbol::new(&env, "TIER_2");
let tier_1 = Symbol::new(&env, "TIER_1");
-
- client.set_risk_tier(&user, &50, &tier_2, &tier_2);
- client.set_risk_tier(&user, &25, &tier_1, &tier_1);
-
- let risk_data = client.get_risk_tier(&user).unwrap();
- assert_eq!(risk_data.score, 25);
- assert_eq!(risk_data.tier, tier_1);
+ client.set_risk_tier(&admin, &user, &50, &tier_2, &tier_2);
+ client.set_risk_tier(&admin, &user, &25, &tier_1, &tier_1);
+ let data = client.get_risk_tier(&user).unwrap();
+ assert_eq!(data.score, 25);
+ assert_eq!(data.tier, tier_1);
}
#[test]
fn test_no_risk_data_returns_zero_score() {
- let env = Env::default();
- let contract_id = env.register_contract(None, RiskTierContract);
+ let (env, contract_id, _admin) = setup();
let client = RiskTierContractClient::new(&env, &contract_id);
-
let user = Address::generate(&env);
-
- let score = client.get_score(&user);
- assert_eq!(score, 0);
+ assert_eq!(client.get_score(&user), 0);
}
#[test]
fn test_no_risk_data_denies_tier_access() {
- let env = Env::default();
- let contract_id = env.register_contract(None, RiskTierContract);
+ let (env, contract_id, _admin) = setup();
let client = RiskTierContractClient::new(&env, &contract_id);
-
let user = Address::generate(&env);
let tier_3 = Symbol::new(&env, "TIER_3");
-
assert!(!client.can_access_tier(&user, &tier_3));
}
#[test]
fn test_multiple_users_different_tiers() {
- let env = Env::default();
- let contract_id = env.register_contract(None, RiskTierContract);
+ let (env, contract_id, admin) = setup();
let client = RiskTierContractClient::new(&env, &contract_id);
-
let user1 = Address::generate(&env);
let user2 = Address::generate(&env);
let user3 = Address::generate(&env);
-
let tier_1 = Symbol::new(&env, "TIER_1");
let tier_2 = Symbol::new(&env, "TIER_2");
let tier_3 = Symbol::new(&env, "TIER_3");
-
- client.set_risk_tier(&user1, &15, &tier_1, &tier_1);
- client.set_risk_tier(&user2, &45, &tier_2, &tier_2);
- client.set_risk_tier(&user3, &90, &tier_3, &tier_3);
-
+ client.set_risk_tier(&admin, &user1, &15, &tier_1, &tier_1);
+ client.set_risk_tier(&admin, &user2, &45, &tier_2, &tier_2);
+ client.set_risk_tier(&admin, &user3, &90, &tier_3, &tier_3);
assert!(client.can_access_tier(&user1, &tier_1));
assert!(!client.can_access_tier(&user2, &tier_1));
assert!(client.can_access_tier(&user2, &tier_2));