Skip to main content

cdk/
fees.rs

1//! Calculate fees
2//!
3//! <https://github.com/cashubtc/nuts/blob/main/02.md>
4
5use std::collections::{BTreeMap, HashMap};
6
7use tracing::instrument;
8
9use crate::nuts::Id;
10use crate::{Amount, Error};
11
12/// Fee breakdown containing total fee and fee per keyset
13#[derive(Debug, Clone, PartialEq)]
14pub struct ProofsFeeBreakdown {
15    /// Total fee across all keysets
16    pub total: Amount,
17    /// Fee collected per keyset
18    pub per_keyset: HashMap<Id, Amount>,
19}
20
21/// Fee required for proof set
22#[instrument(skip_all)]
23pub fn calculate_fee(
24    proofs_count: &HashMap<Id, u64>,
25    keyset_fee: &HashMap<Id, u64>,
26) -> Result<ProofsFeeBreakdown, Error> {
27    let mut sum_fee: u64 = 0;
28    let mut fee_per_keyset_raw: BTreeMap<Id, u64> = BTreeMap::new();
29
30    for (keyset_id, proof_count) in proofs_count {
31        let keyset_fee_ppk = *keyset_fee
32            .get(keyset_id)
33            .ok_or(Error::KeysetUnknown(*keyset_id))?;
34
35        let proofs_fee = keyset_fee_ppk
36            .checked_mul(*proof_count)
37            .ok_or(Error::AmountOverflow)?;
38
39        sum_fee = sum_fee
40            .checked_add(proofs_fee)
41            .ok_or(Error::AmountOverflow)?;
42
43        fee_per_keyset_raw.insert(*keyset_id, proofs_fee);
44    }
45
46    let total_fee = (sum_fee.checked_add(999).ok_or(Error::AmountOverflow)?)
47        .checked_div(1000)
48        .ok_or(Error::AmountOverflow)?;
49
50    // Calculate fee per keyset proportionally based on the total
51    // BTreeMap ensures deterministic iteration order (sorted by keyset ID)
52    let mut per_keyset = HashMap::new();
53    let mut distributed_fee: u64 = 0;
54    let keyset_count = fee_per_keyset_raw.len();
55
56    for (i, (keyset_id, raw_fee)) in fee_per_keyset_raw.iter().enumerate() {
57        if sum_fee == 0 {
58            continue;
59        }
60
61        // Calculate proportional fee, rounding down
62        let keyset_fee = if i == keyset_count - 1 {
63            // Last keyset gets the remainder to ensure total matches
64            total_fee.saturating_sub(distributed_fee)
65        } else {
66            (raw_fee.checked_mul(total_fee))
67                .ok_or(Error::AmountOverflow)?
68                .checked_div(sum_fee)
69                .ok_or(Error::AmountOverflow)?
70        };
71
72        distributed_fee = distributed_fee.saturating_add(keyset_fee);
73        per_keyset.insert(*keyset_id, keyset_fee.into());
74    }
75
76    Ok(ProofsFeeBreakdown {
77        total: total_fee.into(),
78        per_keyset,
79    })
80}
81
82#[cfg(test)]
83mod tests {
84
85    use std::str::FromStr;
86
87    use super::*;
88
89    #[test]
90    fn test_calc_fee() {
91        let keyset_id = Id::from_str("001711afb1de20cb").unwrap();
92
93        let fee = 2;
94
95        let mut keyset_fees = HashMap::new();
96        keyset_fees.insert(keyset_id, fee);
97
98        let mut proofs_count = HashMap::new();
99
100        proofs_count.insert(keyset_id, 1);
101
102        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
103
104        assert_eq!(breakdown.total, 1.into());
105        assert_eq!(breakdown.per_keyset[&keyset_id], 1.into());
106
107        proofs_count.insert(keyset_id, 500);
108
109        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
110
111        assert_eq!(breakdown.total, 1.into());
112        assert_eq!(breakdown.per_keyset[&keyset_id], 1.into());
113
114        proofs_count.insert(keyset_id, 1000);
115
116        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
117
118        assert_eq!(breakdown.total, 2.into());
119        assert_eq!(breakdown.per_keyset[&keyset_id], 2.into());
120
121        proofs_count.insert(keyset_id, 2000);
122        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
123        assert_eq!(breakdown.total, 4.into());
124        assert_eq!(breakdown.per_keyset[&keyset_id], 4.into());
125
126        proofs_count.insert(keyset_id, 3500);
127        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
128        assert_eq!(breakdown.total, 7.into());
129        assert_eq!(breakdown.per_keyset[&keyset_id], 7.into());
130
131        proofs_count.insert(keyset_id, 3501);
132        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
133        assert_eq!(breakdown.total, 8.into());
134        assert_eq!(breakdown.per_keyset[&keyset_id], 8.into());
135    }
136
137    #[test]
138    fn test_fee_calculation_with_ppk_200() {
139        let keyset_id = Id::from_str("001711afb1de20cb").unwrap();
140
141        let fee_ppk = 200;
142
143        let mut keyset_fees = HashMap::new();
144        keyset_fees.insert(keyset_id, fee_ppk);
145
146        let mut proofs_count = HashMap::new();
147
148        proofs_count.insert(keyset_id, 1);
149        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
150        assert_eq!(breakdown.total, 1.into(), "1 proof: ceil(200/1000) = 1 sat");
151
152        proofs_count.insert(keyset_id, 3);
153        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
154        assert_eq!(
155            breakdown.total,
156            1.into(),
157            "3 proofs: ceil(600/1000) = 1 sat"
158        );
159
160        proofs_count.insert(keyset_id, 5);
161        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
162        assert_eq!(
163            breakdown.total,
164            1.into(),
165            "5 proofs: ceil(1000/1000) = 1 sat"
166        );
167
168        proofs_count.insert(keyset_id, 6);
169        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
170        assert_eq!(
171            breakdown.total,
172            2.into(),
173            "6 proofs: ceil(1200/1000) = 2 sats"
174        );
175    }
176
177    #[test]
178    fn test_fee_calculation_with_ppk_1000() {
179        let keyset_id = Id::from_str("001711afb1de20cb").unwrap();
180
181        let fee_ppk = 1000;
182
183        let mut keyset_fees = HashMap::new();
184        keyset_fees.insert(keyset_id, fee_ppk);
185
186        let mut proofs_count = HashMap::new();
187
188        proofs_count.insert(keyset_id, 1);
189        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
190        assert_eq!(breakdown.total, 1.into(), "1 proof at 1000 ppk = 1 sat");
191
192        proofs_count.insert(keyset_id, 2);
193        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
194        assert_eq!(breakdown.total, 2.into(), "2 proofs at 1000 ppk = 2 sats");
195
196        proofs_count.insert(keyset_id, 10);
197        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
198        assert_eq!(
199            breakdown.total,
200            10.into(),
201            "10 proofs at 1000 ppk = 10 sats"
202        );
203    }
204
205    #[test]
206    fn test_fee_calculation_zero_fee() {
207        let keyset_id = Id::from_str("001711afb1de20cb").unwrap();
208
209        let fee_ppk = 0;
210
211        let mut keyset_fees = HashMap::new();
212        keyset_fees.insert(keyset_id, fee_ppk);
213
214        let mut proofs_count = HashMap::new();
215
216        proofs_count.insert(keyset_id, 100);
217        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
218        assert_eq!(
219            breakdown.total,
220            0.into(),
221            "0 ppk means no fee: ceil(0/1000) = 0"
222        );
223    }
224
225    #[test]
226    fn test_fee_calculation_with_ppk_100() {
227        let keyset_id = Id::from_str("001711afb1de20cb").unwrap();
228
229        let fee_ppk = 100;
230
231        let mut keyset_fees = HashMap::new();
232        keyset_fees.insert(keyset_id, fee_ppk);
233
234        let mut proofs_count = HashMap::new();
235
236        proofs_count.insert(keyset_id, 1);
237        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
238        assert_eq!(breakdown.total, 1.into(), "1 proof: ceil(100/1000) = 1 sat");
239
240        proofs_count.insert(keyset_id, 10);
241        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
242        assert_eq!(
243            breakdown.total,
244            1.into(),
245            "10 proofs: ceil(1000/1000) = 1 sat"
246        );
247
248        proofs_count.insert(keyset_id, 11);
249        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
250        assert_eq!(
251            breakdown.total,
252            2.into(),
253            "11 proofs: ceil(1100/1000) = 2 sats"
254        );
255
256        proofs_count.insert(keyset_id, 91);
257        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
258        assert_eq!(
259            breakdown.total,
260            10.into(),
261            "91 proofs: ceil(9100/1000) = 10 sats"
262        );
263    }
264
265    #[test]
266    fn test_fee_calculation_unknown_keyset() {
267        let keyset_id = Id::from_str("001711afb1de20cb").unwrap();
268        let unknown_keyset_id = Id::from_str("001711afb1de20cc").unwrap();
269
270        let mut keyset_fees = HashMap::new();
271        keyset_fees.insert(keyset_id, 100);
272
273        let mut proofs_count = HashMap::new();
274        proofs_count.insert(unknown_keyset_id, 1);
275
276        let result = calculate_fee(&proofs_count, &keyset_fees);
277        assert!(result.is_err(), "Unknown keyset should return error");
278    }
279
280    #[test]
281    fn test_fee_calculation_multiple_keysets() {
282        let keyset_id_1 = Id::from_str("001711afb1de20cb").unwrap();
283        let keyset_id_2 = Id::from_str("001711afb1de20cc").unwrap();
284
285        let mut keyset_fees = HashMap::new();
286        keyset_fees.insert(keyset_id_1, 200);
287        keyset_fees.insert(keyset_id_2, 500);
288
289        let mut proofs_count = HashMap::new();
290        proofs_count.insert(keyset_id_1, 3);
291        proofs_count.insert(keyset_id_2, 2);
292
293        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
294        assert_eq!(
295            breakdown.total,
296            2.into(),
297            "3*200 + 2*500 = 1600, ceil(1600/1000) = 2"
298        );
299    }
300
301    #[test]
302    fn test_per_keyset_fee_sums_to_total() {
303        let keyset_id_1 = Id::from_str("001711afb1de20cb").unwrap();
304        let keyset_id_2 = Id::from_str("001711afb1de20cc").unwrap();
305        let keyset_id_3 = Id::from_str("001711afb1de20cd").unwrap();
306
307        let mut keyset_fees = HashMap::new();
308        keyset_fees.insert(keyset_id_1, 100);
309        keyset_fees.insert(keyset_id_2, 100);
310        keyset_fees.insert(keyset_id_3, 100);
311
312        let mut proofs_count = HashMap::new();
313        proofs_count.insert(keyset_id_1, 1);
314        proofs_count.insert(keyset_id_2, 1);
315        proofs_count.insert(keyset_id_3, 1);
316
317        // 3 proofs * 100 ppk = 300 ppk, ceil(300/1000) = 1 sat total
318        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
319
320        assert_eq!(breakdown.total, 1.into());
321
322        // Sum of per_keyset fees must equal total
323        let per_keyset_sum: u64 = breakdown.per_keyset.values().map(|a| u64::from(*a)).sum();
324        assert_eq!(per_keyset_sum, u64::from(breakdown.total));
325    }
326
327    #[test]
328    fn test_per_keyset_fee_remainder_goes_to_last_sorted_keyset() {
329        // Use keyset IDs where sorting order is predictable
330        let keyset_id_1 = Id::from_str("00aaaaaaaaaaaaa1").unwrap();
331        let keyset_id_2 = Id::from_str("00aaaaaaaaaaaaa2").unwrap();
332        let keyset_id_3 = Id::from_str("00aaaaaaaaaaaaa3").unwrap();
333
334        let mut keyset_fees = HashMap::new();
335        keyset_fees.insert(keyset_id_1, 100);
336        keyset_fees.insert(keyset_id_2, 100);
337        keyset_fees.insert(keyset_id_3, 100);
338
339        let mut proofs_count = HashMap::new();
340        proofs_count.insert(keyset_id_1, 1);
341        proofs_count.insert(keyset_id_2, 1);
342        proofs_count.insert(keyset_id_3, 1);
343
344        // 3 * 100 = 300 ppk, ceil(300/1000) = 1 sat total
345        // Each keyset contributed 100/300 = 1/3 of raw fee
346        // Proportional: (100 * 1) / 300 = 0 for first two (integer division)
347        // Last keyset (keyset_id_3) gets remainder: 1 - 0 - 0 = 1
348        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
349
350        assert_eq!(breakdown.total, 1.into());
351        assert_eq!(breakdown.per_keyset[&keyset_id_1], 0.into());
352        assert_eq!(breakdown.per_keyset[&keyset_id_2], 0.into());
353        assert_eq!(breakdown.per_keyset[&keyset_id_3], 1.into());
354    }
355
356    #[test]
357    fn test_per_keyset_fee_distribution_is_deterministic() {
358        let keyset_id_1 = Id::from_str("001711afb1de20cb").unwrap();
359        let keyset_id_2 = Id::from_str("001711afb1de20cc").unwrap();
360
361        let mut keyset_fees = HashMap::new();
362        keyset_fees.insert(keyset_id_1, 333);
363        keyset_fees.insert(keyset_id_2, 333);
364
365        let mut proofs_count = HashMap::new();
366        proofs_count.insert(keyset_id_1, 1);
367        proofs_count.insert(keyset_id_2, 1);
368
369        // Run multiple times to verify determinism
370        let breakdown1 = calculate_fee(&proofs_count, &keyset_fees).unwrap();
371        let breakdown2 = calculate_fee(&proofs_count, &keyset_fees).unwrap();
372        let breakdown3 = calculate_fee(&proofs_count, &keyset_fees).unwrap();
373
374        // All runs should produce identical per-keyset results
375        assert_eq!(
376            breakdown1.per_keyset[&keyset_id_1],
377            breakdown2.per_keyset[&keyset_id_1]
378        );
379        assert_eq!(
380            breakdown1.per_keyset[&keyset_id_2],
381            breakdown2.per_keyset[&keyset_id_2]
382        );
383        assert_eq!(
384            breakdown2.per_keyset[&keyset_id_1],
385            breakdown3.per_keyset[&keyset_id_1]
386        );
387        assert_eq!(
388            breakdown2.per_keyset[&keyset_id_2],
389            breakdown3.per_keyset[&keyset_id_2]
390        );
391    }
392
393    #[test]
394    fn test_per_keyset_fee_proportional_distribution() {
395        let keyset_id_1 = Id::from_str("001711afb1de20cb").unwrap();
396        let keyset_id_2 = Id::from_str("001711afb1de20cc").unwrap();
397
398        let mut keyset_fees = HashMap::new();
399        keyset_fees.insert(keyset_id_1, 1000); // 1 sat per proof
400        keyset_fees.insert(keyset_id_2, 1000);
401
402        let mut proofs_count = HashMap::new();
403        proofs_count.insert(keyset_id_1, 3); // 3000 ppk = 3 sat raw
404        proofs_count.insert(keyset_id_2, 7); // 7000 ppk = 7 sat raw
405
406        // Total: 10000 ppk = 10 sat
407        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
408
409        assert_eq!(breakdown.total, 10.into());
410        // keyset_id_1: (3000 * 10) / 10000 = 3
411        // keyset_id_2: 10 - 3 = 7 (gets remainder, but happens to be exact)
412        assert_eq!(breakdown.per_keyset[&keyset_id_1], 3.into());
413        assert_eq!(breakdown.per_keyset[&keyset_id_2], 7.into());
414    }
415
416    #[test]
417    fn test_per_keyset_fee_with_uneven_distribution() {
418        let keyset_id_1 = Id::from_str("00aaaaaaaaaaaaa1").unwrap();
419        let keyset_id_2 = Id::from_str("00aaaaaaaaaaaaa2").unwrap();
420
421        let mut keyset_fees = HashMap::new();
422        keyset_fees.insert(keyset_id_1, 100);
423        keyset_fees.insert(keyset_id_2, 100);
424
425        let mut proofs_count = HashMap::new();
426        proofs_count.insert(keyset_id_1, 5); // 500 ppk
427        proofs_count.insert(keyset_id_2, 6); // 600 ppk
428
429        // Total: 1100 ppk, ceil(1100/1000) = 2 sat
430        // keyset_id_1: (500 * 2) / 1100 = 0 (integer division)
431        // keyset_id_2: 2 - 0 = 2 (gets remainder)
432        let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
433
434        assert_eq!(breakdown.total, 2.into());
435
436        // Verify sum equals total
437        let per_keyset_sum: u64 = breakdown.per_keyset.values().map(|a| u64::from(*a)).sum();
438        assert_eq!(per_keyset_sum, 2);
439    }
440}