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,
125 ReadsTransferAmount,
128 ReadsPastInstructionData,
131 NamesUnresolvedAccount,
132 ReadsSwapTimeData,
133 Clash(Clash),
135}
136
137impl fmt::Display for Unresolvable {
138 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
139 f.write_str(match self {
140 Self::MalformedList => "malformed list",
141 Self::NoExecuteEntry => "no Execute entry in the list",
142 Self::ListTooLong => "list too long",
143 Self::MalformedExtra => "malformed extra",
144 Self::MalformedSeed => "malformed seed",
145 Self::TooManySeeds => "a PDA has more seeds than an address allows",
146 Self::SeedTooLong => "seed longer than 32 bytes",
147 Self::ExtraSigns => "an extra must sign",
148 Self::ReadsTransferAmount => "a seed reads the transfer amount",
149 Self::ReadsPastInstructionData => {
150 "an extra reads a pubkey past the end of the hook's instruction data"
151 }
152 Self::NamesUnresolvedAccount => {
153 "an extra names an account not yet resolved"
154 }
155 Self::ReadsSwapTimeData => {
156 "an extra reads account data only the swap can see"
157 }
158 Self::Clash(clash) => return clash.fmt(f),
159 })
160 }
161}
162
163#[derive(Clone, Copy, Debug, PartialEq, Eq)]
164#[non_exhaustive]
165pub enum Clash {
166 WritableDuplicate,
169 ExtraUnresolved,
170}
171
172impl fmt::Display for Clash {
173 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
174 f.write_str(match self {
175 Self::WritableDuplicate => {
176 "transfer hook repeats an account of the trade, writable"
177 }
178 Self::ExtraUnresolved => "transfer hook extra did not resolve",
179 })
180 }
181}
182
183#[derive(Clone, Copy, Debug, PartialEq, Eq)]
184#[non_exhaustive]
185pub enum HookError {
186 NotLoaded,
187 NoList,
190 Unresolvable(Unresolvable),
191 Clash(Clash),
193}
194
195impl fmt::Display for HookError {
196 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
197 match self {
198 Self::NotLoaded => write!(f, "transfer hook accounts not loaded"),
199 Self::NoList => {
200 write!(f, "transfer hook has no ExtraAccountMetaList")
201 }
202 Self::Unresolvable(reason) => write!(
203 f,
204 "transfer hook cannot be resolved ahead of the swap: {reason}"
205 ),
206 Self::Clash(clash) => clash.fmt(f),
207 }
208 }
209}
210
211impl std::error::Error for HookError {}
212
213#[derive(Clone, Debug, PartialEq, Eq)]
214pub struct Hook {
215 program: Address,
216 validation: Address,
217 list: HookList,
218}
219
220#[derive(Clone, Debug, PartialEq, Eq)]
221enum HookList {
222 Pending,
225 Absent,
227 Ready(Vec<u8>),
229 Unsupported(Unresolvable),
230}
231
232impl Hook {
233 #[must_use]
234 pub const fn program(&self) -> &Address {
235 &self.program
236 }
237
238 #[must_use]
239 pub const fn validation(&self) -> &Address {
240 &self.validation
241 }
242
243 pub const fn check_routable(
246 &self,
247 unloaded: Unloaded,
248 ) -> Result<(), HookError> {
249 match &self.list {
250 HookList::Ready(_) => Ok(()),
251 HookList::Pending => match unloaded {
252 Unloaded::StandIn => Ok(()),
253 Unloaded::Refuse => Err(HookError::NotLoaded),
254 },
255 HookList::Absent => Err(HookError::NoList),
256 HookList::Unsupported(reason) => {
257 Err(HookError::Unresolvable(*reason))
258 }
259 }
260 }
261
262 fn load(planned: Planned, account: Option<AccountView<'_>>) -> Self {
265 let list = account
267 .filter(|account| *account.owner == planned.program)
268 .map_or(planned.missing, |account| {
269 match screen_list(account.data) {
270 Ok(()) => HookList::Ready(account.data.to_vec()),
271 Err(reason) => HookList::Unsupported(reason),
272 }
273 });
274 Self {
275 program: planned.program,
276 validation: planned.validation,
277 list,
278 }
279 }
280
281 fn leg_accounts(
283 &self,
284 leg: &Leg,
285 unloaded: Unloaded,
286 ) -> Result<Vec<AccountMeta>, HookError> {
287 let data = match &self.list {
288 HookList::Ready(data) => data,
289 HookList::Pending if unloaded == Unloaded::StandIn => {
290 return Ok(vec![
291 AccountMeta::new_readonly(self.program, false),
292 AccountMeta::new_readonly(self.validation, false),
293 ]);
294 }
295 HookList::Pending | HookList::Absent | HookList::Unsupported(_) => {
296 return Err(self
297 .check_routable(Unloaded::Refuse)
298 .err()
299 .unwrap_or(HookError::NotLoaded));
300 }
301 };
302 let unresolved = HookError::Clash(Clash::ExtraUnresolved);
303 let extras = unpack_list(data).map_err(|_| unresolved)?;
304 let source = leg.token_account_prefix(&leg.owner);
305 let destination = leg.token_account_prefix(&leg.destination_owner);
306 let mut keys = vec![
307 leg.source,
308 leg.mint,
309 leg.destination,
310 leg.owner,
311 self.validation,
312 ];
313 let mut metas = Vec::with_capacity(extras.len());
314 for extra in &extras {
315 let resolved = extra
316 .resolve(&EXECUTE_DISCRIMINATOR, &self.program, |index| {
317 let data = match index {
318 SOURCE_INDEX => Some(source.as_slice()),
319 DESTINATION_INDEX => Some(destination.as_slice()),
320 VALIDATION_INDEX => Some(data.as_slice()),
321 _ => None,
322 };
323 keys.get(index).map(|key| (key, data))
324 })
325 .map_err(|_| unresolved)?;
326 keys.push(resolved.pubkey);
327 metas.push(AccountMeta {
328 pubkey: resolved.pubkey,
329 is_signer: false,
330 is_writable: resolved.is_writable,
331 });
332 }
333 metas.push(AccountMeta::new_readonly(self.program, false));
334 metas.push(AccountMeta::new_readonly(self.validation, false));
335 Ok(metas)
336 }
337}
338
339struct Leg {
341 source: Address,
342 mint: Address,
343 destination: Address,
344 owner: Address,
345 destination_owner: Address,
346}
347
348impl Leg {
349 fn token_account_prefix(&self, owner: &Address) -> [u8; 64] {
350 let mut prefix = [0; TOKEN_ACCOUNT_KNOWN_LEN];
351 let (mint, rest) = prefix.split_at_mut(PUBKEY_LEN);
352 mint.copy_from_slice(self.mint.as_ref());
353 rest.copy_from_slice(owner.as_ref());
354 prefix
355 }
356}
357
358struct Execute;
359
360impl SplDiscriminate for Execute {
361 const SPL_DISCRIMINATOR: ArrayDiscriminator =
362 ArrayDiscriminator::new(EXECUTE_DISCRIMINATOR);
363}
364
365fn unpack_list(data: &[u8]) -> Result<Vec<ExtraAccountMeta>, Unresolvable> {
366 let list_error = |error: ProgramError| {
367 if error == ProgramError::from(TlvError::TypeNotFound) {
368 Unresolvable::NoExecuteEntry
369 } else {
370 Unresolvable::MalformedList
371 }
372 };
373 let state = TlvStateBorrowed::unpack(data).map_err(list_error)?;
374 let extras = ExtraAccountMetaList::unpack_with_tlv_state::<Execute>(&state)
375 .map_err(list_error)?;
376 Ok(extras.iter().copied().collect())
377}
378
379fn screen_list(data: &[u8]) -> Result<(), Unresolvable> {
381 for (index, extra) in unpack_list(data)?.iter().enumerate() {
382 let position = index
383 .checked_add(EXECUTE_ACCOUNTS)
384 .filter(|position| u8::try_from(*position).is_ok())
385 .ok_or(Unresolvable::ListTooLong)?;
386 screen_extra(extra, position)?;
387 }
388 Ok(())
389}
390
391fn screen_extra(
392 extra: &ExtraAccountMeta,
393 position: usize,
394) -> Result<(), Unresolvable> {
395 if bool::from(extra.is_signer) {
396 return Err(Unresolvable::ExtraSigns);
397 }
398 match extra.discriminator {
399 META_FIXED => Ok(()),
400 META_HOOK_PDA => screen_seeds(&extra.address_config, position),
401 META_PUBKEY_DATA => match PubkeyData::unpack(&extra.address_config)
402 .map_err(|_| Unresolvable::MalformedExtra)?
403 {
404 PubkeyData::AccountData {
405 account_index,
406 data_index,
407 } => readable(
408 DataRead {
409 account: account_index,
410 offset: data_index,
411 length: PUBKEY_LEN,
412 },
413 position,
414 ),
415 PubkeyData::InstructionData { .. } => {
416 Err(Unresolvable::ReadsPastInstructionData)
417 }
418 PubkeyData::Uninitialized => Err(Unresolvable::MalformedExtra),
419 },
420 external => match external.checked_sub(META_EXTERNAL_PDA) {
421 Some(program) => {
422 prior(program, position)?;
423 screen_seeds(&extra.address_config, position)
424 }
425 None => Err(Unresolvable::MalformedExtra),
426 },
427 }
428}
429
430fn screen_seeds(
431 config: &[u8; 32],
432 position: usize,
433) -> Result<(), Unresolvable> {
434 let seeds = Seed::unpack_address_config(config)
435 .map_err(|_| Unresolvable::MalformedSeed)?;
436 if seeds.len() >= MAX_SEEDS {
438 return Err(Unresolvable::TooManySeeds);
439 }
440 for seed in &seeds {
441 match *seed {
442 Seed::Uninitialized | Seed::Literal { .. } => {}
444 Seed::InstructionData { index, length } => {
445 let end = usize::from(index)
446 .checked_add(usize::from(length))
447 .ok_or(Unresolvable::MalformedSeed)?;
448 if end > EXECUTE_DISCRIMINATOR.len() {
449 return Err(Unresolvable::ReadsTransferAmount);
450 }
451 }
452 Seed::AccountKey { index } => prior(index, position)?,
453 Seed::AccountData {
454 account_index,
455 data_index,
456 length,
457 } => {
458 let length = usize::from(length);
459 if length > PUBKEY_LEN {
460 return Err(Unresolvable::SeedTooLong);
461 }
462 readable(
463 DataRead {
464 account: account_index,
465 offset: data_index,
466 length,
467 },
468 position,
469 )?;
470 }
471 }
472 }
473 Ok(())
474}
475
476fn prior(index: u8, position: usize) -> Result<(), Unresolvable> {
478 if usize::from(index) < position {
479 Ok(())
480 } else {
481 Err(Unresolvable::NamesUnresolvedAccount)
482 }
483}
484
485struct DataRead {
487 account: u8,
488 offset: u8,
489 length: usize,
490}
491
492fn readable(read: DataRead, position: usize) -> Result<(), Unresolvable> {
495 prior(read.account, position)?;
496 let end = usize::from(read.offset)
497 .checked_add(read.length)
498 .ok_or(Unresolvable::MalformedSeed)?;
499 match usize::from(read.account) {
500 SOURCE_INDEX | DESTINATION_INDEX if end <= TOKEN_ACCOUNT_KNOWN_LEN => {
501 Ok(())
502 }
503 VALIDATION_INDEX => Ok(()),
506 _ => Err(Unresolvable::ReadsSwapTimeData),
507 }
508}
509
510#[derive(Clone, Copy, Debug)]
511pub struct HookSwap<'a> {
512 pub trade: &'a TradeAccounts,
513 pub direction: TradeDirection,
514 pub fixed: &'a [AccountMeta],
516}
517
518#[derive(Clone, Debug, Default, PartialEq, Eq)]
519pub struct MarketHooks {
520 base: Option<Hook>,
521 quote: Option<Hook>,
522}
523
524#[derive(Clone, Debug, PartialEq, Eq)]
525struct Planned {
526 program: Address,
527 validation: Address,
528 missing: HookList,
530}
531
532#[derive(Clone, Debug, PartialEq, Eq)]
533pub struct HookPlan {
534 base: Option<Planned>,
535 quote: Option<Planned>,
536}
537
538impl HookPlan {
539 pub fn accounts(&self) -> impl Iterator<Item = Address> + '_ {
540 [&self.base, &self.quote]
541 .into_iter()
542 .flatten()
543 .map(|planned| planned.validation)
544 }
545
546 pub fn resolve<'a>(
549 self,
550 accounts: impl Fn(&Address) -> Option<AccountView<'a>>,
551 ) -> MarketHooks {
552 let load = |planned: Planned| {
553 let account = accounts(&planned.validation);
554 Hook::load(planned, account)
555 };
556 MarketHooks {
557 base: self.base.map(load),
558 quote: self.quote.map(load),
559 }
560 }
561}
562
563impl MarketHooks {
564 #[must_use]
565 pub const fn base(&self) -> Option<&Hook> {
566 self.base.as_ref()
567 }
568
569 #[must_use]
570 pub const fn quote(&self) -> Option<&Hook> {
571 self.quote.as_ref()
572 }
573
574 #[must_use]
577 pub fn plan(
578 &self,
579 base: (&Address, &MintState),
580 quote: (&Address, &MintState),
581 ) -> HookPlan {
582 let plan = |previous: Option<&Hook>,
583 (mint, state): (&Address, &MintState)| {
584 state.transfer_hook.map(|program| {
585 let asked = previous.filter(|hook| hook.program == program);
586 Planned {
587 program,
588 validation: asked.map_or_else(
589 || find_validation_address(mint, &program),
590 |hook| hook.validation,
591 ),
592 missing: match asked.map(|hook| &hook.list) {
593 Some(HookList::Pending | HookList::Absent) => {
594 HookList::Absent
595 }
596 Some(HookList::Ready(_) | HookList::Unsupported(_))
597 | None => HookList::Pending,
598 },
599 }
600 })
601 };
602 HookPlan {
603 base: plan(self.base(), base),
604 quote: plan(self.quote(), quote),
605 }
606 }
607
608 #[must_use]
612 pub fn screen(self, trade: &StandInTrade, fixed: &[AccountMeta]) -> Self {
613 let screen = |hook: Hook, alone: &Self| {
614 if hook.check_routable(Unloaded::Refuse).is_err() {
615 return hook;
616 }
617 let clash = [TradeDirection::Buy, TradeDirection::Sell]
618 .into_iter()
619 .find_map(|direction| {
620 let swap = HookSwap {
621 trade,
622 direction,
623 fixed,
624 };
625 alone.swap_accounts(swap, Unloaded::Refuse).err()
626 });
627 let reason = match clash {
628 Some(HookError::Unresolvable(reason)) => reason,
629 Some(HookError::Clash(clash)) => Unresolvable::Clash(clash),
630 None | Some(HookError::NotLoaded | HookError::NoList) => {
632 return hook;
633 }
634 };
635 Hook {
636 list: HookList::Unsupported(reason),
637 ..hook
638 }
639 };
640 Self {
641 base: self.base.map(|hook| {
642 let alone = Self {
643 base: Some(hook.clone()),
644 quote: None,
645 };
646 screen(hook, &alone)
647 }),
648 quote: self.quote.map(|hook| {
649 let alone = Self {
650 base: None,
651 quote: Some(hook.clone()),
652 };
653 screen(hook, &alone)
654 }),
655 }
656 }
657
658 fn iter(&self) -> impl Iterator<Item = &Hook> {
659 [&self.base, &self.quote].into_iter().flatten()
660 }
661
662 pub fn validation_accounts(&self) -> impl Iterator<Item = Address> + '_ {
665 self.iter().map(|hook| hook.validation)
666 }
667
668 pub fn check_routable(&self, unloaded: Unloaded) -> Result<(), HookError> {
669 self.iter()
670 .try_for_each(|hook| hook.check_routable(unloaded))
671 }
672
673 #[must_use]
674 pub fn programs(&self) -> Vec<Address> {
675 let mut programs: Vec<Address> = Vec::new();
676 for hook in self.iter() {
677 if !programs.contains(&hook.program) {
678 programs.push(hook.program);
679 }
680 }
681 programs
682 }
683
684 pub fn swap_accounts(
686 &self,
687 swap: HookSwap<'_>,
688 unloaded: Unloaded,
689 ) -> Result<Vec<AccountMeta>, HookError> {
690 let HookSwap {
691 trade,
692 direction,
693 fixed,
694 } = swap;
695 let into_vault = |source: Address, mint: Address, vault: Address| Leg {
696 source,
697 mint,
698 destination: vault,
699 owner: trade.user,
700 destination_owner: trade.market,
701 };
702 let out_of_vault =
703 |vault: Address, mint: Address, destination: Address| Leg {
704 source: vault,
705 mint,
706 destination,
707 owner: trade.market,
708 destination_owner: trade.user,
709 };
710 let legs = match direction {
711 TradeDirection::Buy => [
712 (
713 self.quote(),
714 into_vault(
715 trade.user_quote_account,
716 trade.quote_mint,
717 trade.quote_vault,
718 ),
719 ),
720 (
721 self.base(),
722 out_of_vault(
723 trade.base_vault,
724 trade.base_mint,
725 trade.user_base_account,
726 ),
727 ),
728 ],
729 TradeDirection::Sell => [
730 (
731 self.base(),
732 into_vault(
733 trade.user_base_account,
734 trade.base_mint,
735 trade.base_vault,
736 ),
737 ),
738 (
739 self.quote(),
740 out_of_vault(
741 trade.quote_vault,
742 trade.quote_mint,
743 trade.user_quote_account,
744 ),
745 ),
746 ],
747 };
748 let own = TradeOwn {
749 fixed,
750 market: &trade.market,
751 };
752 trailing_accounts(&legs, &own, unloaded)
753 }
754}
755
756struct TradeOwn<'a> {
758 fixed: &'a [AccountMeta],
759 market: &'a Address,
761}
762
763fn trailing_accounts(
766 legs: &[(Option<&Hook>, Leg)],
767 own: &TradeOwn<'_>,
768 unloaded: Unloaded,
769) -> Result<Vec<AccountMeta>, HookError> {
770 let TradeOwn { fixed, market } = *own;
771 let mut merged: Vec<AccountMeta> = Vec::new();
772 for (hook, leg) in legs {
773 let Some(hook) = hook else {
774 continue;
775 };
776 for meta in hook.leg_accounts(leg, unloaded)? {
777 if let Some(clash) =
778 fixed.iter().find(|account| account.pubkey == meta.pubkey)
779 {
780 let forwarded_read_only = meta.pubkey == *market;
781 if meta.is_writable
782 || (clash.is_writable && !forwarded_read_only)
783 {
784 return Err(HookError::Clash(Clash::WritableDuplicate));
785 }
786 }
787 match merged
788 .iter_mut()
789 .find(|account| account.pubkey == meta.pubkey)
790 {
791 Some(existing) => existing.is_writable |= meta.is_writable,
792 None => merged.push(meta),
793 }
794 }
795 }
796 Ok(merged)
797}
798
799#[cfg(test)]
800mod tests;