1use crate::account::AccountView;
35use crate::error::ProgramError;
36use crate::instruction::{InstructionAccount, InstructionView, Signer};
37use crate::token::{
38 TokenInstruction, TokenProgram, TokenSink, Trailing, MAX_TOKEN_MULTISIG_SIGNERS,
39};
40use crate::ProgramResult;
41use core::mem::MaybeUninit;
42
43pub const BATCH_DISCRIMINATOR: u8 = 255;
45
46pub const BATCH_INSTRUCTION_HEADER_LEN: usize = 2;
49
50pub struct TokenBatch<'a, const DATA: usize = 256, const ACCOUNTS: usize = 16> {
57 data: [MaybeUninit<u8>; DATA],
58 data_len: usize,
59 accounts: [MaybeUninit<InstructionAccount<'a>>; ACCOUNTS],
60 views: [MaybeUninit<&'a AccountView<'a>>; ACCOUNTS],
61 accounts_len: usize,
62 instructions: usize,
63}
64
65impl<'a, const DATA: usize, const ACCOUNTS: usize> Default for TokenBatch<'a, DATA, ACCOUNTS> {
66 fn default() -> Self {
67 Self::new()
68 }
69}
70
71impl<'a, const DATA: usize, const ACCOUNTS: usize> TokenBatch<'a, DATA, ACCOUNTS> {
72 pub const fn new() -> Self {
74 const {
75 assert!(DATA >= 1, "a TokenBatch needs room for its discriminator");
76 assert!(
77 ACCOUNTS <= crate::cpi::MAX_STATIC_CPI_ACCOUNTS,
78 "a TokenBatch cannot carry more accounts than one CPI"
79 );
80 }
81 let mut data = [MaybeUninit::uninit(); DATA];
82 data[0] = MaybeUninit::new(BATCH_DISCRIMINATOR);
83 Self {
84 data,
85 data_len: 1,
86 accounts: [MaybeUninit::uninit(); ACCOUNTS],
87 views: [MaybeUninit::uninit(); ACCOUNTS],
88 accounts_len: 0,
89 instructions: 0,
90 }
91 }
92
93 #[inline]
96 pub fn push(&mut self, instruction: &impl TokenInstruction<'a>) -> ProgramResult {
97 self.push_multisig(instruction, &[])
98 }
99
100 #[inline]
103 pub fn push_multisig(
104 &mut self,
105 instruction: &impl TokenInstruction<'a>,
106 multisig_signers: &[&'a AccountView<'a>],
107 ) -> ProgramResult {
108 let checkpoint = (self.data_len, self.accounts_len, self.instructions);
111 let result = instruction.emit(multisig_signers, self);
112 if result.is_err() {
113 (self.data_len, self.accounts_len, self.instructions) = checkpoint;
114 }
115 result
116 }
117
118 #[inline(always)]
120 pub const fn len(&self) -> usize {
121 self.instructions
122 }
123
124 #[inline(always)]
126 pub const fn is_empty(&self) -> bool {
127 self.instructions == 0
128 }
129
130 #[inline(always)]
133 pub fn data(&self) -> &[u8] {
134 unsafe { core::slice::from_raw_parts(self.data.as_ptr() as *const u8, self.data_len) }
137 }
138
139 #[inline(always)]
141 pub fn account_metas(&self) -> &[InstructionAccount<'a>] {
142 unsafe {
144 core::slice::from_raw_parts(
145 self.accounts.as_ptr() as *const InstructionAccount<'a>,
146 self.accounts_len,
147 )
148 }
149 }
150
151 #[inline(always)]
153 pub fn account_views(&self) -> &[&'a AccountView<'a>] {
154 unsafe {
157 core::slice::from_raw_parts(
158 self.views.as_ptr() as *const &'a AccountView<'a>,
159 self.accounts_len,
160 )
161 }
162 }
163
164 #[inline]
166 pub fn invoke(&self) -> ProgramResult {
167 self.invoke_on(TokenProgram::Legacy, &[])
168 }
169
170 #[inline]
172 pub fn invoke_signed(&self, signers: &[Signer<'_, '_>]) -> ProgramResult {
173 self.invoke_on(TokenProgram::Legacy, signers)
174 }
175
176 #[inline]
179 pub fn invoke_on(&self, program: TokenProgram, signers: &[Signer<'_, '_>]) -> ProgramResult {
180 if self.instructions == 0 {
181 return Err(ProgramError::InvalidArgument);
182 }
183 let instruction = InstructionView {
184 program_id: program.address(),
185 data: self.data(),
186 accounts: self.account_metas(),
187 };
188 crate::cpi::invoke_signed_batch_with_bounds::<{ crate::cpi::MAX_STATIC_CPI_ACCOUNTS }>(
189 &instruction,
190 self.account_views(),
191 signers,
192 )
193 }
194}
195
196impl<'a, const DATA: usize, const ACCOUNTS: usize> TokenSink<'a>
197 for TokenBatch<'a, DATA, ACCOUNTS>
198{
199 #[inline]
200 fn emit<const N: usize>(
201 &mut self,
202 data: &[u8],
203 accounts: [InstructionAccount<'a>; N],
204 views: [&'a AccountView<'a>; N],
205 trailing: &[Trailing<'_, 'a>],
206 ) -> ProgramResult {
207 let mut count = N;
208 for run in trailing {
209 if run.signer && run.views.len() > MAX_TOKEN_MULTISIG_SIGNERS {
210 return Err(ProgramError::InvalidArgument);
211 }
212 count = count
213 .checked_add(run.views.len())
214 .ok_or(ProgramError::ArithmeticOverflow)?;
215 }
216 let Some(new_data_len) = self
217 .data_len
218 .checked_add(BATCH_INSTRUCTION_HEADER_LEN)
219 .and_then(|at| at.checked_add(data.len()))
220 else {
221 return Err(ProgramError::ArithmeticOverflow);
222 };
223 let Some(new_accounts_len) = self.accounts_len.checked_add(count) else {
224 return Err(ProgramError::ArithmeticOverflow);
225 };
226 if count > u8::MAX as usize
227 || data.len() > u8::MAX as usize
228 || new_data_len > DATA
229 || new_accounts_len > ACCOUNTS
230 {
231 return Err(ProgramError::InvalidArgument);
232 }
233
234 let at = self.data_len;
235 self.data[at].write(count as u8);
236 self.data[at + 1].write(data.len() as u8);
237 for (slot, byte) in self.data[at + BATCH_INSTRUCTION_HEADER_LEN..new_data_len]
238 .iter_mut()
239 .zip(data)
240 {
241 slot.write(*byte);
242 }
243
244 let mut index = self.accounts_len;
245 for i in 0..N {
246 self.accounts[index].write(accounts[i]);
247 self.views[index].write(views[i]);
248 index += 1;
249 }
250 for run in trailing {
251 for view in run.views {
252 self.accounts[index].write(InstructionAccount::new(
253 view.address(),
254 run.writable,
255 run.signer,
256 ));
257 self.views[index].write(*view);
258 index += 1;
259 }
260 }
261
262 let previous_accounts_len = self.accounts_len;
266 self.accounts_len = new_accounts_len;
267 let instruction = InstructionView {
268 program_id: TokenProgram::Legacy.address(), data,
270 accounts: &self.account_metas()[previous_accounts_len..],
271 };
272 if let Err(error) = crate::cpi::validate_no_duplicate_writable(
273 &instruction,
274 &self.account_views()[previous_accounts_len..],
275 ) {
276 self.accounts_len = previous_accounts_len;
277 return Err(error);
278 }
279 self.data_len = new_data_len;
280 self.instructions += 1;
281 Ok(())
282 }
283}
284
285#[cfg(test)]
286mod tests {
287 use super::*;
288 use crate::address::Address;
289 use crate::token::{CloseAccount, TransferChecked, TOKEN_PROGRAM_ID};
290 use hopper_native::{
291 AccountView as NativeAccountView, Address as NativeAddress, RuntimeAccount, NOT_BORROWED,
292 };
293
294 fn make_account(address: [u8; 32], signer: bool) -> (std::vec::Vec<u64>, AccountView<'static>) {
295 let mut backing = std::vec![0u64; RuntimeAccount::SIZE.div_ceil(8)];
296 let raw = backing.as_mut_ptr() as *mut RuntimeAccount;
297 unsafe {
300 raw.write(RuntimeAccount {
301 borrow_state: NOT_BORROWED,
302 is_signer: u8::from(signer),
303 is_writable: 1,
304 executable: 0,
305 resize_delta: 0,
306 address: NativeAddress::new_from_array(address),
307 owner: NativeAddress::new_from_array(TOKEN_PROGRAM_ID.to_bytes()),
308 lamports: 1,
309 data_len: 0,
310 });
311 }
312 let backend = unsafe { NativeAccountView::new_unchecked(raw) };
314 (backing, AccountView::from_backend(backend))
315 }
316
317 #[test]
318 fn a_batch_lays_out_headers_data_and_accounts_in_push_order() {
319 let (_b1, from) = make_account([1; 32], false);
320 let (_b2, mint) = make_account([2; 32], false);
321 let (_b3, to) = make_account([3; 32], false);
322 let (_b4, authority) = make_account([4; 32], true);
323 let (_b5, destination) = make_account([5; 32], false);
324 let from = &from;
325 let mint = &mint;
326 let to = &to;
327 let authority = &authority;
328 let destination = &destination;
329
330 let mut batch = TokenBatch::<128, 8>::new();
331 assert!(batch.is_empty());
332 assert!(batch.invoke().is_err(), "an empty batch is refused");
333 batch
334 .push(&TransferChecked {
335 from,
336 mint,
337 to,
338 authority,
339 amount: 5,
340 decimals: 2,
341 })
342 .unwrap();
343 batch
344 .push(&CloseAccount {
345 account: from,
346 destination,
347 authority,
348 })
349 .unwrap();
350 assert_eq!(batch.len(), 2);
351
352 let mut expected = std::vec![255u8];
353 expected.extend_from_slice(&[4, 10, 12, 5, 0, 0, 0, 0, 0, 0, 0, 2]);
354 expected.extend_from_slice(&[3, 1, 9]);
355 assert_eq!(batch.data(), &expected[..]);
356
357 let metas = batch.account_metas();
358 assert_eq!(metas.len(), 7);
359 let flags: std::vec::Vec<(u8, bool, bool)> = metas
360 .iter()
361 .map(|m| (m.address.as_array()[0], m.is_writable, m.is_signer))
362 .collect();
363 assert_eq!(
364 flags,
365 std::vec![
366 (1, true, false),
367 (2, false, false),
368 (3, true, false),
369 (4, false, true),
370 (1, true, false),
371 (5, true, false),
372 (4, false, true),
373 ]
374 );
375 assert_eq!(batch.account_views().len(), 7);
376 assert_eq!(batch.account_views()[3].address(), authority.address());
377 }
378
379 #[test]
380 fn a_push_that_overflows_leaves_the_batch_unchanged() {
381 let (_b1, account) = make_account([1; 32], false);
382 let (_b2, destination) = make_account([5; 32], false);
383 let (_b3, authority) = make_account([4; 32], true);
384 let close = CloseAccount {
385 account: &account,
386 destination: &destination,
387 authority: &authority,
388 };
389 let mut small = TokenBatch::<3, 8>::new();
390 assert_eq!(small.push(&close), Err(ProgramError::InvalidArgument));
391 assert!(small.is_empty());
392 assert_eq!(small.data(), &[255]);
393
394 let mut few = TokenBatch::<64, 2>::new();
395 assert_eq!(few.push(&close), Err(ProgramError::InvalidArgument));
396 assert_eq!(few.account_metas().len(), 0);
397 }
398
399 #[test]
400 fn a_batch_refuses_a_self_transfer_but_reuses_accounts_between_instructions() {
401 let (_b1, from) = make_account([1; 32], false);
402 let (_b2, mint) = make_account([2; 32], false);
403 let (_b3, to) = make_account([3; 32], false);
404 let (_b4, authority) = make_account([4; 32], true);
405 let mut batch = TokenBatch::<64, 12>::new();
406 let mut transfer = TransferChecked {
407 from: &from,
408 mint: &mint,
409 to: &to,
410 authority: &authority,
411 amount: 5,
412 decimals: 2,
413 };
414 batch.push(&transfer).unwrap();
415 let before = batch.data().to_vec();
416 transfer.to = &from;
417 assert_eq!(
418 batch.push(&transfer),
419 Err(ProgramError::AccountBorrowFailed)
420 );
421 assert_eq!(batch.data(), before);
422 assert_eq!(batch.len(), 1);
423 assert_eq!(batch.account_metas().len(), 4);
424 transfer.from = &to;
425 batch.push(&transfer).unwrap();
426 assert_eq!(batch.len(), 2);
427 assert_eq!(batch.account_metas().len(), 8);
428 }
429
430 #[test]
431 fn custom_instruction_failure_rolls_back_all_emitted_instructions() {
432 struct Partial;
433 impl<'a> TokenInstruction<'a> for Partial {
434 fn emit(
435 &self,
436 _: &[&'a AccountView<'a>],
437 sink: &mut impl TokenSink<'a>,
438 ) -> ProgramResult {
439 sink.emit(&[17], [], [], &[])?;
440 Err(ProgramError::InvalidArgument)
441 }
442 }
443 for multisig in [false, true] {
444 let mut batch = TokenBatch::<16, 0>::new();
445 TokenSink::emit(&mut batch, &[20], [], [], &[]).unwrap();
446 let before = batch.data().to_vec();
447 let result = if multisig {
448 batch.push_multisig(&Partial, &[])
449 } else {
450 batch.push(&Partial)
451 };
452 assert_eq!(result, Err(ProgramError::InvalidArgument));
453 assert_eq!(batch.data(), before);
454 assert_eq!(batch.len(), 1);
455 }
456 }
457
458 #[test]
459 fn writable_trailing_alias_is_refused_even_for_distinct_views() {
460 let (_b1, from) = make_account([1; 32], false);
461 let (_b2, alias) = make_account([1; 32], false);
462 let mut batch = TokenBatch::<16, 2>::new();
463 assert_eq!(
464 TokenSink::emit(
465 &mut batch,
466 &[17],
467 [InstructionAccount::writable(from.address())],
468 [&from],
469 &[Trailing::writable(&[&alias])]
470 ),
471 Err(ProgramError::AccountBorrowFailed)
472 );
473 assert!(batch.is_empty());
474 assert_eq!(batch.data(), &[255]);
475 assert!(batch.account_views().is_empty());
476 }
477
478 #[test]
479 fn a_multisig_push_appends_the_signers_after_the_fixed_accounts() {
480 let (_b1, account) = make_account([1; 32], false);
481 let (_b2, destination) = make_account([5; 32], false);
482 let (_b3, multisig) = make_account([6; 32], false);
483 let (_b4, s1) = make_account([7; 32], true);
484 let (_b5, s2) = make_account([8; 32], true);
485 let close = CloseAccount {
486 account: &account,
487 destination: &destination,
488 authority: &multisig,
489 };
490 let mut batch = TokenBatch::<64, 8>::new();
491 batch.push_multisig(&close, &[&s1, &s2]).unwrap();
492 assert_eq!(&batch.data()[1..3], &[5, 1]);
493 let metas = batch.account_metas();
494 assert_eq!(metas.len(), 5);
495 assert!(!metas[2].is_signer, "a multisig authority is not a signer");
496 assert!(metas[3].is_signer && metas[4].is_signer);
497 let _ = Address::default();
498 }
499}