1use std::fmt;
6
7use solana_address::{Address, MAX_SEEDS};
8use solana_instruction::AccountMeta;
9use spl_discriminator::{ArrayDiscriminator, SplDiscriminate};
10use spl_tlv_account_resolution::account::ExtraAccountMeta;
11use spl_tlv_account_resolution::pubkey_data::PubkeyData;
12use spl_tlv_account_resolution::seeds::Seed;
13use spl_tlv_account_resolution::solana_program_error::ProgramError;
14use spl_tlv_account_resolution::state::ExtraAccountMetaList;
15use spl_token_2022_interface::extension::pausable::PausableConfig;
16use spl_token_2022_interface::extension::transfer_hook::TransferHook;
17use spl_type_length_value::error::TlvError;
18use spl_type_length_value::state::TlvStateBorrowed;
19
20use crate::math::amm::MintFee;
21use crate::transfer_fee::{
22 TransferFeeConfig, extension, mint_extensions, mint_fee,
23};
24use crate::utils::{StandInTrade, TradeAccounts, TradeDirection};
25
26pub const EXECUTE_DISCRIMINATOR: [u8; 8] =
29 [105, 37, 101, 197, 75, 251, 102, 26];
30
31pub const VALIDATION_SEED: &[u8] = b"extra-account-metas";
34
35const SOURCE_INDEX: usize = 0;
37const DESTINATION_INDEX: usize = 2;
38const VALIDATION_INDEX: usize = 4;
39const EXECUTE_ACCOUNTS: usize = 5;
40
41const TOKEN_ACCOUNT_KNOWN_LEN: usize = 64;
43
44const PUBKEY_LEN: usize = 32;
45
46const META_FIXED: u8 = 0;
48const META_HOOK_PDA: u8 = 1;
49const META_PUBKEY_DATA: u8 = 2;
50const META_EXTERNAL_PDA: u8 = 128;
51
52#[derive(Clone, Copy, Debug, PartialEq, Eq)]
53pub struct AccountView<'a> {
54 pub owner: &'a Address,
55 pub data: &'a [u8],
56}
57
58#[must_use]
59pub fn find_validation_address(mint: &Address, program: &Address) -> Address {
60 Address::find_program_address(&[VALIDATION_SEED, mint.as_ref()], program).0
61}
62
63#[derive(Clone, Copy, Debug, PartialEq)]
64pub struct MintState {
65 pub token_program: Address,
66 pub transfer_fee: Option<TransferFeeConfig>,
67 pub transfer_hook: Option<Address>,
70 pub paused: bool,
72}
73
74impl MintState {
75 pub fn load(mint: AccountView<'_>) -> Result<Self, ProgramError> {
76 let token_program = *mint.owner;
77 let Some(state) = mint_extensions(mint.data, &token_program)? else {
78 return Ok(Self {
79 token_program,
80 transfer_fee: None,
81 transfer_hook: None,
82 paused: false,
83 });
84 };
85 Ok(Self {
86 token_program,
87 transfer_fee: extension::<TransferFeeConfig>(&state)?.copied(),
88 transfer_hook: extension::<TransferHook>(&state)?
90 .and_then(|hook| Option::from(hook.program_id)),
91 paused: extension::<PausableConfig>(&state)?
92 .is_some_and(|config| config.paused.into()),
93 })
94 }
95
96 #[must_use]
97 pub fn fee_at(&self, epoch: u64) -> Option<MintFee> {
98 self.transfer_fee
99 .as_ref()
100 .map(|config| mint_fee(config.get_epoch_fee(epoch)))
101 }
102}
103
104#[derive(Clone, Copy, Debug, PartialEq, Eq)]
105pub enum Unloaded {
106 Refuse,
107 StandIn,
110}
111
112#[derive(Clone, Copy, Debug, PartialEq, Eq)]
113#[non_exhaustive]
114pub enum Unresolvable {
115 MalformedList,
116 NoExecuteEntry,
117 ListTooLong,
119 MalformedExtra,
120 MalformedSeed,
121 TooManySeeds,
123 SeedTooLong,
124 ExtraSigns,
126 ReadsTransferAmount,
129 ReadsPastInstructionData,
132 NamesUnresolvedAccount,
133 ReadsSwapTimeData,
134 Clash(Clash),
136}
137
138impl fmt::Display for Unresolvable {
139 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
140 f.write_str(match self {
141 Self::MalformedList => "malformed list",
142 Self::NoExecuteEntry => "no Execute entry in the list",
143 Self::ListTooLong => "list too long",
144 Self::MalformedExtra => "malformed extra",
145 Self::MalformedSeed => "malformed seed",
146 Self::TooManySeeds => "a PDA has more seeds than an address allows",
147 Self::SeedTooLong => "seed longer than 32 bytes",
148 Self::ExtraSigns => "an extra must sign",
149 Self::ReadsTransferAmount => "a seed reads the transfer amount",
150 Self::ReadsPastInstructionData => {
151 "an extra reads a pubkey past the end of the hook's instruction data"
152 }
153 Self::NamesUnresolvedAccount => {
154 "an extra names an account not yet resolved"
155 }
156 Self::ReadsSwapTimeData => {
157 "an extra reads account data only the swap can see"
158 }
159 Self::Clash(clash) => return clash.fmt(f),
160 })
161 }
162}
163
164#[derive(Clone, Copy, Debug, PartialEq, Eq)]
165#[non_exhaustive]
166pub enum Clash {
167 SignerNotForwarded,
168 WritableDuplicate,
171 ExtraUnresolved,
172}
173
174impl fmt::Display for Clash {
175 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
176 f.write_str(match self {
177 Self::SignerNotForwarded => {
178 "transfer hook needs a signer the programs do not forward"
179 }
180 Self::WritableDuplicate => {
181 "transfer hook repeats an account of the trade, writable"
182 }
183 Self::ExtraUnresolved => "transfer hook extra did not resolve",
184 })
185 }
186}
187
188#[derive(Clone, Copy, Debug, PartialEq, Eq)]
189#[non_exhaustive]
190pub enum HookError {
191 NotLoaded,
192 NoList,
195 Unresolvable(Unresolvable),
196 Clash(Clash),
198}
199
200impl fmt::Display for HookError {
201 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
202 match self {
203 Self::NotLoaded => write!(f, "transfer hook accounts not loaded"),
204 Self::NoList => {
205 write!(f, "transfer hook has no ExtraAccountMetaList")
206 }
207 Self::Unresolvable(reason) => write!(
208 f,
209 "transfer hook cannot be resolved ahead of the swap: {reason}"
210 ),
211 Self::Clash(clash) => clash.fmt(f),
212 }
213 }
214}
215
216impl std::error::Error for HookError {}
217
218#[derive(Clone, Debug, PartialEq, Eq)]
219pub struct Hook {
220 program: Address,
221 validation: Address,
222 list: HookList,
223}
224
225#[derive(Clone, Debug, PartialEq, Eq)]
226enum HookList {
227 Pending,
230 Absent,
232 Ready(Vec<u8>),
234 Unsupported(Unresolvable),
235}
236
237impl Hook {
238 #[must_use]
239 pub const fn program(&self) -> &Address {
240 &self.program
241 }
242
243 #[must_use]
244 pub const fn validation(&self) -> &Address {
245 &self.validation
246 }
247
248 pub const fn check_routable(
251 &self,
252 unloaded: Unloaded,
253 ) -> Result<(), HookError> {
254 match &self.list {
255 HookList::Ready(_) => Ok(()),
256 HookList::Pending => match unloaded {
257 Unloaded::StandIn => Ok(()),
258 Unloaded::Refuse => Err(HookError::NotLoaded),
259 },
260 HookList::Absent => Err(HookError::NoList),
261 HookList::Unsupported(reason) => {
262 Err(HookError::Unresolvable(*reason))
263 }
264 }
265 }
266
267 fn load(planned: Planned, account: Option<AccountView<'_>>) -> Self {
270 let list = account
272 .filter(|account| *account.owner == planned.program)
273 .map_or(planned.missing, |account| {
274 match screen_list(account.data) {
275 Ok(()) => HookList::Ready(account.data.to_vec()),
276 Err(reason) => HookList::Unsupported(reason),
277 }
278 });
279 Self {
280 program: planned.program,
281 validation: planned.validation,
282 list,
283 }
284 }
285
286 fn leg_accounts(
288 &self,
289 leg: &Leg,
290 unloaded: Unloaded,
291 ) -> Result<Vec<AccountMeta>, HookError> {
292 let data = match &self.list {
293 HookList::Ready(data) => data,
294 HookList::Pending if unloaded == Unloaded::StandIn => {
295 return Ok(vec![
296 AccountMeta::new_readonly(self.program, false),
297 AccountMeta::new_readonly(self.validation, false),
298 ]);
299 }
300 HookList::Pending | HookList::Absent | HookList::Unsupported(_) => {
301 return Err(self
302 .check_routable(Unloaded::Refuse)
303 .err()
304 .unwrap_or(HookError::NotLoaded));
305 }
306 };
307 let unresolved = HookError::Clash(Clash::ExtraUnresolved);
308 let extras = unpack_list(data).map_err(|_| unresolved)?;
309 let source = leg.token_account_prefix(&leg.owner);
310 let destination = leg.token_account_prefix(&leg.destination_owner);
311 let mut keys = vec![
312 leg.source,
313 leg.mint,
314 leg.destination,
315 leg.owner,
316 self.validation,
317 ];
318 let mut metas = Vec::with_capacity(extras.len());
319 for extra in &extras {
320 let resolved = extra
321 .resolve(&EXECUTE_DISCRIMINATOR, &self.program, |index| {
322 let data = match index {
323 SOURCE_INDEX => Some(source.as_slice()),
324 DESTINATION_INDEX => Some(destination.as_slice()),
325 VALIDATION_INDEX => Some(data.as_slice()),
326 _ => None,
327 };
328 keys.get(index).map(|key| (key, data))
329 })
330 .map_err(|_| unresolved)?;
331 keys.push(resolved.pubkey);
332 metas.push(AccountMeta {
333 pubkey: resolved.pubkey,
334 is_signer: false,
335 is_writable: resolved.is_writable,
336 });
337 }
338 metas.push(AccountMeta::new_readonly(self.program, false));
339 metas.push(AccountMeta::new_readonly(self.validation, false));
340 Ok(metas)
341 }
342}
343
344struct Leg {
346 source: Address,
347 mint: Address,
348 destination: Address,
349 owner: Address,
350 destination_owner: Address,
351}
352
353impl Leg {
354 fn token_account_prefix(&self, owner: &Address) -> [u8; 64] {
355 let mut prefix = [0; TOKEN_ACCOUNT_KNOWN_LEN];
356 let (mint, rest) = prefix.split_at_mut(PUBKEY_LEN);
357 mint.copy_from_slice(self.mint.as_ref());
358 rest.copy_from_slice(owner.as_ref());
359 prefix
360 }
361}
362
363struct Execute;
364
365impl SplDiscriminate for Execute {
366 const SPL_DISCRIMINATOR: ArrayDiscriminator =
367 ArrayDiscriminator::new(EXECUTE_DISCRIMINATOR);
368}
369
370fn unpack_list(data: &[u8]) -> Result<Vec<ExtraAccountMeta>, Unresolvable> {
371 let list_error = |error: ProgramError| {
372 if error == ProgramError::from(TlvError::TypeNotFound) {
373 Unresolvable::NoExecuteEntry
374 } else {
375 Unresolvable::MalformedList
376 }
377 };
378 let state = TlvStateBorrowed::unpack(data).map_err(list_error)?;
379 let extras = ExtraAccountMetaList::unpack_with_tlv_state::<Execute>(&state)
380 .map_err(list_error)?;
381 Ok(extras.iter().copied().collect())
382}
383
384fn screen_list(data: &[u8]) -> Result<(), Unresolvable> {
386 for (index, extra) in unpack_list(data)?.iter().enumerate() {
387 let position = index
388 .checked_add(EXECUTE_ACCOUNTS)
389 .filter(|position| u8::try_from(*position).is_ok())
390 .ok_or(Unresolvable::ListTooLong)?;
391 screen_extra(extra, position)?;
392 }
393 Ok(())
394}
395
396fn screen_extra(
397 extra: &ExtraAccountMeta,
398 position: usize,
399) -> Result<(), Unresolvable> {
400 if bool::from(extra.is_signer) {
401 return Err(Unresolvable::ExtraSigns);
402 }
403 match extra.discriminator {
404 META_FIXED => Ok(()),
405 META_HOOK_PDA => screen_seeds(&extra.address_config, position),
406 META_PUBKEY_DATA => match PubkeyData::unpack(&extra.address_config)
407 .map_err(|_| Unresolvable::MalformedExtra)?
408 {
409 PubkeyData::AccountData {
410 account_index,
411 data_index,
412 } => readable(
413 DataRead {
414 account: account_index,
415 offset: data_index,
416 length: PUBKEY_LEN,
417 },
418 position,
419 ),
420 PubkeyData::InstructionData { .. } => {
421 Err(Unresolvable::ReadsPastInstructionData)
422 }
423 PubkeyData::Uninitialized => Err(Unresolvable::MalformedExtra),
424 },
425 external => match external.checked_sub(META_EXTERNAL_PDA) {
426 Some(program) => {
427 prior(program, position)?;
428 screen_seeds(&extra.address_config, position)
429 }
430 None => Err(Unresolvable::MalformedExtra),
431 },
432 }
433}
434
435fn screen_seeds(
436 config: &[u8; 32],
437 position: usize,
438) -> Result<(), Unresolvable> {
439 let seeds = Seed::unpack_address_config(config)
440 .map_err(|_| Unresolvable::MalformedSeed)?;
441 if seeds.len() >= MAX_SEEDS {
443 return Err(Unresolvable::TooManySeeds);
444 }
445 for seed in &seeds {
446 match *seed {
447 Seed::Uninitialized | Seed::Literal { .. } => {}
449 Seed::InstructionData { index, length } => {
450 let end = usize::from(index)
451 .checked_add(usize::from(length))
452 .ok_or(Unresolvable::MalformedSeed)?;
453 if end > EXECUTE_DISCRIMINATOR.len() {
454 return Err(Unresolvable::ReadsTransferAmount);
455 }
456 }
457 Seed::AccountKey { index } => prior(index, position)?,
458 Seed::AccountData {
459 account_index,
460 data_index,
461 length,
462 } => {
463 let length = usize::from(length);
464 if length > PUBKEY_LEN {
465 return Err(Unresolvable::SeedTooLong);
466 }
467 readable(
468 DataRead {
469 account: account_index,
470 offset: data_index,
471 length,
472 },
473 position,
474 )?;
475 }
476 }
477 }
478 Ok(())
479}
480
481fn prior(index: u8, position: usize) -> Result<(), Unresolvable> {
483 if usize::from(index) < position {
484 Ok(())
485 } else {
486 Err(Unresolvable::NamesUnresolvedAccount)
487 }
488}
489
490struct DataRead {
492 account: u8,
493 offset: u8,
494 length: usize,
495}
496
497fn readable(read: DataRead, position: usize) -> Result<(), Unresolvable> {
500 prior(read.account, position)?;
501 let end = usize::from(read.offset)
502 .checked_add(read.length)
503 .ok_or(Unresolvable::MalformedSeed)?;
504 match usize::from(read.account) {
505 SOURCE_INDEX | DESTINATION_INDEX if end <= TOKEN_ACCOUNT_KNOWN_LEN => {
506 Ok(())
507 }
508 VALIDATION_INDEX => Ok(()),
511 _ => Err(Unresolvable::ReadsSwapTimeData),
512 }
513}
514
515#[derive(Clone, Copy, Debug)]
516pub struct HookSwap<'a> {
517 pub trade: &'a TradeAccounts,
518 pub direction: TradeDirection,
519 pub fixed: &'a [AccountMeta],
521}
522
523#[derive(Clone, Debug, Default, PartialEq, Eq)]
524pub struct MarketHooks {
525 base: Option<Hook>,
526 quote: Option<Hook>,
527}
528
529#[derive(Clone, Debug, PartialEq, Eq)]
530struct Planned {
531 program: Address,
532 validation: Address,
533 missing: HookList,
535}
536
537#[derive(Clone, Debug, PartialEq, Eq)]
538pub struct HookPlan {
539 base: Option<Planned>,
540 quote: Option<Planned>,
541}
542
543impl HookPlan {
544 pub fn accounts(&self) -> impl Iterator<Item = Address> + '_ {
545 [&self.base, &self.quote]
546 .into_iter()
547 .flatten()
548 .map(|planned| planned.validation)
549 }
550
551 pub fn resolve<'a>(
554 self,
555 accounts: impl Fn(&Address) -> Option<AccountView<'a>>,
556 ) -> MarketHooks {
557 let load = |planned: Planned| {
558 let account = accounts(&planned.validation);
559 Hook::load(planned, account)
560 };
561 MarketHooks {
562 base: self.base.map(load),
563 quote: self.quote.map(load),
564 }
565 }
566}
567
568impl MarketHooks {
569 #[must_use]
570 pub const fn base(&self) -> Option<&Hook> {
571 self.base.as_ref()
572 }
573
574 #[must_use]
575 pub const fn quote(&self) -> Option<&Hook> {
576 self.quote.as_ref()
577 }
578
579 #[must_use]
582 pub fn plan(
583 &self,
584 base: (&Address, &MintState),
585 quote: (&Address, &MintState),
586 ) -> HookPlan {
587 let plan = |previous: Option<&Hook>,
588 (mint, state): (&Address, &MintState)| {
589 state.transfer_hook.map(|program| {
590 let asked = previous.filter(|hook| hook.program == program);
591 Planned {
592 program,
593 validation: asked.map_or_else(
594 || find_validation_address(mint, &program),
595 |hook| hook.validation,
596 ),
597 missing: match asked.map(|hook| &hook.list) {
598 Some(HookList::Pending | HookList::Absent) => {
599 HookList::Absent
600 }
601 Some(HookList::Ready(_) | HookList::Unsupported(_))
602 | None => HookList::Pending,
603 },
604 }
605 })
606 };
607 HookPlan {
608 base: plan(self.base(), base),
609 quote: plan(self.quote(), quote),
610 }
611 }
612
613 #[must_use]
617 pub fn screen(self, trade: &StandInTrade, fixed: &[AccountMeta]) -> Self {
618 let screen = |hook: Hook, alone: &Self| {
619 if hook.check_routable(Unloaded::Refuse).is_err() {
620 return hook;
621 }
622 let clash = [TradeDirection::Buy, TradeDirection::Sell]
623 .into_iter()
624 .find_map(|direction| {
625 let swap = HookSwap {
626 trade,
627 direction,
628 fixed,
629 };
630 alone.swap_accounts(swap, Unloaded::Refuse).err()
631 });
632 let reason = match clash {
633 Some(HookError::Unresolvable(reason)) => reason,
634 Some(HookError::Clash(clash)) => Unresolvable::Clash(clash),
635 None | Some(HookError::NotLoaded | HookError::NoList) => {
637 return hook;
638 }
639 };
640 Hook {
641 list: HookList::Unsupported(reason),
642 ..hook
643 }
644 };
645 Self {
646 base: self.base.map(|hook| {
647 let alone = Self {
648 base: Some(hook.clone()),
649 quote: None,
650 };
651 screen(hook, &alone)
652 }),
653 quote: self.quote.map(|hook| {
654 let alone = Self {
655 base: None,
656 quote: Some(hook.clone()),
657 };
658 screen(hook, &alone)
659 }),
660 }
661 }
662
663 fn iter(&self) -> impl Iterator<Item = &Hook> {
664 [&self.base, &self.quote].into_iter().flatten()
665 }
666
667 pub fn validation_accounts(&self) -> impl Iterator<Item = Address> + '_ {
670 self.iter().map(|hook| hook.validation)
671 }
672
673 pub fn check_routable(&self, unloaded: Unloaded) -> Result<(), HookError> {
674 self.iter()
675 .try_for_each(|hook| hook.check_routable(unloaded))
676 }
677
678 #[must_use]
679 pub fn programs(&self) -> Vec<Address> {
680 let mut programs: Vec<Address> = Vec::new();
681 for hook in self.iter() {
682 if !programs.contains(&hook.program) {
683 programs.push(hook.program);
684 }
685 }
686 programs
687 }
688
689 pub fn swap_accounts(
691 &self,
692 swap: HookSwap<'_>,
693 unloaded: Unloaded,
694 ) -> Result<Vec<AccountMeta>, HookError> {
695 let HookSwap {
696 trade,
697 direction,
698 fixed,
699 } = swap;
700 let into_vault = |source: Address, mint: Address, vault: Address| Leg {
701 source,
702 mint,
703 destination: vault,
704 owner: trade.user,
705 destination_owner: trade.market,
706 };
707 let out_of_vault =
708 |vault: Address, mint: Address, destination: Address| Leg {
709 source: vault,
710 mint,
711 destination,
712 owner: trade.market,
713 destination_owner: trade.user,
714 };
715 let legs = match direction {
716 TradeDirection::Buy => [
717 (
718 self.quote(),
719 into_vault(
720 trade.user_quote_account,
721 trade.quote_mint,
722 trade.quote_vault,
723 ),
724 ),
725 (
726 self.base(),
727 out_of_vault(
728 trade.base_vault,
729 trade.base_mint,
730 trade.user_base_account,
731 ),
732 ),
733 ],
734 TradeDirection::Sell => [
735 (
736 self.base(),
737 into_vault(
738 trade.user_base_account,
739 trade.base_mint,
740 trade.base_vault,
741 ),
742 ),
743 (
744 self.quote(),
745 out_of_vault(
746 trade.quote_vault,
747 trade.quote_mint,
748 trade.user_quote_account,
749 ),
750 ),
751 ],
752 };
753 trailing_accounts(&legs, fixed, unloaded)
754 }
755}
756
757fn trailing_accounts(
763 legs: &[(Option<&Hook>, Leg)],
764 fixed: &[AccountMeta],
765 unloaded: Unloaded,
766) -> Result<Vec<AccountMeta>, HookError> {
767 let mut merged: Vec<AccountMeta> = Vec::new();
768 for (hook, leg) in legs {
769 let Some(hook) = hook else {
770 continue;
771 };
772 for meta in hook.leg_accounts(leg, unloaded)? {
773 if let Some(clash) =
774 fixed.iter().find(|account| account.pubkey == meta.pubkey)
775 {
776 if clash.is_signer {
777 return Err(HookError::Clash(Clash::SignerNotForwarded));
778 }
779 if clash.is_writable || meta.is_writable {
780 return Err(HookError::Clash(Clash::WritableDuplicate));
781 }
782 }
783 match merged
784 .iter_mut()
785 .find(|account| account.pubkey == meta.pubkey)
786 {
787 Some(existing) => existing.is_writable |= meta.is_writable,
788 None => merged.push(meta),
789 }
790 }
791 }
792 Ok(merged)
793}
794
795#[cfg(test)]
796mod tests;