use sn_data_types::{AccountId, Money, Work};
use std::{cmp::Ordering, collections::HashMap};
pub trait RewardAlgo {
fn set(&mut self, base_cost: Money);
fn work_cost(&self, reward_units: u64) -> Money;
fn total_reward(&self, factor: f64, work_cost: Money) -> Money;
fn distribute(
&self,
total_reward: Money,
accounts_work: HashMap<AccountId, Work>,
) -> HashMap<AccountId, Money>;
}
#[derive(Clone)]
pub struct StorageRewards {
base_cost: Money,
}
impl StorageRewards {
pub fn new(base_cost: Money) -> Self {
Self { base_cost }
}
}
impl RewardAlgo for StorageRewards {
fn set(&mut self, base_cost: Money) {
self.base_cost = base_cost;
}
fn work_cost(&self, num_bytes: u64) -> Money {
Money::from_nano(num_bytes + self.base_cost.as_nano())
}
fn total_reward(&self, factor: f64, work_cost: Money) -> Money {
let amount = factor * work_cost.as_nano() as f64;
Money::from_nano(amount.round() as u64)
}
#[allow(clippy::needless_range_loop)]
fn distribute(
&self,
total_reward: Money,
accounts_work: HashMap<AccountId, Work>,
) -> HashMap<AccountId, Money> {
let total_reward = total_reward.as_nano();
let all_work: Work = accounts_work.values().sum();
let mut shares_sum = 0;
let mut shares: Vec<(AccountId, u64)> = Default::default();
for (id, work) in &accounts_work {
let share = (total_reward as f64 / (all_work as f64 / *work as f64)).round() as u64;
shares.push((*id, share));
shares_sum += share;
}
match total_reward.cmp(&shares_sum) {
Ordering::Greater => {
if !shares.is_empty() {
shares.sort_by_key(|t| t.1);
let index = 0; let (id, share) = shares[index];
let remainder = total_reward - shares_sum;
let new_share = share + remainder;
shares[index] = (id, new_share);
}
}
Ordering::Less => {
let mut diff = shares_sum - total_reward;
shares.sort_by_key(|t| t.1);
while diff > 0 {
for i in 0..shares.len() {
let (id, share) = shares[i];
if 0 == diff {
break;
} else if share >= 1 {
shares[i] = (id, share - 1);
diff -= 1;
}
}
}
}
Ordering::Equal => (),
};
let shares_sum = (&shares).iter().map(|(_, share)| share).sum();
if total_reward != shares_sum {
panic!("total_reward: {}, shares_sum: {}", total_reward, shares_sum);
}
shares
.into_iter()
.map(|(i, s)| (i, Money::from_nano(s)))
.collect()
}
}
#[cfg(test)]
mod test {
use super::*;
use sn_data_types::{Money, PublicKey, Result};
use threshold_crypto::SecretKey;
fn get_random_pk() -> PublicKey {
PublicKey::from(SecretKey::random().public_key())
}
#[test]
fn distributes_proportionally() -> Result<()> {
let calc = StorageRewards::new(Money::from_nano(0));
let accounts_work = (1..8).map(|i| (get_random_pk(), i)).collect();
let mut dist: Vec<Money> = calc
.distribute(Money::from_nano(28), accounts_work)
.into_iter()
.map(|(_, reward)| reward)
.collect();
dist.sort();
for (i, amount) in dist.iter().enumerate().take(7) {
assert_eq!(amount.as_nano(), (i + 1) as u64);
}
Ok(())
}
}