use chrono::{DateTime, Utc};
use rust_decimal::Decimal;
use rust_decimal_macros::dec;
use serde::{Deserialize, Serialize};
use sqlx::FromRow;
use std::collections::HashMap;
use uuid::Uuid;
use crate::error::{CoreError, Result};
#[derive(Debug, Clone, Serialize, Deserialize, FromRow)]
pub struct ReferralCode {
pub code_id: Uuid,
pub owner_user_id: Uuid,
pub code: String,
pub usage_count: u64,
pub total_commission_earned: Decimal,
pub is_active: bool,
pub created_at: DateTime<Utc>,
pub deactivated_at: Option<DateTime<Utc>>,
}
impl ReferralCode {
pub fn new(owner_user_id: Uuid, code: String) -> Self {
Self {
code_id: Uuid::new_v4(),
owner_user_id,
code,
usage_count: 0,
total_commission_earned: Decimal::ZERO,
is_active: true,
created_at: Utc::now(),
deactivated_at: None,
}
}
pub fn deactivate(&mut self) {
self.is_active = false;
self.deactivated_at = Some(Utc::now());
}
pub fn increment_usage(&mut self) {
self.usage_count += 1;
}
pub fn add_commission(&mut self, amount: Decimal) {
self.total_commission_earned += amount;
}
}
#[derive(Debug, Clone, Serialize, Deserialize, FromRow)]
pub struct ReferralRelationship {
pub relationship_id: Uuid,
pub referred_user_id: Uuid,
pub referrer_user_id: Uuid,
pub referral_code: String,
pub referred_at: DateTime<Utc>,
pub total_trading_volume: Decimal,
pub total_commission_earned: Decimal,
pub is_eligible: bool,
}
impl ReferralRelationship {
pub fn new(referred_user_id: Uuid, referrer_user_id: Uuid, referral_code: String) -> Self {
Self {
relationship_id: Uuid::new_v4(),
referred_user_id,
referrer_user_id,
referral_code,
referred_at: Utc::now(),
total_trading_volume: Decimal::ZERO,
total_commission_earned: Decimal::ZERO,
is_eligible: false,
}
}
pub fn add_trading_volume(&mut self, volume: Decimal, min_volume_for_eligibility: Decimal) {
self.total_trading_volume += volume;
if self.total_trading_volume >= min_volume_for_eligibility {
self.is_eligible = true;
}
}
pub fn add_commission(&mut self, amount: Decimal) {
self.total_commission_earned += amount;
}
}
#[derive(Debug, Clone, Serialize, Deserialize, FromRow)]
pub struct CommissionDistribution {
pub distribution_id: Uuid,
pub recipient_user_id: Uuid,
pub referred_user_id: Uuid,
pub tier: u8,
pub fee_amount: Decimal,
pub commission_amount: Decimal,
pub commission_rate: Decimal,
pub token_id: Uuid,
pub distributed_at: DateTime<Utc>,
}
impl CommissionDistribution {
#[allow(clippy::too_many_arguments)]
pub fn new(
recipient_user_id: Uuid,
referred_user_id: Uuid,
tier: u8,
fee_amount: Decimal,
commission_rate: Decimal,
token_id: Uuid,
) -> Self {
let commission_amount = fee_amount * commission_rate;
Self {
distribution_id: Uuid::new_v4(),
recipient_user_id,
referred_user_id,
tier,
fee_amount,
commission_amount,
commission_rate,
token_id,
distributed_at: Utc::now(),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ReferralConfig {
pub tier1_commission_rate: Decimal,
pub tier2_commission_rate: Decimal,
pub tier3_commission_rate: Decimal,
pub min_volume_for_eligibility: Decimal,
pub referral_cooldown_hours: u64,
pub max_referrals_per_user: Option<u64>,
}
impl Default for ReferralConfig {
fn default() -> Self {
Self {
tier1_commission_rate: dec!(0.30), tier2_commission_rate: dec!(0.10), tier3_commission_rate: dec!(0.05), min_volume_for_eligibility: dec!(100.0), referral_cooldown_hours: 24, max_referrals_per_user: Some(1000),
}
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ReferralLeaderboardEntry {
pub user_id: Uuid,
pub username: Option<String>,
pub total_referrals: u64,
pub eligible_referrals: u64,
pub total_commission_earned: Decimal,
pub rank: u64,
}
pub struct ReferralManager {
codes: HashMap<String, ReferralCode>,
relationships: HashMap<Uuid, ReferralRelationship>,
user_codes: HashMap<Uuid, String>,
user_referrals: HashMap<Uuid, Vec<Uuid>>,
config: ReferralConfig,
}
impl ReferralManager {
pub fn new(config: ReferralConfig) -> Self {
Self {
codes: HashMap::new(),
relationships: HashMap::new(),
user_codes: HashMap::new(),
user_referrals: HashMap::new(),
config,
}
}
pub fn generate_code(
&mut self,
user_id: Uuid,
preferred_code: Option<String>,
) -> Result<String> {
if self.user_codes.contains_key(&user_id) {
return Err(CoreError::AlreadyExists(
"User already has a referral code".to_string(),
));
}
let code = if let Some(pref) = preferred_code {
if pref.len() < 4 || pref.len() > 20 {
return Err(CoreError::Validation(
"Code must be between 4 and 20 characters".to_string(),
));
}
if !pref.chars().all(|c| c.is_ascii_alphanumeric()) {
return Err(CoreError::Validation(
"Code must contain only alphanumeric characters".to_string(),
));
}
if self.codes.contains_key(&pref.to_uppercase()) {
return Err(CoreError::AlreadyExists("Code already in use".to_string()));
}
pref.to_uppercase()
} else {
self.auto_generate_code(user_id)?
};
let referral_code = ReferralCode::new(user_id, code.clone());
self.codes.insert(code.clone(), referral_code);
self.user_codes.insert(user_id, code.clone());
Ok(code)
}
fn auto_generate_code(&self, user_id: Uuid) -> Result<String> {
let code = format!(
"REF{}",
&user_id.to_string().replace("-", "")[..8].to_uppercase()
);
if self.codes.contains_key(&code) {
let suffix = chrono::Utc::now().timestamp() % 10000;
Ok(format!("{}{}", code, suffix))
} else {
Ok(code)
}
}
pub fn use_code(&mut self, code: &str, new_user_id: Uuid) -> Result<Uuid> {
let code = code.to_uppercase();
let mut referral_code = self
.codes
.get(&code)
.ok_or_else(|| CoreError::NotFound("Referral code not found".to_string()))?
.clone();
if !referral_code.is_active {
return Err(CoreError::Validation(
"Referral code is inactive".to_string(),
));
}
let referrer_user_id = referral_code.owner_user_id;
if referrer_user_id == new_user_id {
return Err(CoreError::Validation("Cannot refer yourself".to_string()));
}
if let Some(max) = self.config.max_referrals_per_user {
let current_count = self
.user_referrals
.get(&referrer_user_id)
.map(|v| v.len() as u64)
.unwrap_or(0);
if current_count >= max {
return Err(CoreError::Validation(
"Maximum referrals reached".to_string(),
));
}
}
let relationship = ReferralRelationship::new(new_user_id, referrer_user_id, code.clone());
let relationship_id = relationship.relationship_id;
self.relationships.insert(new_user_id, relationship);
referral_code.increment_usage();
self.codes.insert(code, referral_code);
self.user_referrals
.entry(referrer_user_id)
.or_default()
.push(new_user_id);
Ok(relationship_id)
}
pub fn distribute_commissions(
&mut self,
trader_user_id: Uuid,
fee_amount: Decimal,
token_id: Uuid,
) -> Result<Vec<CommissionDistribution>> {
let mut distributions = Vec::new();
let relationship = match self.relationships.get_mut(&trader_user_id) {
Some(r) => r,
None => return Ok(distributions), };
if !relationship.is_eligible {
return Ok(distributions); }
let tier1_user = relationship.referrer_user_id;
let tier1_commission = CommissionDistribution::new(
tier1_user,
trader_user_id,
1,
fee_amount,
self.config.tier1_commission_rate,
token_id,
);
relationship.add_commission(tier1_commission.commission_amount);
if let Some(code_str) = self.user_codes.get(&tier1_user) {
if let Some(code) = self.codes.get_mut(code_str) {
code.add_commission(tier1_commission.commission_amount);
}
}
distributions.push(tier1_commission);
if let Some(tier1_relationship) = self.relationships.get(&tier1_user).cloned() {
if tier1_relationship.is_eligible {
let tier2_user = tier1_relationship.referrer_user_id;
let tier2_commission = CommissionDistribution::new(
tier2_user,
trader_user_id,
2,
fee_amount,
self.config.tier2_commission_rate,
token_id,
);
distributions.push(tier2_commission);
}
}
if let Some(tier1_relationship) = self.relationships.get(&tier1_user) {
if tier1_relationship.is_eligible {
let tier2_user = tier1_relationship.referrer_user_id;
if let Some(tier2_relationship) = self.relationships.get(&tier2_user).cloned() {
if tier2_relationship.is_eligible {
let tier3_user = tier2_relationship.referrer_user_id;
let tier3_commission = CommissionDistribution::new(
tier3_user,
trader_user_id,
3,
fee_amount,
self.config.tier3_commission_rate,
token_id,
);
distributions.push(tier3_commission);
}
}
}
}
Ok(distributions)
}
pub fn record_trading_volume(&mut self, user_id: Uuid, volume: Decimal) -> Result<()> {
if let Some(relationship) = self.relationships.get_mut(&user_id) {
relationship.add_trading_volume(volume, self.config.min_volume_for_eligibility);
}
Ok(())
}
pub fn get_user_code(&self, user_id: Uuid) -> Option<&String> {
self.user_codes.get(&user_id)
}
pub fn get_user_stats(&self, user_id: Uuid) -> ReferralStats {
let code = self.get_user_code(user_id);
let total_referrals = self
.user_referrals
.get(&user_id)
.map(|v| v.len() as u64)
.unwrap_or(0);
let eligible_referrals = self
.user_referrals
.get(&user_id)
.map(|refs| {
refs.iter()
.filter(|&&ref_id| {
self.relationships
.get(&ref_id)
.map(|r| r.is_eligible)
.unwrap_or(false)
})
.count() as u64
})
.unwrap_or(0);
let total_commission = code
.and_then(|c| self.codes.get(c))
.map(|code| code.total_commission_earned)
.unwrap_or(Decimal::ZERO);
ReferralStats {
user_id,
referral_code: code.cloned(),
total_referrals,
eligible_referrals,
total_commission_earned: total_commission,
}
}
pub fn get_leaderboard(&self, limit: usize) -> Vec<ReferralLeaderboardEntry> {
let mut entries: Vec<ReferralLeaderboardEntry> = self
.user_codes
.keys()
.map(|&user_id| {
let stats = self.get_user_stats(user_id);
ReferralLeaderboardEntry {
user_id,
username: None, total_referrals: stats.total_referrals,
eligible_referrals: stats.eligible_referrals,
total_commission_earned: stats.total_commission_earned,
rank: 0, }
})
.collect();
entries.sort_by(|a, b| {
b.total_commission_earned
.cmp(&a.total_commission_earned)
.then(b.total_referrals.cmp(&a.total_referrals))
});
entries
.into_iter()
.take(limit)
.enumerate()
.map(|(idx, mut entry)| {
entry.rank = (idx + 1) as u64;
entry
})
.collect()
}
}
impl Default for ReferralManager {
fn default() -> Self {
Self::new(ReferralConfig::default())
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct ReferralStats {
pub user_id: Uuid,
pub referral_code: Option<String>,
pub total_referrals: u64,
pub eligible_referrals: u64,
pub total_commission_earned: Decimal,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_generate_code() {
let mut manager = ReferralManager::default();
let user_id = Uuid::new_v4();
let code = manager
.generate_code(user_id, Some("ALICE2023".to_string()))
.unwrap();
assert_eq!(code, "ALICE2023");
assert_eq!(
manager.get_user_code(user_id),
Some(&"ALICE2023".to_string())
);
}
#[test]
fn test_use_code() {
let mut manager = ReferralManager::default();
let referrer_id = Uuid::new_v4();
let new_user_id = Uuid::new_v4();
let code = manager
.generate_code(referrer_id, Some("BOB2023".to_string()))
.unwrap();
manager.use_code(&code, new_user_id).unwrap();
let stats = manager.get_user_stats(referrer_id);
assert_eq!(stats.total_referrals, 1);
}
#[test]
fn test_cannot_self_refer() {
let mut manager = ReferralManager::default();
let user_id = Uuid::new_v4();
let code = manager
.generate_code(user_id, Some("SELF".to_string()))
.unwrap();
let result = manager.use_code(&code, user_id);
assert!(result.is_err());
}
#[test]
fn test_commission_distribution() {
let mut manager = ReferralManager::default();
let tier3_id = Uuid::new_v4();
let tier2_id = Uuid::new_v4();
let tier1_id = Uuid::new_v4();
let trader_id = Uuid::new_v4();
manager
.generate_code(tier3_id, Some("TIER3".to_string()))
.unwrap();
manager
.generate_code(tier2_id, Some("TIER2".to_string()))
.unwrap();
manager
.generate_code(tier1_id, Some("TIER1".to_string()))
.unwrap();
manager.use_code("TIER3", tier2_id).unwrap();
manager.use_code("TIER2", tier1_id).unwrap();
manager.use_code("TIER1", trader_id).unwrap();
manager
.record_trading_volume(tier2_id, dec!(100.0))
.unwrap();
manager
.record_trading_volume(tier1_id, dec!(100.0))
.unwrap();
manager
.record_trading_volume(trader_id, dec!(100.0))
.unwrap();
let token_id = Uuid::new_v4();
let commissions = manager
.distribute_commissions(trader_id, dec!(10.0), token_id)
.unwrap();
assert_eq!(commissions.len(), 3); assert_eq!(commissions[0].tier, 1);
assert_eq!(commissions[0].commission_amount, dec!(3.0)); assert_eq!(commissions[1].tier, 2);
assert_eq!(commissions[1].commission_amount, dec!(1.0)); assert_eq!(commissions[2].tier, 3);
assert_eq!(commissions[2].commission_amount, dec!(0.5)); }
#[test]
fn test_eligibility_threshold() {
let mut manager = ReferralManager::default();
let referrer_id = Uuid::new_v4();
let referred_id = Uuid::new_v4();
manager
.generate_code(referrer_id, Some("REFER".to_string()))
.unwrap();
manager.use_code("REFER", referred_id).unwrap();
let token_id = Uuid::new_v4();
let commissions = manager
.distribute_commissions(referred_id, dec!(10.0), token_id)
.unwrap();
assert_eq!(commissions.len(), 0);
manager
.record_trading_volume(referred_id, dec!(100.0))
.unwrap();
let commissions = manager
.distribute_commissions(referred_id, dec!(10.0), token_id)
.unwrap();
assert_eq!(commissions.len(), 1);
}
#[test]
fn test_leaderboard() {
let mut manager = ReferralManager::default();
let alice_id = Uuid::new_v4();
let bob_id = Uuid::new_v4();
manager
.generate_code(alice_id, Some("ALICE".to_string()))
.unwrap();
manager
.generate_code(bob_id, Some("BOBBY".to_string()))
.unwrap();
for _ in 0..3 {
let user = Uuid::new_v4();
manager.use_code("ALICE", user).unwrap();
manager.record_trading_volume(user, dec!(100.0)).unwrap();
manager
.distribute_commissions(user, dec!(10.0), Uuid::new_v4())
.unwrap();
}
let user = Uuid::new_v4();
manager.use_code("BOBBY", user).unwrap();
let leaderboard = manager.get_leaderboard(10);
assert!(!leaderboard.is_empty());
assert_eq!(leaderboard[0].user_id, alice_id); }
}