1use std::collections::{BTreeMap, HashMap};
6
7use tracing::instrument;
8
9use crate::nuts::Id;
10use crate::{Amount, Error};
11
12#[derive(Debug, Clone, PartialEq)]
14pub struct ProofsFeeBreakdown {
15 pub total: Amount,
17 pub per_keyset: HashMap<Id, Amount>,
19}
20
21#[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 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 let keyset_fee = if i == keyset_count - 1 {
63 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 let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
319
320 assert_eq!(breakdown.total, 1.into());
321
322 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 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 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 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 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); keyset_fees.insert(keyset_id_2, 1000);
401
402 let mut proofs_count = HashMap::new();
403 proofs_count.insert(keyset_id_1, 3); proofs_count.insert(keyset_id_2, 7); let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
408
409 assert_eq!(breakdown.total, 10.into());
410 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); proofs_count.insert(keyset_id_2, 6); let breakdown = calculate_fee(&proofs_count, &keyset_fees).unwrap();
433
434 assert_eq!(breakdown.total, 2.into());
435
436 let per_keyset_sum: u64 = breakdown.per_keyset.values().map(|a| u64::from(*a)).sum();
438 assert_eq!(per_keyset_sum, 2);
439 }
440}