1use bytemuck::Pod;
7use solana_address::Address;
8use solana_program::program_error::ProgramError;
9use spl_token_2022_interface::error::TokenError;
10use spl_token_2022_interface::extension::{
11 AccountType, BaseStateWithExtensions, Extension, PodStateWithExtensions,
12};
13use spl_token_2022_interface::pod::PodMint;
14
15use crate::constants::{TOKEN_2022_PROGRAM_ID, TOKEN_PROGRAM_ID};
16use crate::math::amm::MintFee;
17
18pub use spl_token_2022_interface::extension::transfer_fee::{
19 TransferFee, TransferFeeConfig,
20};
21
22const TLV_HEADER_LEN: usize = 4;
24
25pub fn mint_extensions<'a>(
28 data: &'a [u8],
29 owner: &Address,
30) -> Result<Option<PodStateWithExtensions<'a, PodMint>>, ProgramError> {
31 if owner == &TOKEN_PROGRAM_ID {
32 return Ok(None);
33 }
34 if owner != &TOKEN_2022_PROGRAM_ID {
35 return Err(ProgramError::IncorrectProgramId);
36 }
37 PodStateWithExtensions::<PodMint>::unpack(data).map(Some)
38}
39
40pub fn extension<'a, V: Extension + Pod>(
46 mint: &'a PodStateWithExtensions<'_, PodMint>,
47) -> Result<Option<&'a V>, ProgramError> {
48 match mint.get_extension::<V>() {
49 Ok(value) => Ok(Some(value)),
50 Err(error) if error == TokenError::ExtensionNotFound.into() => Ok(None),
52 Err(ProgramError::InvalidAccountData)
56 if V::TYPE.get_account_type() == AccountType::Mint
57 && ends_cleanly(mint.get_tlv_data()) =>
58 {
59 Ok(None)
60 }
61 Err(error) => Err(error),
62 }
63}
64
65fn ends_cleanly(tlv: &[u8]) -> bool {
69 let mut rest = tlv;
70 loop {
71 let Some(&[type_low, type_high]) = rest.get(..2) else {
72 return true;
73 };
74 if u16::from_le_bytes([type_low, type_high]) == 0 {
75 return true;
76 }
77 let Some(&[length_low, length_high]) = rest.get(2..TLV_HEADER_LEN)
78 else {
79 return false;
80 };
81 let length = usize::from(u16::from_le_bytes([length_low, length_high]));
82 let Some(next) = TLV_HEADER_LEN
83 .checked_add(length)
84 .and_then(|end| rest.get(end..))
85 else {
86 return false;
87 };
88 rest = next;
89 }
90}
91
92#[must_use]
93pub fn mint_fee(fee: &TransferFee) -> MintFee {
94 MintFee {
95 bps: fee.transfer_fee_basis_points.into(),
96 maximum_fee: fee.maximum_fee.into(),
97 }
98}
99
100pub fn mint_fee_at_epoch(
103 data: &[u8],
104 owner: &Address,
105 epoch: u64,
106) -> Result<Option<MintFee>, ProgramError> {
107 let Some(mint) = mint_extensions(data, owner)? else {
108 return Ok(None);
109 };
110 Ok(extension::<TransferFeeConfig>(&mint)?
111 .map(|config| mint_fee(config.get_epoch_fee(epoch))))
112}
113
114#[cfg(test)]
115mod tests {
116 use super::*;
117
118 use spl_token_2022_interface::extension::transfer_fee::TransferFeeAmount;
119
120 const IS_INITIALIZED_OFFSET: usize = 45;
122
123 const ACCOUNT_TYPE_OFFSET: usize = 165;
125
126 const TLV_START: usize = 166;
127
128 const FEE_TYPE: u16 = 1;
129
130 const FUTURE_TYPE: u16 = 250;
132
133 const NEWER: Option<MintFee> = Some(MintFee {
134 bps: 250,
135 maximum_fee: 2_000,
136 });
137
138 fn entry(epoch: u64, maximum_fee: u64, basis_points: u16) -> TransferFee {
139 TransferFee {
140 epoch: epoch.into(),
141 maximum_fee: maximum_fee.into(),
142 transfer_fee_basis_points: basis_points.into(),
143 }
144 }
145
146 fn fee_payload(authority: Option<Address>) -> Vec<u8> {
147 bytemuck::bytes_of(&TransferFeeConfig {
148 transfer_fee_config_authority: authority.try_into().unwrap(),
149 older_transfer_fee: entry(5, 1_000, 100),
150 newer_transfer_fee: entry(7, 2_000, 250),
151 ..TransferFeeConfig::default()
152 })
153 .to_vec()
154 }
155
156 fn other_extensions() -> Vec<(u16, Vec<u8>)> {
158 vec![(18, vec![0; 64]), (19, vec![0; 90])]
159 }
160
161 fn tlv_image(entries: &[(u16, Vec<u8>)]) -> Vec<u8> {
162 let mut data = vec![0_u8; ACCOUNT_TYPE_OFFSET];
163 data[IS_INITIALIZED_OFFSET] = 1;
164 data.push(1);
166 for (extension_type, payload) in entries {
167 data.extend_from_slice(&extension_type.to_le_bytes());
168 data.extend_from_slice(
169 &u16::try_from(payload.len()).unwrap().to_le_bytes(),
170 );
171 data.extend_from_slice(payload);
172 }
173 data
174 }
175
176 fn fee(data: &[u8]) -> Result<Option<MintFee>, ProgramError> {
177 mint_fee_at_epoch(data, &TOKEN_2022_PROGRAM_ID, 7)
178 }
179
180 #[test]
181 fn the_walk_finds_the_config_wherever_it_sits() {
182 for at in 0..=2 {
183 let mut entries = other_extensions();
184 entries.insert(at, (FEE_TYPE, fee_payload(None)));
185 assert_eq!(fee(&tlv_image(&entries)), Ok(NEWER), "at {at}");
186 }
187 }
188
189 #[test]
190 fn the_owner_picks_the_reader() {
191 let image = tlv_image(&[(FEE_TYPE, fee_payload(None))]);
192 assert_eq!(fee(&image), Ok(NEWER));
193 assert_eq!(mint_fee_at_epoch(&image, &TOKEN_PROGRAM_ID, 7), Ok(None));
194 assert_eq!(
195 mint_fee_at_epoch(&image, &crate::constants::SYSTEM_PROGRAM_ID, 7),
196 Err(ProgramError::IncorrectProgramId)
197 );
198 }
199
200 #[test]
201 fn the_newer_entry_applies_from_its_own_epoch() {
202 let image = tlv_image(&[(FEE_TYPE, fee_payload(None))]);
203 assert_eq!(
204 mint_fee_at_epoch(&image, &TOKEN_2022_PROGRAM_ID, 6),
205 Ok(Some(MintFee {
206 bps: 100,
207 maximum_fee: 1_000,
208 }))
209 );
210 assert_eq!(fee(&image), Ok(NEWER));
211 }
212
213 #[test]
214 fn a_revoked_authority_reads_as_none() {
215 for authority in [None, Some(Address::new_from_array([9; 32]))] {
216 let image = tlv_image(&[(FEE_TYPE, fee_payload(authority))]);
217 let mint = mint_extensions(&image, &TOKEN_2022_PROGRAM_ID)
218 .unwrap()
219 .unwrap();
220 let config =
221 extension::<TransferFeeConfig>(&mint).unwrap().unwrap();
222 assert_eq!(
223 Option::from(config.transfer_fee_config_authority),
224 authority
225 );
226 }
227 }
228
229 #[test]
230 fn a_bare_mint_has_no_tlv_region_to_walk() {
231 let mut bare = tlv_image(&[]);
232 bare.truncate(82);
233 assert_eq!(fee(&bare), Ok(None));
234 bare[IS_INITIALIZED_OFFSET] = 0;
235 assert_eq!(fee(&bare), Err(ProgramError::UninitializedAccount));
236 }
237
238 #[test]
241 fn a_region_that_ends_cleanly_carries_no_config() {
242 for tail in [
243 vec![],
244 vec![0],
245 vec![0; 2],
246 vec![0; 3],
247 vec![0; 512],
248 vec![7],
249 ] {
250 let mut image = tlv_image(&other_extensions());
251 image.extend_from_slice(&tail);
252 assert_eq!(fee(&image), Ok(None), "tail {tail:?}");
253 }
254 }
255
256 #[test]
259 fn a_cut_entry_past_the_config_is_not_read() {
260 let mut image = tlv_image(&[(FEE_TYPE, fee_payload(None))]);
261 image.extend_from_slice(&[18, 0, 64, 0]);
262 assert_eq!(fee(&image), Ok(NEWER));
263 }
264
265 #[test]
268 fn an_account_extension_is_not_read_off_a_mint() {
269 let image = tlv_image(&other_extensions());
270 let mint = mint_extensions(&image, &TOKEN_2022_PROGRAM_ID)
271 .unwrap()
272 .unwrap();
273 assert_eq!(
274 extension::<TransferFeeAmount>(&mint),
275 Err(ProgramError::InvalidAccountData)
276 );
277 }
278
279 #[test]
282 fn a_future_extension_type_leaves_the_mint_readable() {
283 let future = (FUTURE_TYPE, vec![0; 16]);
284 let fee_free = tlv_image(&[(18, vec![0; 64]), future.clone()]);
285 let mint = mint_extensions(&fee_free, &TOKEN_2022_PROGRAM_ID)
286 .unwrap()
287 .unwrap();
288 assert_eq!(
289 mint.get_extension_types(),
290 Err(ProgramError::InvalidAccountData)
291 );
292 assert_eq!(extension::<TransferFeeConfig>(&mint), Ok(None));
293 let fee_after = tlv_image(&[future, (FEE_TYPE, fee_payload(None))]);
294 assert_eq!(fee(&fee_after), Ok(NEWER));
295 }
296
297 #[test]
298 fn a_cut_region_is_malformed() {
299 let mut short = tlv_image(&[]);
300 short.truncate(ACCOUNT_TYPE_OFFSET);
301 let mut header = tlv_image(&other_extensions());
302 header.truncate(TLV_START + 4 + 64 + 2);
303 let mut overrun = tlv_image(&[(18, vec![0; 64])]);
304 overrun.truncate(TLV_START + 4 + 32);
305 let mut payload = tlv_image(&other_extensions());
306 payload.extend_from_slice(&[18, 0, 64, 0]);
307 let mut config = tlv_image(&[(FEE_TYPE, fee_payload(None))]);
308 config.pop();
309 for (label, image) in [
310 ("no account type", short),
311 ("cut header", header),
312 ("length past the end", overrun),
313 ("payload never arrives", payload),
314 ("config payload cut", config),
315 ] {
316 assert_eq!(
317 fee(&image),
318 Err(ProgramError::InvalidAccountData),
319 "{label}"
320 );
321 }
322 }
323
324 #[test]
325 fn a_config_payload_of_the_wrong_length_is_malformed() {
326 for length in [100_usize, 116] {
327 let image = tlv_image(&[(FEE_TYPE, vec![0; length])]);
328 assert_eq!(
329 fee(&image),
330 Err(ProgramError::InvalidArgument),
331 "payload length {length}"
332 );
333 }
334 }
335
336 #[test]
337 fn a_non_mint_account_type_is_malformed() {
338 let mut image = tlv_image(&[(FEE_TYPE, fee_payload(None))]);
339 image[ACCOUNT_TYPE_OFFSET] = 2;
341 assert_eq!(fee(&image), Err(ProgramError::InvalidAccountData));
342 }
343}