1pub struct FeeSplitArgs {
2 pub fee_amount: u64,
3 pub protocol_bps: u64,
4 pub lp_bps: u64,
5 pub base_total_bps: u64,
6 pub decay_premium_bps: u64,
7}
8
9#[must_use]
12pub fn split_fee_amount(args: FeeSplitArgs) -> Option<(u64, u64, u64, u64)> {
13 let fee_amount = args.fee_amount;
14 if fee_amount == 0 {
16 return Some((0, 0, 0, 0));
17 }
18 let effective_total_bps =
19 args.base_total_bps.checked_add(args.decay_premium_bps)?;
20 if effective_total_bps == 0 {
21 return None;
22 }
23
24 if args.base_total_bps == 0 {
25 return Some((fee_amount, 0, 0, fee_amount));
26 }
27
28 let lp_fee = u128::from(fee_amount)
29 .checked_mul(u128::from(args.lp_bps))?
30 .checked_div(u128::from(effective_total_bps))?;
31 let lp_fee = u64::try_from(lp_fee).ok()?;
32
33 let creator_bps = args
34 .base_total_bps
35 .checked_sub(args.protocol_bps)?
36 .checked_sub(args.lp_bps)?;
37 let creator_fee = u128::from(fee_amount)
38 .checked_mul(u128::from(creator_bps))?
39 .checked_div(u128::from(effective_total_bps))?;
40 let creator_fee = u64::try_from(creator_fee).ok()?;
41
42 let sniper_fee = u128::from(fee_amount)
43 .checked_mul(u128::from(args.decay_premium_bps))?
44 .checked_div(u128::from(effective_total_bps))?;
45 let sniper_fee = u64::try_from(sniper_fee).ok()?;
46
47 let protocol_fee =
48 fee_amount.checked_sub(lp_fee)?.checked_sub(creator_fee)?;
49
50 Some((protocol_fee, lp_fee, creator_fee, sniper_fee))
51}
52
53#[cfg(test)]
54mod tests {
55 use super::*;
56
57 const STANDARD: FeeSplitArgs = FeeSplitArgs {
58 fee_amount: 1000,
59 protocol_bps: 200,
60 lp_bps: 200,
61 base_total_bps: 1000,
62 decay_premium_bps: 0,
63 };
64
65 #[test]
66 fn test_split_fee_amount_no_dust_lost() {
67 let (protocol_fee, lp_fee, creator_fee, sniper_fee) =
68 split_fee_amount(STANDARD).unwrap();
69
70 assert_eq!(protocol_fee, 200);
71 assert_eq!(lp_fee, 200);
72 assert_eq!(creator_fee, 600);
73 assert_eq!(sniper_fee, 0);
74 assert_eq!(protocol_fee + lp_fee + creator_fee, 1000);
75 }
76
77 #[test]
78 fn test_split_fee_amount_protocol_wins_rounding() {
79 let (protocol_fee, lp_fee, creator_fee, sniper_fee) =
80 split_fee_amount(FeeSplitArgs {
81 fee_amount: 1001,
82 ..STANDARD
83 })
84 .unwrap();
85
86 assert_eq!(lp_fee, 200);
87 assert_eq!(creator_fee, 600);
88 assert_eq!(protocol_fee, 201);
89 assert_eq!(sniper_fee, 0);
90 assert_eq!(protocol_fee + lp_fee + creator_fee, 1001);
91 }
92
93 #[test]
94 fn test_split_fee_with_decay_protocol_wins_rounding() {
95 let (protocol_fee, lp_fee, creator_fee, sniper_fee) =
96 split_fee_amount(FeeSplitArgs {
97 fee_amount: 1001,
98 decay_premium_bps: 500,
99 ..STANDARD
100 })
101 .unwrap();
102
103 assert_eq!(lp_fee, 133);
104 assert_eq!(creator_fee, 400);
105 assert_eq!(protocol_fee, 468);
106 assert_eq!(sniper_fee, 333);
107 assert_eq!(protocol_fee + lp_fee + creator_fee, 1001);
108 }
109
110 #[test]
111 fn test_split_fee_max_amount() {
112 let fee_amount = 1_000_000_000_000_000u64;
113 let result = split_fee_amount(FeeSplitArgs {
114 fee_amount,
115 protocol_bps: 100,
116 lp_bps: 100,
117 base_total_bps: 300,
118 decay_premium_bps: 0,
119 });
120 assert!(result.is_some());
121 let (protocol, lp, creator, _decay) = result.unwrap();
122 assert_eq!(protocol + lp + creator, fee_amount);
123 }
124
125 #[test]
126 fn test_split_fee_zero_amount() {
127 assert_eq!(
128 split_fee_amount(FeeSplitArgs {
129 fee_amount: 0,
130 protocol_bps: 100,
131 lp_bps: 100,
132 base_total_bps: 300,
133 decay_premium_bps: 0,
134 }),
135 Some((0, 0, 0, 0)),
136 );
137 }
138
139 #[test]
140 fn test_split_fee_invalid_bps_underflows() {
141 assert_eq!(
142 split_fee_amount(FeeSplitArgs {
143 protocol_bps: 600,
144 lp_bps: 500,
145 ..STANDARD
146 }),
147 None,
148 );
149 }
150
151 #[test]
152 fn test_split_fee_zero_total_bps() {
153 assert_eq!(
154 split_fee_amount(FeeSplitArgs {
155 fee_amount: 1000,
156 protocol_bps: 0,
157 lp_bps: 0,
158 base_total_bps: 0,
159 decay_premium_bps: 0,
160 }),
161 None,
162 );
163 }
164
165 #[test]
166 fn test_split_fee_zero_amount_and_zero_total_bps() {
167 assert_eq!(
168 split_fee_amount(FeeSplitArgs {
169 fee_amount: 0,
170 protocol_bps: 0,
171 lp_bps: 0,
172 base_total_bps: 0,
173 decay_premium_bps: 0,
174 }),
175 Some((0, 0, 0, 0)),
176 );
177 }
178}