From 861f1d1230723d812310b6fd40182b1a8ca1ece9 Mon Sep 17 00:00:00 2001 From: jeffyanta Date: Mon, 28 Sep 2026 09:45:51 -0400 Subject: [PATCH 1/2] Add support for Solana transaction v1 and clean up the solana package --- ocp/data/blockchain.go | 48 -- ocp/data/transaction/transaction.go | 7 +- ocp/rpc/transaction/errors.go | 11 +- ocp/rpc/transaction/stateful_swap.go | 38 +- ocp/rpc/transaction/stateless_swap.go | 13 +- ocp/worker/currency/feeburner/worker.go | 7 +- ocp/worker/currency/feeburner/worker_test.go | 8 +- ocp/worker/currency/launcher/util.go | 6 +- ocp/worker/sequencer/worker.go | 5 +- ocp/worker/sequencer/worker_test.go | 9 +- ocp/worker/swap/util.go | 5 +- solana/client.go | 292 ++----------- solana/client_with_fallback.go | 50 --- solana/client_with_fallback_test.go | 29 -- solana/encoding.go | 333 +++++++++++++- solana/system/program_test.go | 8 +- solana/transaction.go | 141 +++++- solana/transaction_test.go | 434 ++++++++++++++++++- 18 files changed, 1014 insertions(+), 430 deletions(-) diff --git a/ocp/data/blockchain.go b/ocp/data/blockchain.go index 522e896..d361edc 100644 --- a/ocp/data/blockchain.go +++ b/ocp/data/blockchain.go @@ -2,13 +2,11 @@ package data import ( "context" - "crypto/ed25519" "github.com/mr-tron/base58" "github.com/code-payments/ocp-server/database/query" "github.com/code-payments/ocp-server/metrics" - "github.com/code-payments/ocp-server/ocp/config" "github.com/code-payments/ocp-server/solana" "github.com/code-payments/ocp-server/solana/token" ) @@ -24,15 +22,12 @@ type BlockchainData interface { GetBlockchainAccountDataAfterBlock(ctx context.Context, account string, slot uint64) ([]byte, uint64, error) GetBlockchainBalance(ctx context.Context, account string, commitment solana.Commitment) (uint64, uint64, error) GetBlockchainBlock(ctx context.Context, slot uint64) (*solana.Block, error) - GetBlockchainBlockSignatures(ctx context.Context, slot uint64) ([]string, error) - GetBlockchainBlocksWithLimit(ctx context.Context, start uint64, limit uint64) ([]uint64, error) GetBlockchainHistory(ctx context.Context, account string, commitment solana.Commitment, opts ...query.Option) ([]*solana.TransactionSignature, error) GetBlockchainMinimumBalanceForRentExemption(ctx context.Context, size uint64) (uint64, error) GetBlockchainLatestBlockhash(ctx context.Context) (solana.Blockhash, error) GetBlockchainSignatureStatuses(ctx context.Context, signatures []solana.Signature) ([]*solana.SignatureStatus, error) GetBlockchainSlot(ctx context.Context, commitment solana.Commitment) (uint64, error) GetBlockchainTokenAccountInfo(ctx context.Context, account, mint string, commitment solana.Commitment) (*token.Account, error) - GetBlockchainTokenAccountsByOwner(ctx context.Context, account string) ([]ed25519.PublicKey, error) GetBlockchainTransaction(ctx context.Context, sig string, commitment solana.Commitment) (*solana.ConfirmedTransaction, error) GetBlockchainTransactionTokenBalances(ctx context.Context, sig string) (*solana.TransactionTokenBalances, error) GetBlockchainFilteredProgramAccounts(ctx context.Context, program string, offset uint, filterValue []byte) ([]solana.ProgramAccount, uint64, error) @@ -130,22 +125,6 @@ func (dp *BlockchainProvider) GetBlockchainTokenAccountInfo(ctx context.Context, } return res, err } -func (dp *BlockchainProvider) GetBlockchainTokenAccountsByOwner(ctx context.Context, account string) ([]ed25519.PublicKey, error) { - tracer := metrics.TraceMethodCall(ctx, blockchainProviderMetricsName, "GetBlockchainTokenAccountsByOwner") - defer tracer.End() - - accountId, err := base58.Decode(account) - if err != nil { - return nil, err - } - - res, err := dp.sc.GetTokenAccountsByOwner(accountId, config.CoreMintPublicKeyBytes) - - if err != nil { - tracer.OnError(err) - } - return res, err -} func (dp *BlockchainProvider) GetBlockchainSlot(ctx context.Context, commitment solana.Commitment) (uint64, error) { tracer := metrics.TraceMethodCall(ctx, blockchainProviderMetricsName, "GetBlockchainSlot") defer tracer.End() @@ -158,21 +137,6 @@ func (dp *BlockchainProvider) GetBlockchainSlot(ctx context.Context, commitment return res, err } -func (dp *BlockchainProvider) GetBlockchainBlocksWithLimit(ctx context.Context, start uint64, limit uint64) ([]uint64, error) { - tracer := metrics.TraceMethodCall(ctx, blockchainProviderMetricsName, "GetBlockchainBlocksWithLimit") - defer tracer.End() - - // TODO: this call is deprecated, remove it - // https://docs.solana.com/developing/clients/jsonrpc-api#getconfirmedblockswithlimit - - res, err := dp.sc.GetConfirmedBlocksWithLimit(start, limit) - - if err != nil { - tracer.OnError(err) - } - return res, err -} - func (dp *BlockchainProvider) GetBlockchainBlock(ctx context.Context, slot uint64) (*solana.Block, error) { tracer := metrics.TraceMethodCall(ctx, blockchainProviderMetricsName, "GetBlockchainBlock") defer tracer.End() @@ -185,18 +149,6 @@ func (dp *BlockchainProvider) GetBlockchainBlock(ctx context.Context, slot uint6 return res, err } -func (dp *BlockchainProvider) GetBlockchainBlockSignatures(ctx context.Context, slot uint64) ([]string, error) { - tracer := metrics.TraceMethodCall(ctx, blockchainProviderMetricsName, "GetBlockchainBlockSignatures") - defer tracer.End() - - res, err := dp.sc.GetBlockSignatures(slot) - - if err != nil { - tracer.OnError(err) - } - return res, err -} - func (dp *BlockchainProvider) GetBlockchainHistory(ctx context.Context, account string, commitment solana.Commitment, opts ...query.Option) ([]*solana.TransactionSignature, error) { tracer := metrics.TraceMethodCall(ctx, blockchainProviderMetricsName, "GetBlockchainHistory") defer tracer.End() diff --git a/ocp/data/transaction/transaction.go b/ocp/data/transaction/transaction.go index 66f7183..7880276 100644 --- a/ocp/data/transaction/transaction.go +++ b/ocp/data/transaction/transaction.go @@ -92,11 +92,16 @@ func FromConfirmedTransaction(tx *solana.ConfirmedTransaction) (*Record, error) return nil, errors.New("unsupported transaction version") } + data, err := tx.Transaction.Marshal() + if err != nil { + return nil, err + } + sig := tx.Transaction.Signature() res := &Record{ Signature: base58.Encode(sig), Slot: tx.Slot, - Data: tx.Transaction.Marshal(), + Data: data, HasErrors: tx.Err != nil, ConfirmationState: ConfirmationFinalized, CreatedAt: time.Now(), diff --git a/ocp/rpc/transaction/errors.go b/ocp/rpc/transaction/errors.go index 10fa26d..d574f9b 100644 --- a/ocp/rpc/transaction/errors.go +++ b/ocp/rpc/transaction/errors.go @@ -162,26 +162,31 @@ func toInvalidTxnSignatureErrorDetails( actionId uint32, txn solana.Transaction, signature *commonpb.Signature, -) *transactionpb.ErrorDetails { +) (*transactionpb.ErrorDetails, error) { // Clear out all signatures, so clients have no way of submitting this transaction var emptySig solana.Signature for i := range txn.Signatures { copy(txn.Signatures[i][:], emptySig[:]) } + marshalledTxn, err := txn.Marshal() + if err != nil { + return nil, err + } + return &transactionpb.ErrorDetails{ Type: &transactionpb.ErrorDetails_InvalidSignature{ InvalidSignature: &transactionpb.InvalidSignatureErrorDetails{ ActionId: actionId, ExpectedBlob: &transactionpb.InvalidSignatureErrorDetails_ExpectedTransaction{ ExpectedTransaction: &commonpb.Transaction{ - Value: txn.Marshal(), + Value: marshalledTxn, }, }, ProvidedSignature: signature, }, }, - } + }, nil } func toInvalidVirtualIxnSignatureErrorDetails( diff --git a/ocp/rpc/transaction/stateful_swap.go b/ocp/rpc/transaction/stateful_swap.go index 015127c..a24e57e 100644 --- a/ocp/rpc/transaction/stateful_swap.go +++ b/ocp/rpc/transaction/stateful_swap.go @@ -593,7 +593,11 @@ func (s *transactionServer) handleReserveStatefulSwap( txn.SetBlockhash(selectedNonce.Blockhash) - marshalledTxnMessage := txn.Message.Marshal() + marshalledTxnMessage, err := txn.Message.Marshal() + if err != nil { + log.With(zap.Error(err)).Warn("failure marshalling transaction message") + return handleStatefulSwapError(streamer, err) + } // // Section: Server parameters @@ -657,10 +661,15 @@ func (s *transactionServer) handleReserveStatefulSwap( marshalledTxnMessage, protoSignature.Value, ) { + errorDetails, err := toInvalidTxnSignatureErrorDetails(0, txn, protoSignature) + if err != nil { + log.With(zap.Error(err)).Warn("failure creating error details") + return handleStatefulSwapError(streamer, err) + } return handleStatefulSwapStructuredError( streamer, transactionpb.StatefulSwapResponse_Error_SIGNATURE_ERROR, - toInvalidTxnSignatureErrorDetails(0, txn, protoSignature), + errorDetails, ) } @@ -683,7 +692,11 @@ func (s *transactionServer) handleReserveStatefulSwap( return handleStatefulSwapError(streamer, err) } - marshalledTxn := txn.Marshal() + marshalledTxn, err := txn.Marshal() + if err != nil { + log.With(zap.Error(err)).Warn("failure marshalling transaction") + return handleStatefulSwapError(streamer, err) + } txnSignature := base58.Encode(txn.Signature()) @@ -1022,7 +1035,11 @@ func (s *transactionServer) handleStablecoinStatefulSwap( txn.SetBlockhash(selectedNonce.Blockhash) - marshalledTxnMessage := txn.Message.Marshal() + marshalledTxnMessage, err := txn.Message.Marshal() + if err != nil { + log.With(zap.Error(err)).Warn("failure marshalling transaction message") + return handleStatefulSwapError(streamer, err) + } // // Section: Server parameters @@ -1082,10 +1099,15 @@ func (s *transactionServer) handleStablecoinStatefulSwap( marshalledTxnMessage, protoSignature.Value, ) { + errorDetails, err := toInvalidTxnSignatureErrorDetails(0, txn, protoSignature) + if err != nil { + log.With(zap.Error(err)).Warn("failure creating error details") + return handleStatefulSwapError(streamer, err) + } return handleStatefulSwapStructuredError( streamer, transactionpb.StatefulSwapResponse_Error_SIGNATURE_ERROR, - toInvalidTxnSignatureErrorDetails(0, txn, protoSignature), + errorDetails, ) } @@ -1101,7 +1123,11 @@ func (s *transactionServer) handleStablecoinStatefulSwap( return handleStatefulSwapError(streamer, err) } - marshalledTxn := txn.Marshal() + marshalledTxn, err := txn.Marshal() + if err != nil { + log.With(zap.Error(err)).Warn("failure marshalling transaction") + return handleStatefulSwapError(streamer, err) + } txnSignature := base58.Encode(txn.Signature()) diff --git a/ocp/rpc/transaction/stateless_swap.go b/ocp/rpc/transaction/stateless_swap.go index e950387..fc464b6 100644 --- a/ocp/rpc/transaction/stateless_swap.go +++ b/ocp/rpc/transaction/stateless_swap.go @@ -274,7 +274,11 @@ func (s *transactionServer) handleStablecoinStatelessSwap( } txn.SetBlockhash(blockhash) - marshalledTxnMessage := txn.Message.Marshal() + marshalledTxnMessage, err := txn.Message.Marshal() + if err != nil { + log.With(zap.Error(err)).Warn("failure marshalling transaction message") + return handleStatelessSwapError(streamer, err) + } // // Section: Server parameters @@ -322,10 +326,15 @@ func (s *transactionServer) handleStablecoinStatelessSwap( marshalledTxnMessage, protoSignature.Value, ) { + errorDetails, err := toInvalidTxnSignatureErrorDetails(0, txn, protoSignature) + if err != nil { + log.With(zap.Error(err)).Warn("failure creating error details") + return handleStatelessSwapError(streamer, err) + } return handleStatelessSwapStructuredError( streamer, transactionpb.StatelessSwapResponse_Error_SIGNATURE_ERROR, - toInvalidTxnSignatureErrorDetails(0, txn, protoSignature), + errorDetails, ) } diff --git a/ocp/worker/currency/feeburner/worker.go b/ocp/worker/currency/feeburner/worker.go index e8d7086..092bb14 100644 --- a/ocp/worker/currency/feeburner/worker.go +++ b/ocp/worker/currency/feeburner/worker.go @@ -149,7 +149,12 @@ func (p *runtime) packBurnBatches(targets []*burnTarget, maxBurnsPerBatch int) [ candidate := append(current, target) txn := p.makeBurnTransaction(candidate) - if len(txn.Marshal()) > solana.MaxTransactionSize { + marshalledTxn, err := txn.Marshal() + if err != nil { + p.log.With(zap.Error(err), zap.String("mint", target.mint)).Warn("skipping currency with unmarshallable burn transaction") + continue + } + if len(marshalledTxn) > solana.MaxLegacyTransactionSize { if len(current) == 0 { p.log.With(zap.String("mint", target.mint)).Warn("skipping currency with oversized burn transaction") continue diff --git a/ocp/worker/currency/feeburner/worker_test.go b/ocp/worker/currency/feeburner/worker_test.go index de0f728..af12593 100644 --- a/ocp/worker/currency/feeburner/worker_test.go +++ b/ocp/worker/currency/feeburner/worker_test.go @@ -45,7 +45,9 @@ func TestPackBurnBatches(t *testing.T) { for i, batch := range batches { txn := p.makeBurnTransaction(batch) - assert.LessOrEqual(t, len(txn.Marshal()), solana.MaxTransactionSize, fmt.Sprintf("batch %d exceeds size limit", i)) + marshalledTxn, err := txn.Marshal() + require.NoError(t, err) + assert.LessOrEqual(t, len(marshalledTxn), solana.MaxLegacyTransactionSize, fmt.Sprintf("batch %d exceeds size limit", i)) assert.LessOrEqual(t, len(batch), defaultMaxBurnsPerBatch, fmt.Sprintf("batch %d exceeds max burns", i)) } @@ -57,7 +59,9 @@ func TestPackBurnBatches(t *testing.T) { } overfilled := append(append([]*burnTarget{}, batches[i]...), batches[i+1][0]) txn := p.makeBurnTransaction(overfilled) - assert.Greater(t, len(txn.Marshal()), solana.MaxTransactionSize, fmt.Sprintf("batch %d is not fully packed", i)) + marshalledTxn, err := txn.Marshal() + require.NoError(t, err) + assert.Greater(t, len(marshalledTxn), solana.MaxLegacyTransactionSize, fmt.Sprintf("batch %d is not fully packed", i)) } assert.Greater(t, len(batches[0]), 1) diff --git a/ocp/worker/currency/launcher/util.go b/ocp/worker/currency/launcher/util.go index 63dbe92..5304cf1 100644 --- a/ocp/worker/currency/launcher/util.go +++ b/ocp/worker/currency/launcher/util.go @@ -757,7 +757,11 @@ func (p *runtime) resizeAndExtendBlockchainAccounts(ctx context.Context, account ixns..., ) - if len(txn.Marshal()) > solana.MaxTransactionSize { + marshalledTxn, err := txn.Marshal() + if err != nil { + return errors.Wrap(err, "error marshalling transaction") + } + if len(marshalledTxn) > solana.MaxLegacyTransactionSize { return errors.New("transaction exceeds maximum size") } diff --git a/ocp/worker/sequencer/worker.go b/ocp/worker/sequencer/worker.go index 7d81345..2a948c2 100644 --- a/ocp/worker/sequencer/worker.go +++ b/ocp/worker/sequencer/worker.go @@ -261,7 +261,10 @@ func (p *runtime) handlePending(ctx context.Context, record *fulfillment.Record) record.Signature = pointer.String(base58.Encode(txn.Signature())) record.Nonce = pointer.String(selectedSolanaNonce.Account.PublicKey().ToBase58()) record.Blockhash = pointer.String(base58.Encode(selectedSolanaNonce.Blockhash[:])) - record.Data = txn.Marshal() + record.Data, err = txn.Marshal() + if err != nil { + return err + } err = selectedSolanaNonce.MarkReservedWithSignature(ctx, *record.Signature) if err != nil { diff --git a/ocp/worker/sequencer/worker_test.go b/ocp/worker/sequencer/worker_test.go index b85bb43..daadfbe 100644 --- a/ocp/worker/sequencer/worker_test.go +++ b/ocp/worker/sequencer/worker_test.go @@ -288,13 +288,16 @@ func (e *workerTestEnv) createAnyFulfillmentInState(t *testing.T, state fulfillm txn.Sign(fakeCodeAccouht.PrivateKey().ToBytes()) + marshalledTxn, err := txn.Marshal() + require.NoError(t, err) + fulfillmentRecord := &fulfillment.Record{ Intent: testutil.NewRandomAccount(t).PublicKey().ToBase58(), IntentType: intent.OpenAccounts, ActionId: 3, ActionType: action.OpenAccount, FulfillmentType: fulfillment.InitializeLockedTimelockAccount, - Data: txn.Marshal(), + Data: marshalledTxn, Signature: pointer.String(base58.Encode(txn.Signature())), Source: "source", Nonce: pointer.String(fakeNonceAccount.PublicKey().ToBase58()), @@ -372,7 +375,9 @@ func (e *workerTestEnv) assertFulfillmentCreatedOnDemand(t *testing.T, id uint64 assert.Equal(t, expectedSignature, *fulfillmentRecord.Signature) assert.Equal(t, nonceAddress, *fulfillmentRecord.Nonce) assert.Equal(t, blockhash, *fulfillmentRecord.Blockhash) - assert.Equal(t, expectedTxn.Marshal(), fulfillmentRecord.Data) + expectedData, err := expectedTxn.Marshal() + require.NoError(t, err) + assert.Equal(t, expectedData, fulfillmentRecord.Data) e.assertNonceState(t, nonceAddress, nonce.StateReserved, expectedSignature, blockhash) } diff --git a/ocp/worker/swap/util.go b/ocp/worker/swap/util.go index e154c5f..040735e 100644 --- a/ocp/worker/swap/util.go +++ b/ocp/worker/swap/util.go @@ -399,7 +399,10 @@ func (p *runtime) markSwapCancelling( swapRecord.Nonce = cancelNonce.Account.PublicKey().ToBase58() swapRecord.Blockhash = base58.Encode(cancelNonce.Blockhash[:]) swapRecord.TransactionSignature = cancelTransactionSignature - swapRecord.TransactionBlob = cancelTxn.Marshal() + swapRecord.TransactionBlob, err = cancelTxn.Marshal() + if err != nil { + return err + } swapRecord.State = swap.StateCancelling return p.data.SaveSwap(ctx, swapRecord) }) diff --git a/solana/client.go b/solana/client.go index 8d50492..e9aa342 100644 --- a/solana/client.go +++ b/solana/client.go @@ -5,10 +5,7 @@ import ( "crypto/ed25519" "encoding/base64" "encoding/json" - "fmt" - "math/rand" "strconv" - "sync" "time" "github.com/mr-tron/base58" @@ -20,25 +17,14 @@ import ( ) const ( - // todo: we can retrieve these from the Syscall account - // but they're unlikely to change. - ticksPerSec = 160 - ticksPerSlot = 64 - slotsPerSec = ticksPerSec / ticksPerSlot - - // PollRate is the rate at which blocks should be polled at. - PollRate = (time.Second / slotsPerSec) / 2 - - // Poll rate is ~2x the slot rate, and we want to wait ~32 slots - sigStatusPollLimit = 2 * 32 - - // Reference: https://github.com/solana-labs/solana/blob/14d793b22c1571fb092d5822189d5b64f32605e6/client/src/rpc_custom_error.rs#L10 - blockNotAvailableCode = -32004 - // Reference: https://github.com/solana-labs/solana/blob/71e9958e061493d7545bd28d4ac7a85aaed6ffbb/client/src/rpc_custom_error.rs#L11 rpcNodeUnhealthyCode = -32005 invalidParamCode = -32602 + + // Highest transaction version the client can parse. RPC methods returning + // transactions fail entirely when a newer version is encountered. + maxSupportedTransactionVersion = 1 ) type Commitment struct { @@ -177,20 +163,13 @@ type Client interface { GetAccountDataAfterBlock(ed25519.PublicKey, uint64) ([]byte, uint64, error) GetBalance(ed25519.PublicKey) (uint64, error) GetBlock(slot uint64) (*Block, error) - GetBlockSignatures(slot uint64) ([]string, error) - GetBlockTime(block uint64) (time.Time, error) - GetConfirmationStatus(Signature, Commitment) (bool, error) - GetConfirmedBlock(slot uint64) (*Block, error) - GetConfirmedBlocksWithLimit(start, limit uint64) ([]uint64, error) GetFilteredProgramAccounts(program ed25519.PublicKey, offset uint, filterValue []byte) ([]ProgramAccount, uint64, error) GetLatestBlockhash() (Blockhash, error) GetMinimumBalanceForRentExemption(size uint64) (lamports uint64, err error) - GetSignatureStatus(Signature, Commitment) (*SignatureStatus, error) GetSignatureStatuses([]Signature) ([]*SignatureStatus, error) GetSignaturesForAddress(owner ed25519.PublicKey, commitment Commitment, limit uint64, before, until string) ([]*TransactionSignature, error) GetSlot(Commitment) (uint64, error) GetTokenAccountBalance(ed25519.PublicKey, Commitment) (uint64, uint64, error) - GetTokenAccountsByOwner(owner, mint ed25519.PublicKey) ([]ed25519.PublicKey, error) GetTransaction(Signature, Commitment) (ConfirmedTransaction, error) GetTransactionTokenBalances(Signature) (TransactionTokenBalances, error) SubmitTransaction(Transaction, Commitment) (Signature, error) @@ -211,10 +190,6 @@ type rpcResponse struct { type client struct { client jsonrpc.RPCClient retrier retry.Retrier - - blockMu sync.RWMutex - blockhash Blockhash - lastWrite time.Time } // New returns a client using the specified endpoint. @@ -307,21 +282,6 @@ func (c *client) GetSlot(commitment Commitment) (slot uint64, err error) { } func (c *client) GetLatestBlockhash() (hash Blockhash, err error) { - // To avoid having thrashing around a similar periodic interval, we - // randomize when we refresh our block hash. This is mostly only a - // concern when running a batch migrator with a _ton_ of goroutines. - window := time.Duration(float64(2*time.Second) * (0.8 + rand.Float64())) - - c.blockMu.RLock() - if time.Since(c.lastWrite) < window { - hash = c.blockhash - } - c.blockMu.RUnlock() - - if hash != (Blockhash{}) { - return hash, nil - } - type response struct { Value struct { Blockhash string `json:"blockhash"` @@ -340,94 +300,9 @@ func (c *client) GetLatestBlockhash() (hash Blockhash, err error) { copy(hash[:], hashBytes) - c.blockMu.Lock() - c.blockhash = hash - c.lastWrite = time.Now() - c.blockMu.Unlock() - return hash, nil } -func (c *client) GetBlockTime(slot uint64) (time.Time, error) { - var unixTs int64 - if err := c.call(&unixTs, "getBlockTime", slot); err != nil { - jsonRPCErr, ok := err.(*jsonrpc.RPCError) - if !ok { - return time.Time{}, errors.Wrapf(err, "getBlockTime() failed to send request") - } - - if jsonRPCErr.Code == blockNotAvailableCode { - return time.Time{}, ErrBlockNotAvailable - } - } - - return time.Unix(unixTs, 0), nil -} - -func (c *client) GetConfirmedBlock(slot uint64) (block *Block, err error) { - type rawBlock struct { - Hash string `json:"blockhash"` // Since this value is in base58, we can't []byte - PrevHash string `json:"previousBlockhash"` - ParentSlot uint64 `json:"parentSlot"` - - RawTransactions []struct { - Transaction []string `json:"transaction"` // [string,encoding] - Meta *struct { - Err interface{} `json:"err"` - } `json:"meta"` - } `json:"transactions"` - } - - var rb *rawBlock - if err := c.call(&rb, "getConfirmedBlock", slot, "base64"); err != nil { - return nil, err - } - - // Not all slots contain a block, which manifests itself as having a nil block - if rb == nil { - return nil, nil - } - - block = &Block{ - ParentSlot: rb.ParentSlot, - Slot: slot, - } - - if block.Hash, err = base58.Decode(rb.Hash); err != nil { - return nil, errors.Wrap(err, "invalid base58 encoding for hash") - } - if block.PrevHash, err = base58.Decode(rb.PrevHash); err != nil { - return nil, errors.Wrapf(err, "invalid base58 encoding for prevHash: %s", rb.PrevHash) - } - - for i, txn := range rb.RawTransactions { - txnBytes, err := base64.StdEncoding.DecodeString(txn.Transaction[0]) - if err != nil { - return nil, errors.Wrapf(err, "invalid base58 encoding for transaction %d", i) - } - - var t Transaction - if err := t.Unmarshal(txnBytes); err != nil { - return nil, errors.Wrapf(err, "invalid bytes for transaction %d", i) - } - - var txErr *TransactionError - if txn.Meta != nil { - txErr, err = ParseTransactionError(txn.Meta.Err) - if err != nil { - return nil, errors.Wrap(err, "failed to parse transaction meta") - } - } - - block.Transactions = append(block.Transactions, BlockTransaction{ - Transaction: t, - Err: txErr, - }) - } - - return block, nil -} - func (c *client) GetBlock(slot uint64) (block *Block, err error) { type rawBlock struct { Hash string `json:"blockhash"` // Since this value is in base58, we can't []byte @@ -443,8 +318,16 @@ func (c *client) GetBlock(slot uint64) (block *Block, err error) { BlockTime *int64 `json:"blockTime"` } + config := struct { + Encoding string `json:"encoding"` + MaxSupportedTransactionVersion int `json:"maxSupportedTransactionVersion"` + }{ + Encoding: "base64", + MaxSupportedTransactionVersion: maxSupportedTransactionVersion, + } + var rb *rawBlock - if err := c.call(&rb, "getBlock", slot, "base64"); err != nil { + if err := c.call(&rb, "getBlock", slot, config); err != nil { return nil, err } @@ -499,34 +382,6 @@ func (c *client) GetBlock(slot uint64) (block *Block, err error) { return block, nil } -func (c *client) GetBlockSignatures(slot uint64) ([]string, error) { - type rawBlock struct { - Signatures []string `json:"signatures"` - } - - config := struct { - TransactionDetails string `json:"transactionDetails"` - }{ - TransactionDetails: "signatures", - } - - var rb *rawBlock - if err := c.call(&rb, "getBlock", slot, config); err != nil { - return nil, err - } - - // Not all slots contain a block, which manifests itself as having a nil block - if rb == nil { - return nil, nil - } - - return rb.Signatures, nil -} - -func (c *client) GetConfirmedBlocksWithLimit(start, limit uint64) (slots []uint64, err error) { - return slots, c.call(&slots, "getConfirmedBlocksWithLimit", start, limit) -} - func (c *client) GetTransaction(sig Signature, commitment Commitment) (ConfirmedTransaction, error) { type rpcResponse struct { Slot uint64 `json:"slot"` @@ -542,7 +397,7 @@ func (c *client) GetTransaction(sig Signature, commitment Commitment) (Confirmed }{ Commitment: commitment.Commitment, Encoding: "base64", - MaxSupportedTransactionVersion: 0, + MaxSupportedTransactionVersion: maxSupportedTransactionVersion, } var resp *rpcResponse @@ -592,7 +447,7 @@ func (c *client) GetTransactionTokenBalances(sig Signature) (TransactionTokenBal MaxSupportedTransactionVersion int `json:"maxSupportedTransactionVersion"` }{ Encoding: "json", // Easier to use json in the event of ever-changing transaction versions - MaxSupportedTransactionVersion: 0, + MaxSupportedTransactionVersion: maxSupportedTransactionVersion, } type rpcResp struct { @@ -687,18 +542,25 @@ func (c *client) GetTokenAccountBalance(account ed25519.PublicKey, commitment Co func (c *client) SubmitTransaction(txn Transaction, commitment Commitment) (Signature, error) { sig := txn.Signatures[0] - txnBytes := txn.Marshal() + txnBytes, err := txn.Marshal() + if err != nil { + return sig, errors.Wrap(err, "failed to marshal transaction") + } + // Base64 is used because base58 is size limited by RPC nodes, which fails + // for v1 transactions larger than the legacy size limit. config := struct { + Encoding string `json:"encoding"` SkipPreflight bool `json:"skipPreflight"` PreflightCommitment string `json:"preflightCommitment"` }{ + Encoding: "base64", SkipPreflight: true, PreflightCommitment: commitment.Commitment, } var sigStr string - err := c.call(&sigStr, "sendTransaction", base58.Encode(txnBytes), config) + err = c.call(&sigStr, "sendTransaction", base64.StdEncoding.EncodeToString(txnBytes), config) if err != nil { jsonRPCErr, ok := err.(*jsonrpc.RPCError) if !ok { @@ -710,8 +572,6 @@ func (c *client) SubmitTransaction(txn Transaction, commitment Commitment) (Sign return sig, err } - fmt.Printf("%+v\n", txResult) - if txResult != nil { if txResult.transactionError != nil { return sig, txResult.transactionError @@ -788,13 +648,15 @@ func (c *client) GetAccountDataAfterBlock(account ed25519.PublicKey, slot uint64 getBlockHeightRequest := jsonrpc.NewRequest("getBlockHeight", []interface{}{CommitmentFinalized}) getBlockRpcConfig := struct { - Encoding string `json:"encoding"` - TransactionDetails string `json:"transactionDetails"` - Rewards bool `json:"rewards"` + Encoding string `json:"encoding"` + TransactionDetails string `json:"transactionDetails"` + Rewards bool `json:"rewards"` + MaxSupportedTransactionVersion int `json:"maxSupportedTransactionVersion"` }{ - Encoding: "base64", - TransactionDetails: "none", - Rewards: false, + Encoding: "base64", + TransactionDetails: "none", + Rewards: false, + MaxSupportedTransactionVersion: maxSupportedTransactionVersion, } getBlockRequest := jsonrpc.NewRequest("getBlock", slot, getBlockRpcConfig) @@ -900,61 +762,6 @@ func (c *client) GetAccountDataAfterBlock(account ed25519.PublicKey, slot uint64 return rawData, unmarshalledGetAccountInfoResp.Context.Slot, nil } -func (c *client) GetConfirmationStatus(sig Signature, commitment Commitment) (bool, error) { - type response struct { - Value bool `json:"value"` - } - - var resp response - if err := c.call(&resp, "confirmTransaction", base58.Encode(sig[:]), commitment); err != nil { - return false, err - } - - return resp.Value, nil -} - -func (c *client) GetSignatureStatus(sig Signature, commitment Commitment) (*SignatureStatus, error) { - var s *SignatureStatus - errConfirmationsNotReached := errors.New("confirmations not reached") - _, err := retry.Retry( - func() error { - statuses, err := c.GetSignatureStatuses([]Signature{sig}) - if err != nil { - return err - } - - s = statuses[0] - if s == nil { - return ErrSignatureNotFound - } - - if s.ErrorResult != nil { - return err - } - - switch commitment { - case CommitmentProcessed: - return nil - case CommitmentConfirmed: - if s.Confirmed() { - return nil - } - case CommitmentFinalized: - if s.Finalized() { - return nil - } - } - - return errConfirmationsNotReached - }, - retry.RetriableErrors(ErrSignatureNotFound, errConfirmationsNotReached), - retry.Limit(sigStatusPollLimit), - retry.Backoff(backoff.Constant(PollRate), PollRate), - ) - - return s, err -} - func (c *client) GetSignaturesForAddress(account ed25519.PublicKey, commitment Commitment, limit uint64, before, until string) ([]*TransactionSignature, error) { req := struct { Commitment string `json:"commitment"` @@ -1088,41 +895,6 @@ func (c *client) GetSignatureStatuses(sigs []Signature) ([]*SignatureStatus, err return statuses, nil } -func (c *client) GetTokenAccountsByOwner(owner, mint ed25519.PublicKey) ([]ed25519.PublicKey, error) { - mintObject := struct { - Mint string `json:"mint"` - }{ - Mint: base58.Encode(mint), - } - config := struct { - Encoding string `json:"encoding"` - Commitment Commitment - }{ - Encoding: "base64", - Commitment: CommitmentConfirmed, - } - - var resp struct { - Value []struct { - PubKey string `json:"pubkey"` - } `json:"value"` - } - if err := c.call(&resp, "getTokenAccountsByOwner", base58.Encode(owner), mintObject, config); err != nil { - return nil, err - } - - keys := make([]ed25519.PublicKey, len(resp.Value)) - for i := range resp.Value { - var err error - keys[i], err = base58.Decode(resp.Value[i].PubKey) - if err != nil { - return nil, errors.Wrap(err, "failed to decode token account public key") - } - } - - return keys, nil -} - func (c *client) GetFilteredProgramAccounts(program ed25519.PublicKey, offset uint, filterValue []byte) ([]ProgramAccount, uint64, error) { type memcmpFilter struct { Offset uint `json:"offset"` diff --git a/solana/client_with_fallback.go b/solana/client_with_fallback.go index 3f52f16..d1866c4 100644 --- a/solana/client_with_fallback.go +++ b/solana/client_with_fallback.go @@ -2,7 +2,6 @@ package solana import ( "crypto/ed25519" - "time" "github.com/pkg/errors" ) @@ -105,41 +104,6 @@ func (c *clientWithFallback) GetBlock(slot uint64) (*Block, error) { ) } -func (c *clientWithFallback) GetBlockSignatures(slot uint64) ([]string, error) { - return withFallback( - func() ([]string, error) { return c.primary.GetBlockSignatures(slot) }, - func() ([]string, error) { return c.fallback.GetBlockSignatures(slot) }, - ) -} - -func (c *clientWithFallback) GetBlockTime(block uint64) (time.Time, error) { - return withFallback( - func() (time.Time, error) { return c.primary.GetBlockTime(block) }, - func() (time.Time, error) { return c.fallback.GetBlockTime(block) }, - ) -} - -func (c *clientWithFallback) GetConfirmationStatus(sig Signature, commitment Commitment) (bool, error) { - return withFallback( - func() (bool, error) { return c.primary.GetConfirmationStatus(sig, commitment) }, - func() (bool, error) { return c.fallback.GetConfirmationStatus(sig, commitment) }, - ) -} - -func (c *clientWithFallback) GetConfirmedBlock(slot uint64) (*Block, error) { - return withFallback( - func() (*Block, error) { return c.primary.GetConfirmedBlock(slot) }, - func() (*Block, error) { return c.fallback.GetConfirmedBlock(slot) }, - ) -} - -func (c *clientWithFallback) GetConfirmedBlocksWithLimit(start, limit uint64) ([]uint64, error) { - return withFallback( - func() ([]uint64, error) { return c.primary.GetConfirmedBlocksWithLimit(start, limit) }, - func() ([]uint64, error) { return c.fallback.GetConfirmedBlocksWithLimit(start, limit) }, - ) -} - func (c *clientWithFallback) GetFilteredProgramAccounts(program ed25519.PublicKey, offset uint, filterValue []byte) ([]ProgramAccount, uint64, error) { return withFallback2( func() ([]ProgramAccount, uint64, error) { @@ -165,13 +129,6 @@ func (c *clientWithFallback) GetMinimumBalanceForRentExemption(size uint64) (uin ) } -func (c *clientWithFallback) GetSignatureStatus(sig Signature, commitment Commitment) (*SignatureStatus, error) { - return withFallback( - func() (*SignatureStatus, error) { return c.primary.GetSignatureStatus(sig, commitment) }, - func() (*SignatureStatus, error) { return c.fallback.GetSignatureStatus(sig, commitment) }, - ) -} - func (c *clientWithFallback) GetSignatureStatuses(sigs []Signature) ([]*SignatureStatus, error) { return withFallback( func() ([]*SignatureStatus, error) { return c.primary.GetSignatureStatuses(sigs) }, @@ -204,13 +161,6 @@ func (c *clientWithFallback) GetTokenAccountBalance(account ed25519.PublicKey, c ) } -func (c *clientWithFallback) GetTokenAccountsByOwner(owner, mint ed25519.PublicKey) ([]ed25519.PublicKey, error) { - return withFallback( - func() ([]ed25519.PublicKey, error) { return c.primary.GetTokenAccountsByOwner(owner, mint) }, - func() ([]ed25519.PublicKey, error) { return c.fallback.GetTokenAccountsByOwner(owner, mint) }, - ) -} - func (c *clientWithFallback) GetTransaction(sig Signature, commitment Commitment) (ConfirmedTransaction, error) { return withFallback( func() (ConfirmedTransaction, error) { return c.primary.GetTransaction(sig, commitment) }, diff --git a/solana/client_with_fallback_test.go b/solana/client_with_fallback_test.go index 6774df0..790cfaf 100644 --- a/solana/client_with_fallback_test.go +++ b/solana/client_with_fallback_test.go @@ -3,7 +3,6 @@ package solana import ( "crypto/ed25519" "testing" - "time" "github.com/pkg/errors" "github.com/stretchr/testify/assert" @@ -75,26 +74,6 @@ func (m *mockClient) GetAccountDataAfterBlock(ed25519.PublicKey, uint64) ([]byte m.callCount++ return nil, 0, nil } -func (m *mockClient) GetBlockSignatures(uint64) ([]string, error) { - m.callCount++ - return nil, nil -} -func (m *mockClient) GetBlockTime(uint64) (time.Time, error) { - m.callCount++ - return time.Time{}, nil -} -func (m *mockClient) GetConfirmationStatus(Signature, Commitment) (bool, error) { - m.callCount++ - return false, nil -} -func (m *mockClient) GetConfirmedBlock(uint64) (*Block, error) { - m.callCount++ - return nil, nil -} -func (m *mockClient) GetConfirmedBlocksWithLimit(uint64, uint64) ([]uint64, error) { - m.callCount++ - return nil, nil -} func (m *mockClient) GetFilteredProgramAccounts(ed25519.PublicKey, uint, []byte) ([]ProgramAccount, uint64, error) { m.callCount++ return nil, 0, nil @@ -103,10 +82,6 @@ func (m *mockClient) GetMinimumBalanceForRentExemption(uint64) (uint64, error) { m.callCount++ return 0, nil } -func (m *mockClient) GetSignatureStatus(Signature, Commitment) (*SignatureStatus, error) { - m.callCount++ - return nil, nil -} func (m *mockClient) GetSignatureStatuses([]Signature) ([]*SignatureStatus, error) { m.callCount++ return nil, nil @@ -119,10 +94,6 @@ func (m *mockClient) GetTokenAccountBalance(ed25519.PublicKey, Commitment) (uint m.callCount++ return 0, 0, nil } -func (m *mockClient) GetTokenAccountsByOwner(ed25519.PublicKey, ed25519.PublicKey) ([]ed25519.PublicKey, error) { - m.callCount++ - return nil, nil -} func (m *mockClient) GetTransaction(Signature, Commitment) (ConfirmedTransaction, error) { m.callCount++ return ConfirmedTransaction{}, nil diff --git a/solana/encoding.go b/solana/encoding.go index b707bd3..2d2e05a 100644 --- a/solana/encoding.go +++ b/solana/encoding.go @@ -3,7 +3,10 @@ package solana import ( "bytes" "crypto/ed25519" + "encoding/binary" "io" + "math" + "math/bits" "github.com/mr-tron/base58" "github.com/pkg/errors" @@ -13,15 +16,43 @@ import ( const ( messageVersionSerializationOffset = 127 + + // Each config mask bit corresponds to 4 bytes of config values, ordered by + // bit position. The priority fee is the only field spanning two bits. + transactionConfigBitPriorityFee uint8 = 0 + transactionConfigBitComputeUnitLimit uint8 = 2 + transactionConfigBitLoadedAccountsDataSizeLimit uint8 = 3 + transactionConfigBitHeapSize uint8 = 4 + + transactionConfigMaskPriorityFee uint32 = 0b11 << transactionConfigBitPriorityFee ) func (s TransactionSignature) ToBase58() string { return base58.Encode(s.Signature[:]) } -func (t Transaction) Marshal() []byte { +func (t Transaction) Marshal() ([]byte, error) { + message, err := t.Message.Marshal() + if err != nil { + return nil, err + } + b := bytes.NewBuffer(nil) + // v1 transactions place the message first, followed by signatures with no + // length prefix, since the count is implied by the message header. + if t.Message.Version == MessageVersion1 { + if len(t.Signatures) != int(t.Message.Header.NumSignatures) { + return nil, errors.Errorf("transaction has %d signatures, but header requires %d", len(t.Signatures), t.Message.Header.NumSignatures) + } + + _, _ = b.Write(message) + for _, s := range t.Signatures { + _, _ = b.Write(s[:]) + } + return b.Bytes(), nil + } + // Signatures _, _ = shortvec.EncodeLen(b, len(t.Signatures)) for _, s := range t.Signatures { @@ -29,12 +60,18 @@ func (t Transaction) Marshal() []byte { } // Message - _, _ = b.Write(t.Message.Marshal()) + _, _ = b.Write(message) - return b.Bytes() + return b.Bytes(), nil } func (t *Transaction) Unmarshal(b []byte) error { + // A legacy or v0 transaction starts with its signature count, which can + // never collide with the v1 version byte. + if len(b) > 0 && b[0] == byte(MessageVersion1+messageVersionSerializationOffset) { + return t.unmarshalV1(b) + } + buf := bytes.NewBuffer(b) sigLen, err := shortvec.DecodeLen(buf) @@ -52,17 +89,42 @@ func (t *Transaction) Unmarshal(b []byte) error { return (&t.Message).Unmarshal(buf.Bytes()) } -func (m Message) Marshal() []byte { +func (t *Transaction) unmarshalV1(b []byte) error { + buf := bytes.NewBuffer(b) + + if err := t.Message.unmarshalV1(buf); err != nil { + return err + } + + t.Signatures = make([]Signature, t.Message.Header.NumSignatures) + for i := range t.Signatures { + if _, err := io.ReadFull(buf, t.Signatures[i][:]); err != nil { + return errors.Wrapf(err, "failed to read signature at %d", i) + } + } + + if buf.Len() > 0 { + return errors.New("unexpected trailing data after signatures") + } + + return nil +} + +func (m Message) Marshal() ([]byte, error) { buf := bytes.NewBuffer(nil) switch m.Version { case MessageVersionLegacy: m.marshalLegacy(buf) case MessageVersion0: m.marshalV0(buf) + case MessageVersion1: + if err := m.marshalV1(buf); err != nil { + return nil, err + } default: - panic("unsupported message version") + return nil, errors.New("unsupported message version") } - return buf.Bytes() + return buf.Bytes(), nil } func (m *Message) marshalLegacy(b *bytes.Buffer) { @@ -117,6 +179,108 @@ func (m *Message) marshalV0(b *bytes.Buffer) { } } +func (m *Message) marshalV1(b *bytes.Buffer) error { + // Counts are fixed width in v1, so values that don't fit would otherwise + // silently wrap and produce a different transaction than was signed. + if len(m.Instructions) > math.MaxUint8 { + return errors.New("too many instructions for v1 message") + } + if len(m.Accounts) > math.MaxUint8 { + return errors.New("too many accounts for v1 message") + } + for _, i := range m.Instructions { + if len(i.Accounts) > math.MaxUint8 { + return errors.New("too many instruction accounts for v1 message") + } + if len(i.Data) > math.MaxUint16 { + return errors.New("instruction data too large for v1 message") + } + } + + // Version Number + _ = b.WriteByte(byte(m.Version + messageVersionSerializationOffset)) + + // Header + _ = b.WriteByte(m.Header.NumSignatures) + _ = b.WriteByte(m.Header.NumReadonlySigned) + _ = b.WriteByte(m.Header.NumReadOnly) + + // Transaction Config Mask + mask, configValues := m.Config.marshal() + _ = binary.Write(b, binary.LittleEndian, mask) + + // Lifetime Specifier + _, _ = b.Write(m.RecentBlockhash[:]) + + // Counts + _ = b.WriteByte(byte(len(m.Instructions))) + _ = b.WriteByte(byte(len(m.Accounts))) + + // Addresses + for _, a := range m.Accounts { + _, _ = b.Write(a) + } + + // Config Values + _, _ = b.Write(configValues) + + // Instruction Headers + for _, i := range m.Instructions { + _ = b.WriteByte(i.ProgramIndex) + _ = b.WriteByte(byte(len(i.Accounts))) + _ = binary.Write(b, binary.LittleEndian, uint16(len(i.Data))) + } + + // Instruction Payloads + for _, i := range m.Instructions { + _, _ = b.Write(i.Accounts) + _, _ = b.Write(i.Data) + } + + return nil +} + +// marshal returns the config mask and the concatenated config values, which +// are ordered by their bit position in the mask. +func (c TransactionConfig) marshal() (uint32, []byte) { + var mask uint32 + var values []byte + + for bit := uint8(0); bit < 32; bit++ { + switch bit { + case transactionConfigBitPriorityFee: + if c.PriorityFeeLamports != nil { + mask |= transactionConfigMaskPriorityFee + values = binary.LittleEndian.AppendUint64(values, *c.PriorityFeeLamports) + } + case transactionConfigBitPriorityFee + 1: + // Second half of the priority fee, written above + case transactionConfigBitComputeUnitLimit: + if c.ComputeUnitLimit != nil { + mask |= 1 << bit + values = binary.LittleEndian.AppendUint32(values, *c.ComputeUnitLimit) + } + case transactionConfigBitLoadedAccountsDataSizeLimit: + if c.LoadedAccountsDataSizeLimit != nil { + mask |= 1 << bit + values = binary.LittleEndian.AppendUint32(values, *c.LoadedAccountsDataSizeLimit) + } + case transactionConfigBitHeapSize: + if c.HeapSize != nil { + mask |= 1 << bit + values = binary.LittleEndian.AppendUint32(values, *c.HeapSize) + } + default: + if v, ok := c.unknown[bit]; ok { + mask |= 1 << bit + values = append(values, v[:]...) + } + } + } + + return mask, values +} + func (m *Message) Unmarshal(b []byte) (err error) { if len(b) == 0 { return errors.New("invalid byte buffer") @@ -126,6 +290,8 @@ func (m *Message) Unmarshal(b []byte) (err error) { m.Version = MessageVersionLegacy } else if b[0] == byte(MessageVersion0+messageVersionSerializationOffset) { m.Version = MessageVersion0 + } else if b[0] == byte(MessageVersion1+messageVersionSerializationOffset) { + m.Version = MessageVersion1 } else { return errors.New("unsupported message version") } @@ -137,6 +303,16 @@ func (m *Message) Unmarshal(b []byte) (err error) { return m.unmarshalLegacy(buf) case MessageVersion0: return m.unmarshalV0(buf) + case MessageVersion1: + if err := m.unmarshalV1(buf); err != nil { + return err + } + // Matches Transaction.Unmarshal, so a full v1 transaction isn't + // accepted as a message with its signatures silently ignored. + if buf.Len() > 0 { + return errors.New("unexpected trailing data after message") + } + return nil default: return errors.New("unsupported message version") } @@ -276,3 +452,148 @@ func (m *Message) unmarshalV0(buf *bytes.Buffer) (err error) { return nil } + +func (m *Message) unmarshalV1(buf *bytes.Buffer) (err error) { + // Message Version + version, err := buf.ReadByte() + if err != nil { + return errors.Wrap(err, "failed to read version byte") + } + if version != byte(MessageVersion1+messageVersionSerializationOffset) { + return errors.New("message version is not v1") + } + m.Version = MessageVersion1 + + // Header + if m.Header.NumSignatures, err = buf.ReadByte(); err != nil { + return errors.Wrap(err, "failed to read num signatures") + } + if m.Header.NumReadonlySigned, err = buf.ReadByte(); err != nil { + return errors.Wrap(err, "failed to read num readonly signatures") + } + if m.Header.NumReadOnly, err = buf.ReadByte(); err != nil { + return errors.Wrap(err, "failed to read num readonly") + } + + // Transaction Config Mask + var mask uint32 + if err = binary.Read(buf, binary.LittleEndian, &mask); err != nil { + return errors.Wrap(err, "failed to read transaction config mask") + } + if bits.OnesCount32(mask&transactionConfigMaskPriorityFee) == 1 { + return errors.New("transaction config mask sets only one priority fee bit") + } + + // Lifetime Specifier + if _, err = io.ReadFull(buf, m.RecentBlockhash[:]); err != nil { + return errors.Wrap(err, "failed to read lifetime specifier") + } + + // Counts + numInstructions, err := buf.ReadByte() + if err != nil { + return errors.Wrap(err, "failed to read num instructions") + } + numAddresses, err := buf.ReadByte() + if err != nil { + return errors.Wrap(err, "failed to read num addresses") + } + + // Addresses + m.Accounts = make([]ed25519.PublicKey, numAddresses) + for i := range m.Accounts { + m.Accounts[i] = make([]byte, ed25519.PublicKeySize) + if _, err = io.ReadFull(buf, m.Accounts[i]); err != nil { + return errors.Wrapf(err, "failed to read address at index %d", i) + } + } + + // Config Values + m.Config = TransactionConfig{} + for bit := uint8(0); bit < 32; bit++ { + if mask&(1<= len(m.Accounts) { + return errors.Errorf("program index out of range: %d:%d", i, headers[i].ProgramIndex) + } + } + + // Instruction Payloads + m.Instructions = make([]CompiledInstruction, numInstructions) + for i, header := range headers { + c := CompiledInstruction{ + ProgramIndex: header.ProgramIndex, + Accounts: make([]byte, header.NumAccounts), + Data: make([]byte, header.NumDataBytes), + } + + if _, err = io.ReadFull(buf, c.Accounts); err != nil { + return errors.Wrapf(err, "failed to read instruction[%d] accounts", i) + } + for _, index := range c.Accounts { + if int(index) >= len(m.Accounts) { + return errors.Errorf("account index out of range: %d:%d", i, index) + } + } + + if _, err = io.ReadFull(buf, c.Data); err != nil { + return errors.Wrapf(err, "failed to read instruction[%d] data", i) + } + + m.Instructions[i] = c + } + + return nil +} diff --git a/solana/system/program_test.go b/solana/system/program_test.go index 7267379..5b7c666 100644 --- a/solana/system/program_test.go +++ b/solana/system/program_test.go @@ -28,8 +28,10 @@ func TestCreateAccount(t *testing.T) { assert.Equal(t, size, instruction.Data[12:20]) assert.Equal(t, []byte(keys[2]), instruction.Data[20:52]) + marshalled, err := solana.NewLegacyTransaction(keys[0], instruction).Marshal() + require.NoError(t, err) var tx solana.Transaction - require.NoError(t, tx.Unmarshal(solana.NewLegacyTransaction(keys[0], instruction).Marshal())) + require.NoError(t, tx.Unmarshal(marshalled)) decompiled, err := DecompileCreateAccount(tx.Message, 0) require.NoError(t, err) @@ -80,8 +82,10 @@ func TestTransfer(t *testing.T) { assert.Equal(t, command, instruction.Data[0:4]) assert.Equal(t, lamports, instruction.Data[4:12]) + marshalled, err := solana.NewLegacyTransaction(keys[0], instruction).Marshal() + require.NoError(t, err) var tx solana.Transaction - require.NoError(t, tx.Unmarshal(solana.NewLegacyTransaction(keys[0], instruction).Marshal())) + require.NoError(t, tx.Unmarshal(marshalled)) decompiled, err := DecompileTransfer(tx.Message, 0) require.NoError(t, err) diff --git a/solana/transaction.go b/solana/transaction.go index b315e9a..4780fd1 100644 --- a/solana/transaction.go +++ b/solana/transaction.go @@ -5,16 +5,37 @@ import ( "crypto/ed25519" "crypto/sha256" "fmt" + "maps" "sort" "strings" "github.com/mr-tron/base58/base58" "github.com/pkg/errors" + + "github.com/code-payments/ocp-server/pointer" ) const ( - // MaxTransactionSize taken from: https://github.com/solana-labs/solana/blob/39b3ac6a8d29e14faa1de73d8b46d390ad41797b/sdk/src/packet.rs#L9-L13 - MaxTransactionSize = 1232 + // MaxLegacyTransactionSize applies to both legacy and v0 transactions. + // Taken from: https://github.com/solana-labs/solana/blob/39b3ac6a8d29e14faa1de73d8b46d390ad41797b/sdk/src/packet.rs#L9-L13 + MaxLegacyTransactionSize = 1232 + + // MaxV1TransactionSize taken from: https://github.com/solana-foundation/solana-improvement-documents/blob/main/proposals/0296-larger-transactions.md + MaxV1TransactionSize = 4096 + + // Sanitization constraints for v1 transactions from SIMD-0385 + maxV1Signatures = 12 + maxV1Accounts = 64 + maxV1Instructions = 64 + minV1HeapSize = 32 * 1024 + maxV1HeapSize = 256 * 1024 + v1HeapSizeMultiple = 1024 + + // Runtime maximums for v1 config requests. Larger values aren't a + // sanitization failure, and are silently clamped by the runtime. + // Taken from: https://github.com/anza-xyz/agave/blob/443ba0e197adafd0206e651f86934d22a44d2729/program-runtime/src/execution_budget.rs#L26-L41 + maxV1ComputeUnitLimit = 1_400_000 + maxV1LoadedAccountsDataSizeLimit = 64 * 1024 * 1024 ) type Signature [ed25519.SignatureSize]byte @@ -25,6 +46,7 @@ type MessageVersion uint8 const ( MessageVersionLegacy MessageVersion = iota MessageVersion0 + MessageVersion1 ) type Header struct { @@ -40,6 +62,41 @@ type Message struct { RecentBlockhash Blockhash Instructions []CompiledInstruction AddressTableLookups []MessageAddressTableLookup + Config TransactionConfig +} + +// TransactionConfig is the set of fee and resource requests carried in a v1 +// message, replacing ComputeBudgetProgram instructions (which are ignored for +// configuration in v1 transactions). +// +// Unset fields use the minimum allowed value. Notably, an unset compute unit +// limit or loaded accounts data size limit is zero, not the runtime default. +// Fields are optional here so decoded transactions round trip exactly, but +// NewV1Transaction requires both limits. +type TransactionConfig struct { + PriorityFeeLamports *uint64 + ComputeUnitLimit *uint32 + LoadedAccountsDataSizeLimit *uint32 + HeapSize *uint32 + + // Values for mask bits this package doesn't understand, keyed by bit, so + // transactions using config fields from future SIMDs still decode and + // round trip exactly. + unknown map[uint8][4]byte +} + +func (c TransactionConfig) clone() TransactionConfig { + cloned := TransactionConfig{ + PriorityFeeLamports: pointer.Uint64Copy(c.PriorityFeeLamports), + ComputeUnitLimit: pointer.Uint32Copy(c.ComputeUnitLimit), + LoadedAccountsDataSizeLimit: pointer.Uint32Copy(c.LoadedAccountsDataSizeLimit), + HeapSize: pointer.Uint32Copy(c.HeapSize), + } + if c.unknown != nil { + cloned.unknown = make(map[uint8][4]byte, len(c.unknown)) + maps.Copy(cloned.unknown, c.unknown) + } + return cloned } type MessageAddressTableLookup struct { @@ -144,6 +201,62 @@ func NewLegacyTransaction(payer ed25519.PublicKey, instructions ...Instruction) } } +// NewV1Transaction builds a v1 transaction as specified in SIMD-0385. Account +// ordering is unchanged from legacy transactions, but address lookup tables +// are not supported and all accounts are included inline. +// +// The config must set nonzero compute unit and loaded accounts data size +// limits, because v1 treats them as zero when unset, which fails every +// transaction. Limits above the runtime maximum are rejected rather than left +// for the runtime to clamp, since they indicate a misconfigured budget. +// Transactions violating the SIMD-0385 sanitization constraints are rejected, +// since they would never be included in a block. +func NewV1Transaction(payer ed25519.PublicKey, config TransactionConfig, instructions ...Instruction) (Transaction, error) { + // Copy so later changes to the caller's config can't alter the signed message + config = config.clone() + + if config.ComputeUnitLimit == nil { + return Transaction{}, errors.New("compute unit limit is required") + } + if *config.ComputeUnitLimit == 0 || *config.ComputeUnitLimit > maxV1ComputeUnitLimit { + return Transaction{}, errors.Errorf("compute unit limit must be between 1 and %d", maxV1ComputeUnitLimit) + } + if config.LoadedAccountsDataSizeLimit == nil { + return Transaction{}, errors.New("loaded accounts data size limit is required") + } + if *config.LoadedAccountsDataSizeLimit == 0 || *config.LoadedAccountsDataSizeLimit > maxV1LoadedAccountsDataSizeLimit { + return Transaction{}, errors.Errorf("loaded accounts data size limit must be between 1 and %d", maxV1LoadedAccountsDataSizeLimit) + } + if config.HeapSize != nil { + if *config.HeapSize < minV1HeapSize || *config.HeapSize > maxV1HeapSize || *config.HeapSize%v1HeapSizeMultiple != 0 { + return Transaction{}, errors.Errorf("heap size must be a multiple of %d between %d and %d", v1HeapSizeMultiple, minV1HeapSize, maxV1HeapSize) + } + } + + txn := NewLegacyTransaction(payer, instructions...) + txn.Message.Version = MessageVersion1 + txn.Message.Config = config + + if txn.Message.Header.NumSignatures > maxV1Signatures { + return Transaction{}, errors.Errorf("transaction exceeds %d signatures", maxV1Signatures) + } + if len(txn.Message.Accounts) > maxV1Accounts { + return Transaction{}, errors.Errorf("transaction exceeds %d accounts", maxV1Accounts) + } + if len(txn.Message.Instructions) > maxV1Instructions { + return Transaction{}, errors.Errorf("transaction exceeds %d instructions", maxV1Instructions) + } + marshalled, err := txn.Marshal() + if err != nil { + return Transaction{}, err + } + if len(marshalled) > MaxV1TransactionSize { + return Transaction{}, errors.Errorf("transaction size %d exceeds %d bytes", len(marshalled), MaxV1TransactionSize) + } + + return txn, nil +} + func NewV0Transaction(payer ed25519.PublicKey, addressLookupTables []AddressLookupTable, instructions []Instruction) Transaction { accounts := []AccountMeta{ { @@ -314,7 +427,22 @@ func (t *Transaction) String() string { sb.WriteString(fmt.Sprintf(" Accounts: %v\n", t.Message.Instructions[i].Accounts)) sb.WriteString(fmt.Sprintf(" Data: %v\n", t.Message.Instructions[i].Data)) } - if t.Message.Version >= MessageVersion0 { + if t.Message.Version == MessageVersion1 { + sb.WriteString(" Config:\n") + if t.Message.Config.PriorityFeeLamports != nil { + sb.WriteString(fmt.Sprintf(" PriorityFeeLamports: %d\n", *t.Message.Config.PriorityFeeLamports)) + } + if t.Message.Config.ComputeUnitLimit != nil { + sb.WriteString(fmt.Sprintf(" ComputeUnitLimit: %d\n", *t.Message.Config.ComputeUnitLimit)) + } + if t.Message.Config.LoadedAccountsDataSizeLimit != nil { + sb.WriteString(fmt.Sprintf(" LoadedAccountsDataSizeLimit: %d\n", *t.Message.Config.LoadedAccountsDataSizeLimit)) + } + if t.Message.Config.HeapSize != nil { + sb.WriteString(fmt.Sprintf(" HeapSize: %d\n", *t.Message.Config.HeapSize)) + } + } + if t.Message.Version == MessageVersion0 { sb.WriteString(" Address Table Lookups:\n") for i := range t.Message.AddressTableLookups { sb.WriteString(fmt.Sprintf(" %s:\n", base58.Encode(t.Message.AddressTableLookups[i].PublicKey))) @@ -331,7 +459,10 @@ func (t *Transaction) SetBlockhash(bh Blockhash) { } func (t *Transaction) Sign(signers ...ed25519.PrivateKey) error { - messageBytes := t.Message.Marshal() + messageBytes, err := t.Message.Marshal() + if err != nil { + return err + } for _, s := range signers { pub := s.Public().(ed25519.PublicKey) @@ -397,6 +528,8 @@ func (v MessageVersion) String() string { return "legacy" case MessageVersion0: return "v0" + case MessageVersion1: + return "v1" } return "unknown" } diff --git a/solana/transaction_test.go b/solana/transaction_test.go index 011adc3..ef934ab 100644 --- a/solana/transaction_test.go +++ b/solana/transaction_test.go @@ -4,12 +4,15 @@ import ( "bytes" "crypto/ed25519" "encoding/base64" + "math" "math/rand" "sort" "testing" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + + "github.com/code-payments/ocp-server/pointer" ) // Taken from: https://github.com/solana-labs/solana/blob/14339dec0a960e8161d1165b6a8e5cfb73e78f23/sdk/src/transaction.rs#L523 @@ -41,7 +44,7 @@ func TestLegacyTransaction_CrossImpl(t *testing.T) { generated, err := base64.StdEncoding.DecodeString(rustGenerated) require.NoError(t, err) - assert.Equal(t, generated, tx.Marshal()) + assert.Equal(t, generated, mustMarshal(t, tx)) } func TestLegacyTransaction_GenerateValidCrossImpl(t *testing.T) { @@ -61,7 +64,7 @@ func TestLegacyTransaction_GenerateValidCrossImpl(t *testing.T) { ), ) require.NoError(t, tx.Sign(keypair)) - assert.Equal(t, rustGeneratedAdjusted, base64.StdEncoding.EncodeToString(tx.Marshal())) + assert.Equal(t, rustGeneratedAdjusted, base64.StdEncoding.EncodeToString(mustMarshal(t, tx))) } func TestLegacyTransaction_EmptyAccount(t *testing.T) { @@ -81,7 +84,7 @@ func TestLegacyTransaction_EmptyAccount(t *testing.T) { assert.NoError(t, tx.Sign(priv)) var rtt Transaction - assert.NoError(t, rtt.Unmarshal(tx.Marshal())) + assert.NoError(t, rtt.Unmarshal(mustMarshal(t, tx))) } func TestLegacyTransaction_MarshalRoundTrip(t *testing.T) { @@ -90,7 +93,7 @@ func TestLegacyTransaction_MarshalRoundTrip(t *testing.T) { require.NoError(t, err) var txn Transaction require.NoError(t, txn.Unmarshal(decoded)) - assert.Equal(t, decoded, txn.Marshal()) + assert.Equal(t, decoded, mustMarshal(t, txn)) } func TestLegacyTransaction_MissingBlockhash(t *testing.T) { @@ -110,7 +113,7 @@ func TestLegacyTransaction_MissingBlockhash(t *testing.T) { assert.NoError(t, tx.Sign(priv)) var rtt Transaction - assert.NoError(t, rtt.Unmarshal(tx.Marshal())) + assert.NoError(t, rtt.Unmarshal(mustMarshal(t, tx))) } func TestLegacyTransaction_InvalidAccounts(t *testing.T) { @@ -124,7 +127,7 @@ func TestLegacyTransaction_InvalidAccounts(t *testing.T) { ), ) tx.Message.Instructions[0].ProgramIndex = 2 - assert.Error(t, tx.Unmarshal(tx.Marshal())) + assert.Error(t, tx.Unmarshal(mustMarshal(t, tx))) tx = NewLegacyTransaction( public(keys[0]), @@ -166,7 +169,7 @@ func TestLegacyTransaction_SingleInstruction(t *testing.T) { assert.EqualValues(t, 1, tx.Message.Header.NumReadonlySigned) assert.EqualValues(t, 2, tx.Message.Header.NumReadOnly) - message := tx.Message.Marshal() + message := mustMarshalMessage(t, tx.Message) assert.True(t, ed25519.Verify(public(payer), message, tx.Signatures[0][:])) assert.True(t, ed25519.Verify(public(keys[3]), message, tx.Signatures[1][:])) @@ -235,7 +238,7 @@ func TestLegacyTransaction_DuplicateKeys(t *testing.T) { assert.EqualValues(t, 1, tx.Message.Header.NumReadonlySigned) assert.EqualValues(t, 1, tx.Message.Header.NumReadOnly) - message := tx.Message.Marshal() + message := mustMarshalMessage(t, tx.Message) assert.True(t, ed25519.Verify(public(payer), message, tx.Signatures[0][:])) assert.True(t, ed25519.Verify(public(keys[0]), message, tx.Signatures[1][:])) @@ -321,7 +324,7 @@ func TestLegacyTransaction_MultiInstruction(t *testing.T) { assert.EqualValues(t, 0, tx.Message.Header.NumReadonlySigned) assert.EqualValues(t, 3, tx.Message.Header.NumReadOnly) - message := tx.Message.Marshal() + message := mustMarshalMessage(t, tx.Message) assert.True(t, ed25519.Verify(public(payer), message, tx.Signatures[0][:])) assert.True(t, ed25519.Verify(public(keys[0]), message, tx.Signatures[1][:])) @@ -440,7 +443,7 @@ func TestV0Transaction_MultipleAlts(t *testing.T) { assert.Equal(t, bh, tx.Message.RecentBlockhash) - message := tx.Message.Marshal() + message := mustMarshalMessage(t, tx.Message) assert.True(t, ed25519.Verify(public(payer), message, tx.Signatures[0][:])) assert.True(t, ed25519.Verify(public(accountSigner), message, tx.Signatures[1][:])) @@ -479,7 +482,328 @@ func TestV0Transaction_MarshalRoundTrip(t *testing.T) { require.NoError(t, err) var txn Transaction require.NoError(t, txn.Unmarshal(decoded)) - assert.Equal(t, decoded, txn.Marshal()) + assert.Equal(t, decoded, mustMarshal(t, txn)) +} + +// Generated with @solana/kit 8 (see TestV1Transaction_CrossImpl for the inputs) +var kitGeneratedV1 = map[string]string{ + "full": "gQIBAR8AAAAFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQIEiojj3XQJ8ZX9UtstPLpdcspnCb8dlBIb83SIAbQPb1yBOXcOqH0XX1ajVGbDTH7My42KkbTuN6Jd9g9bj8mzlAMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBASIEwAAAAAAAEANAwCghgEAAAABAAMDAwADASwBAAECAQIDAgABAgMEBQYHCAkKCwwNDg8QERITFBUWFxgZGhscHR4fICEiIyQlJicoKSorLC0uLzAxMjM0NTY3ODk6Ozw9Pj9AQUJDREVGR0hJSktMTU5PUFFSU1RVVldYWVpbXF1eX2BhYmNkZWZnaGlqa2xtbm9wcXJzdHV2d3h5ent8fX5/gIGCg4SFhoeIiYqLjI2Oj5CRkpOUlZaXmJmam5ydnp+goaKjpKWmp6ipqqusra6vsLGys7S1tre4ubq7vL2+v8DBwsPExcbHyMnKy8zNzs/Q0dLT1NXW19jZ2tvc3d7f4OHi4+Tl5ufo6err7O3u7/Dx8vP09fb3+Pn6+/z9/v8AAQIDBAUGBwgJCgsMDQ4PEBESExQVFhcYGRobHB0eHyAhIiMkJSYnKCkqKwadBkWJOSTqIre0TpnD2/pEhHIliTVG9q0VmMxZ5Mz5+BhjNNIomISK6k39PbmWxc6/z2bFjYF5UKDjzbEHWQRYnYl9XMa2so4QHF4tQfq45lXaxyETBhfbL3V64qt3bvHqAomU9jRkhk89YPkuvmIZN+mJaoH/DIY8R4GcxygF", + "limits-only": "gQIBAQwAAAAFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQIEiojj3XQJ8ZX9UtstPLpdcspnCb8dlBIb83SIAbQPb1yBOXcOqH0XX1ajVGbDTH7My42KkbTuN6Jd9g9bj8mzlAMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBARADQMAoIYBAAMDAwADASwBAAECAQIDAgABAgMEBQYHCAkKCwwNDg8QERITFBUWFxgZGhscHR4fICEiIyQlJicoKSorLC0uLzAxMjM0NTY3ODk6Ozw9Pj9AQUJDREVGR0hJSktMTU5PUFFSU1RVVldYWVpbXF1eX2BhYmNkZWZnaGlqa2xtbm9wcXJzdHV2d3h5ent8fX5/gIGCg4SFhoeIiYqLjI2Oj5CRkpOUlZaXmJmam5ydnp+goaKjpKWmp6ipqqusra6vsLGys7S1tre4ubq7vL2+v8DBwsPExcbHyMnKy8zNzs/Q0dLT1NXW19jZ2tvc3d7f4OHi4+Tl5ufo6err7O3u7/Dx8vP09fb3+Pn6+/z9/v8AAQIDBAUGBwgJCgsMDQ4PEBESExQVFhcYGRobHB0eHyAhIiMkJSYnKCkqKwE8y4qiLSYJrKSBQnDToZidgFSVl7PhZv5Gwbq0tmlW4g6d42zt8uQLEsdSANu13g41ITmR+fooHRO2CB3w0AS2t7xsuZBA4mrDP7O+Aq52TyOMK2QTm+RSHX9IRdKNMEHLiASI3fj13ipTPUY5vrUBDL20WjXNFgwtrPuQ44QI", + "limits-and-heap": "gQIBARwAAAAFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQIEiojj3XQJ8ZX9UtstPLpdcspnCb8dlBIb83SIAbQPb1yBOXcOqH0XX1ajVGbDTH7My42KkbTuN6Jd9g9bj8mzlAMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBARADQMAoIYBAAAAAQADAwMAAwEsAQABAgECAwIAAQIDBAUGBwgJCgsMDQ4PEBESExQVFhcYGRobHB0eHyAhIiMkJSYnKCkqKywtLi8wMTIzNDU2Nzg5Ojs8PT4/QEFCQ0RFRkdISUpLTE1OT1BRUlNUVVZXWFlaW1xdXl9gYWJjZGVmZ2hpamtsbW5vcHFyc3R1dnd4eXp7fH1+f4CBgoOEhYaHiImKi4yNjo+QkZKTlJWWl5iZmpucnZ6foKGio6SlpqeoqaqrrK2ur7CxsrO0tba3uLm6u7y9vr/AwcLDxMXGx8jJysvMzc7P0NHS09TV1tfY2drb3N3e3+Dh4uPk5ebn6Onq6+zt7u/w8fLz9PX29/j5+vv8/f7/AAECAwQFBgcICQoLDA0ODxAREhMUFRYXGBkaGxwdHh8gISIjJCUmJygpKisjFDtun4YLWCyoXjNiGkdZqTNhrKcegvxZZJmLerDD3X4uzI+ShZm8rJghk5SBkFAwUG1UeLxc96xPZKCgysQLlaE5PLTk7rlpzjkfFni+07lAdfaB+lJYV+tyHMu/BGE+g3aTmHZO+xeYJg67HrNfhRJkAEA33w2dpkbuaFtaCQ==", +} + +// Generated with @solana/kit 8 using the same inputs as kitGeneratedV1, but with +// limits unset. These can't be built with NewV1Transaction, but must decode. +var kitGeneratedV1WithoutLimits = map[string]string{ + "cu-limit-and-heap": "gQIBARQAAAAFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQIEiojj3XQJ8ZX9UtstPLpdcspnCb8dlBIb83SIAbQPb1yBOXcOqH0XX1ajVGbDTH7My42KkbTuN6Jd9g9bj8mzlAMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBARADQMAAAABAAMDAwADASwBAAECAQIDAgABAgMEBQYHCAkKCwwNDg8QERITFBUWFxgZGhscHR4fICEiIyQlJicoKSorLC0uLzAxMjM0NTY3ODk6Ozw9Pj9AQUJDREVGR0hJSktMTU5PUFFSU1RVVldYWVpbXF1eX2BhYmNkZWZnaGlqa2xtbm9wcXJzdHV2d3h5ent8fX5/gIGCg4SFhoeIiYqLjI2Oj5CRkpOUlZaXmJmam5ydnp+goaKjpKWmp6ipqqusra6vsLGys7S1tre4ubq7vL2+v8DBwsPExcbHyMnKy8zNzs/Q0dLT1NXW19jZ2tvc3d7f4OHi4+Tl5ufo6err7O3u7/Dx8vP09fb3+Pn6+/z9/v8AAQIDBAUGBwgJCgsMDQ4PEBESExQVFhcYGRobHB0eHyAhIiMkJSYnKCkqK5Vge5MRSqEEShsH3fq8JxeGd6NGbSH3emDtsNI38/tF1zwrIClU2Q4k/0LLCSU3JoUGDrZhOeody++FF+kMSAZHeHJIGGKcvfgZzlJFs86l4tBomCdR0LqdQVq+rTJHgawAvzFnXQgmrT8gOBbbR6nxnPDBlNNRYD+PmTI+WmAH", + "none": "gQIBAQAAAAAFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQUFBQIEiojj3XQJ8ZX9UtstPLpdcspnCb8dlBIb83SIAbQPb1yBOXcOqH0XX1ajVGbDTH7My42KkbTuN6Jd9g9bj8mzlAMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDAwMDBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQEBAQDAwMAAwEsAQABAgECAwIAAQIDBAUGBwgJCgsMDQ4PEBESExQVFhcYGRobHB0eHyAhIiMkJSYnKCkqKywtLi8wMTIzNDU2Nzg5Ojs8PT4/QEFCQ0RFRkdISUpLTE1OT1BRUlNUVVZXWFlaW1xdXl9gYWJjZGVmZ2hpamtsbW5vcHFyc3R1dnd4eXp7fH1+f4CBgoOEhYaHiImKi4yNjo+QkZKTlJWWl5iZmpucnZ6foKGio6SlpqeoqaqrrK2ur7CxsrO0tba3uLm6u7y9vr/AwcLDxMXGx8jJysvMzc7P0NHS09TV1tfY2drb3N3e3+Dh4uPk5ebn6Onq6+zt7u/w8fLz9PX29/j5+vv8/f7/AAECAwQFBgcICQoLDA0ODxAREhMUFRYXGBkaGxwdHh8gISIjJCUmJygpKivDPe8U3WKLcDb+WjfwvtMcYMg6Vb6h3gPBYfcmIpybp1b4FpmQF9XPqRjnOzi5YJCgcXeZrp5lSLWgsvWeJlQAZXRbwxrznjGCqvhnbzKRIAEp7jqpmMkmxWQg8uXHgacrBJOYhRn3s8A4q0xykHvuSTnPHLJq4uUiD5j4Xhv/Aw==", +} + +func TestV1Transaction_CrossImpl(t *testing.T) { + filled := func(b byte) []byte { return bytes.Repeat([]byte{b}, 32) } + + payer := ed25519.NewKeyFromSeed(filled(1)) + signer := ed25519.NewKeyFromSeed(filled(2)) + writable := ed25519.PublicKey(filled(3)) + program := ed25519.PublicKey(filled(4)) + + var bh Blockhash + copy(bh[:], filled(5)) + + largeData := make([]byte, 300) + for i := range largeData { + largeData[i] = byte(i) + } + + for name, config := range map[string]TransactionConfig{ + "full": { + PriorityFeeLamports: pointer.Uint64(5000), + ComputeUnitLimit: pointer.Uint32(200_000), + LoadedAccountsDataSizeLimit: pointer.Uint32(100_000), + HeapSize: pointer.Uint32(64 * 1024), + }, + "limits-only": { + ComputeUnitLimit: pointer.Uint32(200_000), + LoadedAccountsDataSizeLimit: pointer.Uint32(100_000), + }, + "limits-and-heap": { + ComputeUnitLimit: pointer.Uint32(200_000), + LoadedAccountsDataSizeLimit: pointer.Uint32(100_000), + HeapSize: pointer.Uint32(64 * 1024), + }, + } { + t.Run(name, func(t *testing.T) { + tx, err := NewV1Transaction( + public(payer), + config, + NewInstruction( + program, + []byte{1, 2, 3}, + NewAccountMeta(public(payer), true), + NewReadonlyAccountMeta(public(signer), true), + NewAccountMeta(writable, false), + ), + NewInstruction( + program, + largeData, + NewAccountMeta(writable, false), + ), + ) + require.NoError(t, err) + tx.SetBlockhash(bh) + require.NoError(t, tx.Sign(payer, signer)) + + assert.Equal(t, kitGeneratedV1[name], base64.StdEncoding.EncodeToString(mustMarshal(t, tx))) + + message := mustMarshalMessage(t, tx.Message) + assert.True(t, ed25519.Verify(public(payer), message, tx.Signatures[0][:])) + assert.True(t, ed25519.Verify(public(signer), message, tx.Signatures[1][:])) + + var decoded Transaction + require.NoError(t, decoded.Unmarshal(mustMarshal(t, tx))) + assert.Equal(t, tx, decoded) + }) + } +} + +func TestV1Transaction_MarshalRoundTripWithoutLimits(t *testing.T) { + for name, encoded := range kitGeneratedV1WithoutLimits { + t.Run(name, func(t *testing.T) { + decoded, err := base64.StdEncoding.DecodeString(encoded) + require.NoError(t, err) + + var tx Transaction + require.NoError(t, tx.Unmarshal(decoded)) + assert.Nil(t, tx.Message.Config.LoadedAccountsDataSizeLimit) + assert.Equal(t, decoded, mustMarshal(t, tx)) + }) + } +} + +func TestV1Transaction_Builder(t *testing.T) { + keys := generateKeys(t, 4) + payer, program, writable, readonly := keys[0], keys[1], keys[2], keys[3] + + ixn := NewInstruction( + public(program), + []byte{1}, + NewReadonlyAccountMeta(public(readonly), false), + NewAccountMeta(public(writable), false), + NewAccountMeta(public(writable), false), + ) + + // Both limits are required, since v1 treats unset limits as zero + for _, config := range []TransactionConfig{ + {}, + {ComputeUnitLimit: pointer.Uint32(10_000)}, + {LoadedAccountsDataSizeLimit: pointer.Uint32(20_000)}, + } { + _, err := NewV1Transaction(public(payer), config, ixn) + assert.Error(t, err) + } + + config := TransactionConfig{ + PriorityFeeLamports: pointer.Uint64(1), + ComputeUnitLimit: pointer.Uint32(10_000), + LoadedAccountsDataSizeLimit: pointer.Uint32(20_000), + } + tx, err := NewV1Transaction(public(payer), config, ixn) + require.NoError(t, err) + + assert.Equal(t, MessageVersion1, tx.Message.Version) + assert.Equal(t, config, tx.Message.Config) + assert.Empty(t, tx.Message.AddressTableLookups) + require.Len(t, tx.Signatures, 1) + + // Addresses are deduplicated, which v1 requires + require.Len(t, tx.Message.Accounts, 4) + assert.Equal(t, public(payer), tx.Message.Accounts[0]) + assert.Equal(t, public(writable), tx.Message.Accounts[1]) + assert.EqualValues(t, 1, tx.Message.Header.NumSignatures) + assert.EqualValues(t, 0, tx.Message.Header.NumReadonlySigned) + assert.EqualValues(t, 2, tx.Message.Header.NumReadOnly) + + // The wire format leads with the version byte and ends with signatures + require.NoError(t, tx.Sign(payer)) + marshalled := mustMarshal(t, tx) + assert.EqualValues(t, 129, marshalled[0]) + assert.Equal(t, tx.Signatures[0][:], marshalled[len(marshalled)-ed25519.SignatureSize:]) + assert.Len(t, marshalled, len(mustMarshalMessage(t, tx.Message))+ed25519.SignatureSize) + + // Mutating the caller's config doesn't alter the signed transaction + *config.PriorityFeeLamports = 2 + *config.ComputeUnitLimit = 30_000 + *config.LoadedAccountsDataSizeLimit = 40_000 + assert.EqualValues(t, 1, *tx.Message.Config.PriorityFeeLamports) + assert.EqualValues(t, 10_000, *tx.Message.Config.ComputeUnitLimit) + assert.EqualValues(t, 20_000, *tx.Message.Config.LoadedAccountsDataSizeLimit) + assert.Equal(t, marshalled, mustMarshal(t, tx)) + assert.True(t, ed25519.Verify(public(payer), mustMarshalMessage(t, tx.Message), tx.Signatures[0][:])) +} + +func TestV1Transaction_BuilderConstraints(t *testing.T) { + keys := generateKeys(t, 2) + payer, program := keys[0], keys[1] + + limits := TransactionConfig{ + ComputeUnitLimit: pointer.Uint32(10_000), + LoadedAccountsDataSizeLimit: pointer.Uint32(20_000), + } + withHeap := func(size uint32) TransactionConfig { + config := limits + config.HeapSize = pointer.Uint32(size) + return config + } + withComputeUnitLimit := func(limit uint32) TransactionConfig { + config := limits + config.ComputeUnitLimit = pointer.Uint32(limit) + return config + } + withLoadedAccountsDataSizeLimit := func(limit uint32) TransactionConfig { + config := limits + config.LoadedAccountsDataSizeLimit = pointer.Uint32(limit) + return config + } + accountsInstruction := func(n int, signer bool) Instruction { + metas := make([]AccountMeta, n) + for i := range metas { + metas[i] = NewAccountMeta(public(generateKeys(t, 1)[0]), signer) + } + return NewInstruction(public(program), []byte{1}, metas...) + } + repeatInstruction := func(n int) []Instruction { + ixns := make([]Instruction, n) + for i := range ixns { + ixns[i] = NewInstruction(public(program), []byte{byte(i)}) + } + return ixns + } + dataInstruction := func(n int) Instruction { + return NewInstruction(public(program), make([]byte, n)) + } + + empty, err := NewV1Transaction(public(payer), limits, dataInstruction(0)) + require.NoError(t, err) + maxData := MaxV1TransactionSize - len(mustMarshal(t, empty)) + + // The payer and program each occupy an account, and the payer is a signer + for name, tc := range map[string]struct { + config TransactionConfig + instructions []Instruction + valid bool + }{ + "max signatures": {limits, []Instruction{accountsInstruction(maxV1Signatures-1, true)}, true}, + "too many signatures": {limits, []Instruction{accountsInstruction(maxV1Signatures, true)}, false}, + "max accounts": {limits, []Instruction{accountsInstruction(maxV1Accounts-2, false)}, true}, + "too many accounts": {limits, []Instruction{accountsInstruction(maxV1Accounts-1, false)}, false}, + "max instructions": {limits, repeatInstruction(maxV1Instructions), true}, + "too many instructions": {limits, repeatInstruction(maxV1Instructions + 1), false}, + "min heap size": {withHeap(minV1HeapSize), repeatInstruction(1), true}, + "max heap size": {withHeap(maxV1HeapSize), repeatInstruction(1), true}, + "heap size too small": {withHeap(minV1HeapSize - v1HeapSizeMultiple), repeatInstruction(1), false}, + "heap size too large": {withHeap(maxV1HeapSize + v1HeapSizeMultiple), repeatInstruction(1), false}, + "heap size unaligned": {withHeap(minV1HeapSize + 1), repeatInstruction(1), false}, + + "zero compute unit limit": {withComputeUnitLimit(0), repeatInstruction(1), false}, + "max compute unit limit": {withComputeUnitLimit(maxV1ComputeUnitLimit), repeatInstruction(1), true}, + "compute unit limit too high": {withComputeUnitLimit(maxV1ComputeUnitLimit + 1), repeatInstruction(1), false}, + + "zero loaded accounts data size limit": {withLoadedAccountsDataSizeLimit(0), repeatInstruction(1), false}, + "max loaded accounts data size limit": {withLoadedAccountsDataSizeLimit(maxV1LoadedAccountsDataSizeLimit), repeatInstruction(1), true}, + "loaded accounts data size limit too high": {withLoadedAccountsDataSizeLimit(maxV1LoadedAccountsDataSizeLimit + 1), repeatInstruction(1), false}, + + "max size": {limits, []Instruction{dataInstruction(maxData)}, true}, + "too large": {limits, []Instruction{dataInstruction(maxData + 1)}, false}, + } { + t.Run(name, func(t *testing.T) { + tx, err := NewV1Transaction(public(payer), tc.config, tc.instructions...) + if !tc.valid { + assert.Error(t, err) + return + } + require.NoError(t, err) + assert.LessOrEqual(t, len(mustMarshal(t, tx)), MaxV1TransactionSize) + }) + } +} + +func TestV1Transaction_UnmarshalInvalid(t *testing.T) { + valid, err := base64.StdEncoding.DecodeString(kitGeneratedV1["full"]) + require.NoError(t, err) + + // version + header, then mask + lifetime + counts, then 4 addresses and + // 20 bytes of config values + const maskOffset = 1 + 3 + const programIndexOffset = maskOffset + 4 + 32 + 2 + 4*32 + 20 + + for name, mutate := range map[string]func([]byte) []byte{ + "truncated signature": func(b []byte) []byte { return b[:len(b)-1] }, + "trailing data": func(b []byte) []byte { return append(b, 0) }, + "truncated message": func(b []byte) []byte { return b[:100] }, + "single priority fee bit": func(b []byte) []byte { + b[maskOffset] &^= 0b10 + return b + }, + "program index out of range": func(b []byte) []byte { + b[programIndexOffset] = 4 + return b + }, + } { + t.Run(name, func(t *testing.T) { + var tx Transaction + assert.Error(t, tx.Unmarshal(mutate(bytes.Clone(valid)))) + }) + } +} + +func TestV1Message_UnmarshalTrailingData(t *testing.T) { + full, err := base64.StdEncoding.DecodeString(kitGeneratedV1["full"]) + require.NoError(t, err) + + var tx Transaction + require.NoError(t, tx.Unmarshal(full)) + message := mustMarshalMessage(t, tx.Message) + + var decoded Message + require.NoError(t, decoded.Unmarshal(message)) + assert.Equal(t, tx.Message, decoded) + + // A full transaction, or any extra bytes, isn't a valid message + assert.Error(t, (&Message{}).Unmarshal(full)) + assert.Error(t, (&Message{}).Unmarshal(append(bytes.Clone(message), 0))) +} + +func TestV1Transaction_UnknownConfigBits(t *testing.T) { + valid, err := base64.StdEncoding.DecodeString(kitGeneratedV1["full"]) + require.NoError(t, err) + + // Every mask bit carries 4 bytes, so config fields from future SIMDs can + // be carried through without understanding them. Set bits 5 and 31, and + // insert their values after the 20 bytes of known config values. + const maskOffset = 1 + 3 + const unknownValuesOffset = maskOffset + 4 + 32 + 2 + 4*32 + 20 + + mutated := bytes.Clone(valid) + mutated[maskOffset] |= 1 << 5 + mutated[maskOffset+3] |= 1 << 7 + unknownValues := []byte{0xa, 0xb, 0xc, 0xd, 0x1, 0x2, 0x3, 0x4} + mutated = append(mutated[:unknownValuesOffset], append(unknownValues, mutated[unknownValuesOffset:]...)...) + + var tx Transaction + require.NoError(t, tx.Unmarshal(mutated)) + assert.Equal(t, mutated, mustMarshal(t, tx)) + + assert.EqualValues(t, 5000, *tx.Message.Config.PriorityFeeLamports) + assert.EqualValues(t, 200_000, *tx.Message.Config.ComputeUnitLimit) + assert.EqualValues(t, 100_000, *tx.Message.Config.LoadedAccountsDataSizeLimit) + assert.EqualValues(t, 64*1024, *tx.Message.Config.HeapSize) + assert.Equal(t, map[uint8][4]byte{5: {0xa, 0xb, 0xc, 0xd}, 31: {0x1, 0x2, 0x3, 0x4}}, tx.Message.Config.unknown) + + // Instructions and signatures are unaffected by the extra config values + var expected Transaction + require.NoError(t, expected.Unmarshal(valid)) + assert.Equal(t, expected.Message.Instructions, tx.Message.Instructions) + assert.Equal(t, expected.Signatures, tx.Signatures) } func public(priv ed25519.PrivateKey) ed25519.PublicKey { @@ -497,3 +821,91 @@ func generateKeys(t *testing.T, amount int) []ed25519.PrivateKey { return keys } + +func TestV1Transaction_MarshalOverflow(t *testing.T) { + newMessage := func() Message { + return Message{ + Version: MessageVersion1, + Header: Header{NumSignatures: 1}, + Accounts: []ed25519.PublicKey{make([]byte, ed25519.PublicKeySize)}, + Instructions: []CompiledInstruction{ + {ProgramIndex: 0}, + }, + } + } + + for name, tc := range map[string]struct { + mutate func(m *Message, n int) + max int + }{ + "instructions": { + mutate: func(m *Message, n int) { m.Instructions = make([]CompiledInstruction, n) }, + max: math.MaxUint8, + }, + "accounts": { + mutate: func(m *Message, n int) { + m.Accounts = make([]ed25519.PublicKey, n) + for i := range m.Accounts { + m.Accounts[i] = make([]byte, ed25519.PublicKeySize) + } + }, + max: math.MaxUint8, + }, + "instruction accounts": { + mutate: func(m *Message, n int) { m.Instructions[0].Accounts = make([]byte, n) }, + max: math.MaxUint8, + }, + "instruction data": { + mutate: func(m *Message, n int) { m.Instructions[0].Data = make([]byte, n) }, + max: math.MaxUint16, + }, + } { + t.Run(name, func(t *testing.T) { + m := newMessage() + tc.mutate(&m, tc.max) + _, err := m.Marshal() + assert.NoError(t, err) + + m = newMessage() + tc.mutate(&m, tc.max+1) + _, err = m.Marshal() + assert.Error(t, err) + }) + } +} + +func TestV1Transaction_MarshalSignatureCount(t *testing.T) { + message := Message{ + Version: MessageVersion1, + Header: Header{NumSignatures: 2}, + Accounts: []ed25519.PublicKey{make([]byte, ed25519.PublicKeySize), make([]byte, ed25519.PublicKeySize)}, + Instructions: []CompiledInstruction{ + {ProgramIndex: 1}, + }, + } + + // v1 has no signature count prefix, so it must match the header to decode + for _, n := range []int{0, 1, 3} { + _, err := Transaction{Signatures: make([]Signature, n), Message: message}.Marshal() + assert.Error(t, err) + } + + marshalled, err := Transaction{Signatures: make([]Signature, 2), Message: message}.Marshal() + require.NoError(t, err) + + var decoded Transaction + require.NoError(t, decoded.Unmarshal(marshalled)) + assert.Len(t, decoded.Signatures, 2) +} + +func mustMarshal(t *testing.T, tx Transaction) []byte { + b, err := tx.Marshal() + require.NoError(t, err) + return b +} + +func mustMarshalMessage(t *testing.T, m Message) []byte { + b, err := m.Marshal() + require.NoError(t, err) + return b +} From 06a0810e8d8e66b3ee204c84ec232612e0c18feb Mon Sep 17 00:00:00 2001 From: jeffyanta Date: Mon, 28 Sep 2026 10:02:49 -0400 Subject: [PATCH 2/2] Add test to ensure account ordering matches legacy --- solana/transaction_test.go | 79 ++++++++++++++++++++++++++++++++++++++ 1 file changed, 79 insertions(+) diff --git a/solana/transaction_test.go b/solana/transaction_test.go index ef934ab..2ece7bd 100644 --- a/solana/transaction_test.go +++ b/solana/transaction_test.go @@ -6,6 +6,7 @@ import ( "encoding/base64" "math" "math/rand" + "slices" "sort" "testing" @@ -641,6 +642,84 @@ func TestV1Transaction_Builder(t *testing.T) { assert.True(t, ed25519.Verify(public(payer), mustMarshalMessage(t, tx.Message), tx.Signatures[0][:])) } +func TestV1Transaction_MatchesLegacyAccountOrdering(t *testing.T) { + keys := generateKeys(t, 9) + payer, program, program2 := keys[0], keys[1], keys[2] + writableSigner, readonlySigner, writable, readonly, upgraded, added := keys[3], keys[4], keys[5], keys[6], keys[7], keys[8] + + config := TransactionConfig{ + ComputeUnitLimit: pointer.Uint32(10_000), + LoadedAccountsDataSizeLimit: pointer.Uint32(20_000), + } + + for name, instructions := range map[string][]Instruction{ + "single instruction": { + NewInstruction( + public(program), + []byte{1}, + NewReadonlyAccountMeta(public(readonly), false), + NewAccountMeta(public(writable), false), + NewReadonlyAccountMeta(public(readonlySigner), true), + NewAccountMeta(public(writableSigner), true), + ), + }, + "duplicate accounts with permission upgrades": { + NewInstruction( + public(program), + []byte{1}, + NewReadonlyAccountMeta(public(upgraded), false), + NewReadonlyAccountMeta(public(readonly), false), + NewAccountMeta(public(upgraded), true), + NewReadonlyAccountMeta(public(payer), false), + ), + }, + "multiple instructions": { + NewInstruction( + public(program2), + []byte{1, 2}, + NewReadonlyAccountMeta(public(readonlySigner), true), + NewAccountMeta(public(writable), false), + ), + NewInstruction( + public(program), + []byte{3}, + NewReadonlyAccountMeta(public(writable), false), + NewAccountMeta(public(readonlySigner), false), + NewAccountMeta(public(added), true), + NewReadonlyAccountMeta(public(readonly), false), + ), + NewInstruction( + public(program2), + nil, + NewReadonlyAccountMeta(public(program), false), + ), + }, + } { + t.Run(name, func(t *testing.T) { + legacy := NewLegacyTransaction(public(payer), instructions...) + v1, err := NewV1Transaction(public(payer), config, instructions...) + require.NoError(t, err) + + assert.Equal(t, legacy.Message.Header, v1.Message.Header) + assert.Equal(t, legacy.Message.Accounts, v1.Message.Accounts) + assert.Equal(t, legacy.Message.Instructions, v1.Message.Instructions) + assert.Len(t, v1.Signatures, len(legacy.Signatures)) + + // Account ordering doesn't depend on the order instructions list them + reversed := make([]Instruction, len(instructions)) + for i, ixn := range instructions { + metas := slices.Clone(ixn.Accounts) + slices.Reverse(metas) + reversed[len(instructions)-1-i] = NewInstruction(ixn.Program, ixn.Data, metas...) + } + reordered, err := NewV1Transaction(public(payer), config, reversed...) + require.NoError(t, err) + assert.Equal(t, v1.Message.Header, reordered.Message.Header) + assert.Equal(t, v1.Message.Accounts, reordered.Message.Accounts) + }) + } +} + func TestV1Transaction_BuilderConstraints(t *testing.T) { keys := generateKeys(t, 2) payer, program := keys[0], keys[1]