fix: Handle context cancellation (#24551)
Co-authored-by: Alex | Interchain Labs <alex@interchainlabs.io>
This commit is contained in:
co-authored by
Alex | Interchain Labs
parent
7f1f655900
commit
245a5faf02
@@ -43,6 +43,10 @@ Ref: https://keepachangelog.com/en/1.0.0/
|
||||
* (x/nft) [#24575](https://github.com/cosmos/cosmos-sdk/pull/24575) Deprecate the `x/nft` module in the Cosmos SDK repository. This module will not be maintained to the extent that our core modules will and will be kept in a [legacy repo](https://github.com/cosmos/cosmos-legacy).
|
||||
* (x/group) [#24571](https://github.com/cosmos/cosmos-sdk/pull/24571) Deprecate the `x/group` module in the Cosmos SDK repository. This module will not be maintained to the extent that our core modules will and will be kept in a [legacy repo](https://github.com/cosmos/cosmos-legacy).
|
||||
|
||||
### Bug Fixes
|
||||
|
||||
* (client, client/rpc, x/auth/tx) [#24551](https://github.com/cosmos/cosmos-sdk/pull/24551) Handle cancellation properly when supplying context to client methods.
|
||||
|
||||
## [v0.53.0](https://github.com/cosmos/cosmos-sdk/releases/tag/v0.53.0) - 2025-04-29
|
||||
|
||||
### Features
|
||||
|
||||
+2
-2
@@ -91,7 +91,7 @@ func (ctx Context) BroadcastTxSync(txBytes []byte) (*sdk.TxResponse, error) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
res, err := node.BroadcastTxSync(context.Background(), txBytes)
|
||||
res, err := node.BroadcastTxSync(ctx.GetCmdContextWithFallback(), txBytes)
|
||||
if errRes := CheckCometError(err, txBytes); errRes != nil {
|
||||
return errRes, nil
|
||||
}
|
||||
@@ -107,7 +107,7 @@ func (ctx Context) BroadcastTxAsync(txBytes []byte) (*sdk.TxResponse, error) {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
res, err := node.BroadcastTxAsync(context.Background(), txBytes)
|
||||
res, err := node.BroadcastTxAsync(ctx.GetCmdContextWithFallback(), txBytes)
|
||||
if errRes := CheckCometError(err, txBytes); errRes != nil {
|
||||
return errRes, nil
|
||||
}
|
||||
|
||||
+17
-16
@@ -7,28 +7,12 @@ import (
|
||||
|
||||
"github.com/cometbft/cometbft/crypto/tmhash"
|
||||
"github.com/cometbft/cometbft/mempool"
|
||||
"github.com/cometbft/cometbft/rpc/client/mock"
|
||||
coretypes "github.com/cometbft/cometbft/rpc/core/types"
|
||||
cmttypes "github.com/cometbft/cometbft/types"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/cosmos/cosmos-sdk/client/flags"
|
||||
sdkerrors "github.com/cosmos/cosmos-sdk/types/errors"
|
||||
)
|
||||
|
||||
type MockClient struct {
|
||||
mock.Client
|
||||
err error
|
||||
}
|
||||
|
||||
func (c MockClient) BroadcastTxAsync(_ context.Context, _ cmttypes.Tx) (*coretypes.ResultBroadcastTx, error) {
|
||||
return nil, c.err
|
||||
}
|
||||
|
||||
func (c MockClient) BroadcastTxSync(_ context.Context, _ cmttypes.Tx) (*coretypes.ResultBroadcastTx, error) {
|
||||
return nil, c.err
|
||||
}
|
||||
|
||||
func CreateContextWithErrorAndMode(err error, mode string) Context {
|
||||
return Context{
|
||||
Client: MockClient{err: err},
|
||||
@@ -63,3 +47,20 @@ func TestBroadcastError(t *testing.T) {
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestBroadcastCancellation(t *testing.T) {
|
||||
modes := []string{
|
||||
flags.BroadcastAsync,
|
||||
flags.BroadcastSync,
|
||||
}
|
||||
|
||||
txBytes := []byte{0xA, 0xB}
|
||||
cmdCtx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
for _, mode := range modes {
|
||||
ctx := CreateContextWithErrorAndMode(nil, mode).WithCmdContext(cmdCtx)
|
||||
_, err := ctx.BroadcastTx(txBytes)
|
||||
require.ErrorIs(t, err, context.Canceled)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -308,6 +308,15 @@ func (ctx Context) WithPreprocessTxHook(preprocessFn PreprocessTxFn) Context {
|
||||
return ctx
|
||||
}
|
||||
|
||||
// GetCmdContextWithFallback returns the CmdContext if it is not nil, otherwise it
|
||||
// returns context.Background()
|
||||
func (ctx Context) GetCmdContextWithFallback() context.Context {
|
||||
if ctx.CmdContext != nil {
|
||||
return ctx.CmdContext
|
||||
}
|
||||
return context.Background()
|
||||
}
|
||||
|
||||
// PrintString prints the raw string to ctx.Output if it's defined, otherwise to os.Stdout
|
||||
func (ctx Context) PrintString(str string) error {
|
||||
return ctx.PrintBytes([]byte(str))
|
||||
|
||||
+1
-2
@@ -1,7 +1,6 @@
|
||||
package client
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
@@ -95,7 +94,7 @@ func (ctx Context) queryABCI(req abci.RequestQuery) (abci.ResponseQuery, error)
|
||||
Prove: req.Prove,
|
||||
}
|
||||
|
||||
result, err := node.ABCIQueryWithOptions(context.Background(), req.Path, req.Data, opts)
|
||||
result, err := node.ABCIQueryWithOptions(ctx.GetCmdContextWithFallback(), req.Path, req.Data, opts)
|
||||
if err != nil {
|
||||
return abci.ResponseQuery{}, err
|
||||
}
|
||||
|
||||
@@ -0,0 +1,20 @@
|
||||
package client_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
abci "github.com/cometbft/cometbft/abci/types"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/cosmos/cosmos-sdk/client"
|
||||
)
|
||||
|
||||
func TestQueryABCICancellation(t *testing.T) {
|
||||
cmdCtx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
ctx := client.Context{}.WithClient(client.MockClient{}).WithCmdContext(cmdCtx)
|
||||
_, err := ctx.QueryABCI(abci.RequestQuery{})
|
||||
require.ErrorIs(t, err, context.Canceled)
|
||||
}
|
||||
+4
-5
@@ -1,7 +1,6 @@
|
||||
package rpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/hex"
|
||||
"fmt"
|
||||
"time"
|
||||
@@ -20,7 +19,7 @@ func GetChainHeight(clientCtx client.Context) (int64, error) {
|
||||
return -1, err
|
||||
}
|
||||
|
||||
status, err := node.Status(context.Background())
|
||||
status, err := node.Status(clientCtx.GetCmdContextWithFallback())
|
||||
if err != nil {
|
||||
return -1, err
|
||||
}
|
||||
@@ -53,7 +52,7 @@ func QueryBlocks(clientCtx client.Context, page, limit int, query, orderBy strin
|
||||
return nil, err
|
||||
}
|
||||
|
||||
resBlocks, err := node.BlockSearch(context.Background(), query, &page, &limit, orderBy)
|
||||
resBlocks, err := node.BlockSearch(clientCtx.GetCmdContextWithFallback(), query, &page, &limit, orderBy)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -79,7 +78,7 @@ func GetBlockByHeight(clientCtx client.Context, height *int64) (*cmt.Block, erro
|
||||
// header -> BlockchainInfo
|
||||
// header, tx -> Block
|
||||
// results -> BlockResults
|
||||
resBlock, err := node.Block(context.Background(), height)
|
||||
resBlock, err := node.Block(clientCtx.GetCmdContextWithFallback(), height)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -104,7 +103,7 @@ func GetBlockByHash(clientCtx client.Context, hashHexString string) (*cmt.Block,
|
||||
return nil, err
|
||||
}
|
||||
|
||||
resBlock, err := node.BlockByHash(context.Background(), hash)
|
||||
resBlock, err := node.BlockByHash(clientCtx.GetCmdContextWithFallback(), hash)
|
||||
|
||||
if err != nil {
|
||||
return nil, err
|
||||
|
||||
@@ -0,0 +1,59 @@
|
||||
package rpc_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/cosmos/cosmos-sdk/client"
|
||||
"github.com/cosmos/cosmos-sdk/client/rpc"
|
||||
)
|
||||
|
||||
func TestContextCancellation(t *testing.T) {
|
||||
testCases := []struct {
|
||||
name string
|
||||
query func(ctx client.Context) error
|
||||
}{
|
||||
{
|
||||
name: "get chain height cancellation",
|
||||
query: func(ctx client.Context) error {
|
||||
_, err := rpc.GetChainHeight(ctx)
|
||||
return err
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "query blocks cancellation",
|
||||
query: func(ctx client.Context) error {
|
||||
_, err := rpc.QueryBlocks(ctx, 1, 100, "", "")
|
||||
return err
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "get block by height cancellation",
|
||||
query: func(ctx client.Context) error {
|
||||
height := int64(1)
|
||||
_, err := rpc.GetBlockByHeight(ctx, &height)
|
||||
return err
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "get block by hash cancellation",
|
||||
query: func(ctx client.Context) error {
|
||||
_, err := rpc.GetBlockByHash(ctx, "")
|
||||
return err
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
cmdCtx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
ctx := client.Context{}.WithClient(client.MockClient{}).WithCmdContext(cmdCtx)
|
||||
err := tc.query(ctx)
|
||||
require.ErrorIs(t, err, context.Canceled)
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -1,8 +1,15 @@
|
||||
package client
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
"github.com/cometbft/cometbft/libs/bytes"
|
||||
rpcclient "github.com/cometbft/cometbft/rpc/client"
|
||||
"github.com/cometbft/cometbft/rpc/client/mock"
|
||||
coretypes "github.com/cometbft/cometbft/rpc/core/types"
|
||||
cmttypes "github.com/cometbft/cometbft/types"
|
||||
|
||||
cryptotypes "github.com/cosmos/cosmos-sdk/crypto/types"
|
||||
sdk "github.com/cosmos/cosmos-sdk/types"
|
||||
)
|
||||
@@ -81,3 +88,69 @@ func (t TestAccountRetriever) GetAccountNumberSequence(_ Context, addr sdk.AccAd
|
||||
}
|
||||
return acc.Num, acc.Seq, nil
|
||||
}
|
||||
|
||||
type MockClient struct {
|
||||
mock.Client
|
||||
err error
|
||||
}
|
||||
|
||||
func (c MockClient) ABCIQueryWithOptions(
|
||||
ctx context.Context,
|
||||
_ string,
|
||||
_ bytes.HexBytes,
|
||||
_ rpcclient.ABCIQueryOptions,
|
||||
) (*coretypes.ResultABCIQuery, error) {
|
||||
return handleError[*coretypes.ResultABCIQuery](ctx, c.err)
|
||||
}
|
||||
|
||||
func (c MockClient) BlockSearch(
|
||||
ctx context.Context,
|
||||
_ string,
|
||||
_, _ *int,
|
||||
_ string,
|
||||
) (*coretypes.ResultBlockSearch, error) {
|
||||
return handleError[*coretypes.ResultBlockSearch](ctx, c.err)
|
||||
}
|
||||
|
||||
func (c MockClient) BroadcastTxAsync(ctx context.Context, _ cmttypes.Tx) (*coretypes.ResultBroadcastTx, error) {
|
||||
return handleError[*coretypes.ResultBroadcastTx](ctx, c.err)
|
||||
}
|
||||
|
||||
func (c MockClient) BroadcastTxSync(ctx context.Context, _ cmttypes.Tx) (*coretypes.ResultBroadcastTx, error) {
|
||||
return handleError[*coretypes.ResultBroadcastTx](ctx, c.err)
|
||||
}
|
||||
|
||||
func (c MockClient) Block(ctx context.Context, _ *int64) (*coretypes.ResultBlock, error) {
|
||||
return handleError[*coretypes.ResultBlock](ctx, c.err)
|
||||
}
|
||||
|
||||
func (c MockClient) BlockByHash(ctx context.Context, _ []byte) (*coretypes.ResultBlock, error) {
|
||||
return handleError[*coretypes.ResultBlock](ctx, c.err)
|
||||
}
|
||||
|
||||
func (c MockClient) Status(ctx context.Context) (*coretypes.ResultStatus, error) {
|
||||
return handleError[*coretypes.ResultStatus](ctx, c.err)
|
||||
}
|
||||
|
||||
func (c MockClient) Tx(ctx context.Context, _ []byte, _ bool) (*coretypes.ResultTx, error) {
|
||||
return handleError[*coretypes.ResultTx](ctx, c.err)
|
||||
}
|
||||
|
||||
func (c MockClient) TxSearch(
|
||||
ctx context.Context,
|
||||
_ string,
|
||||
_ bool,
|
||||
_, _ *int,
|
||||
_ string,
|
||||
) (*coretypes.ResultTxSearch, error) {
|
||||
return handleError[*coretypes.ResultTxSearch](ctx, c.err)
|
||||
}
|
||||
|
||||
func handleError[T any](ctx context.Context, err error) (T, error) {
|
||||
var ret T
|
||||
if ctx != nil && ctx.Err() != nil {
|
||||
return ret, ctx.Err()
|
||||
} else {
|
||||
return ret, err
|
||||
}
|
||||
}
|
||||
|
||||
+1
-1
@@ -126,7 +126,7 @@ func BroadcastTx(clientCtx client.Context, txf Factory, msgs ...sdk.Msg) error {
|
||||
}
|
||||
}
|
||||
|
||||
if err = Sign(clientCtx.CmdContext, txf, clientCtx.FromName, tx, true); err != nil {
|
||||
if err = Sign(clientCtx.GetCmdContextWithFallback(), txf, clientCtx.FromName, tx, true); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
|
||||
+2
-2
@@ -62,7 +62,7 @@ func SignTx(txFactory tx.Factory, clientCtx client.Context, name string, txBuild
|
||||
}
|
||||
}
|
||||
|
||||
return tx.Sign(clientCtx.CmdContext, txFactory, name, txBuilder, overwriteSig)
|
||||
return tx.Sign(clientCtx.GetCmdContextWithFallback(), txFactory, name, txBuilder, overwriteSig)
|
||||
}
|
||||
|
||||
// SignTxWithSignerAddress attaches a signature to a transaction.
|
||||
@@ -85,7 +85,7 @@ func SignTxWithSignerAddress(txFactory tx.Factory, clientCtx client.Context, add
|
||||
}
|
||||
}
|
||||
|
||||
return tx.Sign(clientCtx.CmdContext, txFactory, name, txBuilder, overwrite)
|
||||
return tx.Sign(clientCtx.GetCmdContextWithFallback(), txFactory, name, txBuilder, overwrite)
|
||||
}
|
||||
|
||||
// Read and decode a StdTx from the given filename. Can pass "-" to read from stdin.
|
||||
|
||||
+3
-4
@@ -1,7 +1,6 @@
|
||||
package tx
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
@@ -42,7 +41,7 @@ func QueryTxsByEvents(clientCtx client.Context, page, limit int, query, orderBy
|
||||
return nil, err
|
||||
}
|
||||
|
||||
resTxs, err := node.TxSearch(context.Background(), query, false, &page, &limit, orderBy)
|
||||
resTxs, err := node.TxSearch(clientCtx.GetCmdContextWithFallback(), query, false, &page, &limit, orderBy)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to search for txs: %w", err)
|
||||
}
|
||||
@@ -75,7 +74,7 @@ func QueryTx(clientCtx client.Context, hashHexStr string) (*sdk.TxResponse, erro
|
||||
|
||||
// TODO: this may not always need to be proven
|
||||
// https://github.com/cosmos/cosmos-sdk/issues/6807
|
||||
resTx, err := node.Tx(context.Background(), hash, true)
|
||||
resTx, err := node.Tx(clientCtx.GetCmdContextWithFallback(), hash, true)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -117,7 +116,7 @@ func getBlocksForTxResults(clientCtx client.Context, resTxs []*coretypes.ResultT
|
||||
|
||||
for _, resTx := range resTxs {
|
||||
if _, ok := resBlocks[resTx.Height]; !ok {
|
||||
resBlock, err := node.Block(context.Background(), &resTx.Height)
|
||||
resBlock, err := node.Block(clientCtx.GetCmdContextWithFallback(), &resTx.Height)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -0,0 +1,44 @@
|
||||
package tx_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/cosmos/cosmos-sdk/client"
|
||||
"github.com/cosmos/cosmos-sdk/x/auth/tx"
|
||||
)
|
||||
|
||||
func TestContextCancellation(t *testing.T) {
|
||||
testCases := []struct {
|
||||
name string
|
||||
query func(ctx client.Context) error
|
||||
}{
|
||||
{
|
||||
name: "query tx cancellation",
|
||||
query: func(ctx client.Context) error {
|
||||
_, err := tx.QueryTx(ctx, "")
|
||||
return err
|
||||
},
|
||||
},
|
||||
{
|
||||
name: "query txs by events cancellation",
|
||||
query: func(ctx client.Context) error {
|
||||
_, err := tx.QueryTxsByEvents(ctx, 1, 100, "query", "")
|
||||
return err
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
for _, tc := range testCases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
cmdCtx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
ctx := client.Context{}.WithClient(client.MockClient{}).WithCmdContext(cmdCtx)
|
||||
err := tc.query(ctx)
|
||||
require.ErrorIs(t, err, context.Canceled)
|
||||
})
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user