From 740a933a5c94c276c61df328ddd0dbfce88d9653 Mon Sep 17 00:00:00 2001 From: YimingZang Date: Mon, 17 Aug 2026 11:04:57 -0700 Subject: [PATCH 1/6] Add context cancellation to SS DB layer --- evmrpc/simulate.go | 8 + evmrpc/trace_cancellation_test.go | 30 +++ evmrpc/trace_transaction_timeout_test.go | 239 ++++++++++++++++++ evmrpc/tracers.go | 16 +- sei-cosmos/baseapp/baseapp.go | 1 + sei-cosmos/baseapp/recovery.go | 19 ++ sei-cosmos/store/cachekv/store.go | 25 +- sei-cosmos/store/ctxkv/store.go | 56 ++++ sei-cosmos/store/ctxkv/store_test.go | 87 +++++++ sei-cosmos/store/types/store.go | 24 ++ sei-cosmos/storev2/state/store.go | 25 +- sei-cosmos/types/context.go | 9 +- sei-db/db_engine/pebbledb/mvcc/db.go | 28 +- .../db_engine/pebbledb/mvcc/db_ascending.go | 8 +- sei-db/db_engine/pebbledb/mvcc/iterator.go | 51 +++- .../pebbledb/mvcc/iterator_ascending.go | 36 ++- .../pebbledb/mvcc/iterator_context_test.go | 64 +++++ sei-db/db_engine/rocksdb/mvcc/db.go | 9 + sei-db/db_engine/types/types.go | 23 ++ sei-db/state_db/ss/composite/store.go | 16 ++ sei-db/state_db/ss/cosmos/store.go | 11 + sei-db/state_db/ss/evm/store.go | 24 ++ 22 files changed, 770 insertions(+), 39 deletions(-) create mode 100644 evmrpc/trace_cancellation_test.go create mode 100644 evmrpc/trace_transaction_timeout_test.go create mode 100644 sei-cosmos/store/ctxkv/store.go create mode 100644 sei-cosmos/store/ctxkv/store_test.go create mode 100644 sei-db/db_engine/pebbledb/mvcc/iterator_context_test.go diff --git a/evmrpc/simulate.go b/evmrpc/simulate.go index 2be879e60f..27941a87e7 100644 --- a/evmrpc/simulate.go +++ b/evmrpc/simulate.go @@ -669,6 +669,9 @@ func (b *Backend) replayTransactionTillIndex(ctx context.Context, block *ethtype } _ = b.app.DeliverTx(sdkCtx, abci.RequestDeliverTxV2{Tx: tx}, sdkTx, sha256.Sum256(tx)) } + if err := ctx.Err(); err != nil { + return nil, nil, emptyRelease, err + } success = true return state.NewDBImpl(sdkCtx.WithIsEVM(true), b.keeper, true), tmBlock.Block.Txs, release, nil } @@ -702,6 +705,11 @@ func (b *Backend) initializeBlock(ctx context.Context, block *ethtypes.Block, ct reqBeginBlock.Simulate = true baseCtx, baseRelease := ctxProvider(prevBlockHeight) sdkCtx := baseCtx.WithBlockHeight(blockNumber).WithBlockTime(tmBlock.Block.Time) + if ctx != nil { + // The RPC/trace deadline must be on the SDK context so KVStore + // iteration can pass it into the SS MVCC skip loops. + sdkCtx = sdkCtx.WithContext(ctx) + } legacyabci.BeginBlock(sdkCtx, blockNumber, reqBeginBlock.LastCommitInfo.Votes, tmBlock.Block.Evidence.ToABCI(), b.beginBlockKeepers) nextCtx, nextRelease := ctxProvider(sdkCtx.BlockHeight()) sdkCtx = sdkCtx.WithNextMs( diff --git a/evmrpc/trace_cancellation_test.go b/evmrpc/trace_cancellation_test.go new file mode 100644 index 0000000000..f6e550cf75 --- /dev/null +++ b/evmrpc/trace_cancellation_test.go @@ -0,0 +1,30 @@ +package evmrpc + +import ( + "context" + "errors" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestResultUnlessExpiredReportsTheDeadline(t *testing.T) { + live, cancel := context.WithCancel(context.Background()) + defer cancel() + + traced := map[string]string{"gas": "0x1"} + result, err := resultUnlessExpired(live, traced, nil) + require.NoError(t, err) + require.Equal(t, traced, result) + + underlying := errors.New("tracer failed") + result, err = resultUnlessExpired(live, nil, underlying) + require.ErrorIs(t, err, underlying) + require.Nil(t, result) + + expired, cancelExpired := context.WithCancel(context.Background()) + cancelExpired() + result, err = resultUnlessExpired(expired, map[string]string{"error": "store access cancelled"}, nil) + require.ErrorIs(t, err, context.Canceled) + require.Nil(t, result) +} diff --git a/evmrpc/trace_transaction_timeout_test.go b/evmrpc/trace_transaction_timeout_test.go new file mode 100644 index 0000000000..af892fd4b8 --- /dev/null +++ b/evmrpc/trace_transaction_timeout_test.go @@ -0,0 +1,239 @@ +package evmrpc + +import ( + "context" + "errors" + "io" + "math/big" + "sync" + "testing" + "time" + + "github.com/ethereum/go-ethereum/common" + "github.com/ethereum/go-ethereum/consensus" + ethtypes "github.com/ethereum/go-ethereum/core/types" + "github.com/ethereum/go-ethereum/core/vm" + "github.com/ethereum/go-ethereum/eth/tracers" + "github.com/ethereum/go-ethereum/eth/tracers/tracersutils" + "github.com/ethereum/go-ethereum/ethdb" + "github.com/ethereum/go-ethereum/export" + "github.com/ethereum/go-ethereum/params" + "github.com/ethereum/go-ethereum/rpc" + "github.com/ethereum/go-ethereum/trie" + storetypes "github.com/sei-protocol/sei-chain/sei-cosmos/store/types" + sdk "github.com/sei-protocol/sei-chain/sei-cosmos/types" + tmproto "github.com/sei-protocol/sei-chain/sei-tendermint/proto/tendermint/types" + "github.com/sei-protocol/sei-chain/x/evm/keeper" + "github.com/stretchr/testify/require" +) + +// TestTraceTransactionTimeoutReleasesSemaphore sends debug_traceTransaction +// through a StateAtTransaction that blocks in an SS-shaped skip loop. The +// first call returns the deadline, a concurrent call is rejected as busy, and +// the semaphore slot is free afterwards. +func TestTraceTransactionTimeoutReleasesSemaphore(t *testing.T) { + t.Parallel() + + store := newSlowSkipStore() + backend := newSlowIterTraceBackend(store) + api := &DebugAPI{ + tracersAPI: tracers.NewAPI(backend), + keeper: &keeper.Keeper{}, + ctxProvider: func(int64) sdk.Context { return sdk.Context{} }, + traceCallSemaphore: make(chan struct{}, 1), + traceTimeout: 50 * time.Millisecond, + } + + firstErr := make(chan error, 1) + go func() { + _, err := api.TraceTransaction(context.Background(), backend.tx.Hash(), nil) + firstErr <- err + }() + + select { + case <-store.entered: + case <-time.After(2 * time.Second): + t.Fatal("StateAtTransaction did not enter the skip loop") + } + + _, busyErr := api.TraceTransaction(context.Background(), backend.tx.Hash(), nil) + require.ErrorIs(t, busyErr, errTraceConcurrencyLimit) + + select { + case err := <-firstErr: + require.ErrorIs(t, err, context.DeadlineExceeded) + case <-time.After(2 * time.Second): + t.Fatal("timed-out trace did not return") + } + + select { + case api.traceCallSemaphore <- struct{}{}: + <-api.traceCallSemaphore + default: + t.Fatal("expected the timed-out trace to release the semaphore") + } + + _, err := api.TraceTransaction(context.Background(), backend.tx.Hash(), nil) + require.NotErrorIs(t, err, errTraceConcurrencyLimit) + require.ErrorIs(t, err, context.DeadlineExceeded) +} + +type slowSkipStore struct { + entered chan struct{} + once sync.Once +} + +func newSlowSkipStore() *slowSkipStore { + return &slowSkipStore{entered: make(chan struct{})} +} + +var ( + _ sdk.KVStore = (*slowSkipStore)(nil) + _ storetypes.ContextIterator = (*slowSkipStore)(nil) +) + +func (s *slowSkipStore) skipUntilCancelled(ctx context.Context) { + s.once.Do(func() { close(s.entered) }) + for { + if err := ctx.Err(); err != nil { + panic(err) + } + time.Sleep(time.Millisecond) + } +} + +func (s *slowSkipStore) IteratorWithContext(ctx context.Context, _, _ []byte) storetypes.Iterator { + s.skipUntilCancelled(ctx) + return &emptyTraceIter{} +} + +func (s *slowSkipStore) ReverseIteratorWithContext(ctx context.Context, _, _ []byte) storetypes.Iterator { + s.skipUntilCancelled(ctx) + return &emptyTraceIter{} +} + +func (s *slowSkipStore) Iterator(_, _ []byte) storetypes.Iterator { + panic("Iterator called without context; expected IteratorWithContext") +} + +func (s *slowSkipStore) ReverseIterator(_, _ []byte) storetypes.Iterator { + panic("ReverseIterator called without context; expected ReverseIteratorWithContext") +} + +func (s *slowSkipStore) GetStoreType() sdk.StoreType { return sdk.StoreTypeDB } +func (s *slowSkipStore) CacheWrap(sdk.StoreKey) sdk.CacheWrap { + panic("not implemented") +} +func (s *slowSkipStore) CacheWrapWithTrace(sdk.StoreKey, io.Writer, sdk.TraceContext) sdk.CacheWrap { + panic("not implemented") +} +func (s *slowSkipStore) Get([]byte) []byte { return nil } +func (s *slowSkipStore) Has([]byte) bool { return false } +func (s *slowSkipStore) Set(_, _ []byte) {} +func (s *slowSkipStore) Delete([]byte) {} +func (s *slowSkipStore) GetWorkingHash() ([]byte, error) { return nil, nil } +func (s *slowSkipStore) VersionExists(int64) bool { return true } +func (s *slowSkipStore) DeleteAll(_, _ []byte) error { return nil } +func (s *slowSkipStore) GetAllKeyStrsInRange(_, _ []byte) []string { + return nil +} + +type emptyTraceIter struct{} + +func (emptyTraceIter) Domain() ([]byte, []byte) { return nil, nil } +func (emptyTraceIter) Valid() bool { return false } +func (emptyTraceIter) Next() {} +func (emptyTraceIter) Key() []byte { return nil } +func (emptyTraceIter) Value() []byte { return nil } +func (emptyTraceIter) Error() error { return nil } +func (emptyTraceIter) Close() error { return nil } + +type oneStoreMS struct { + fakeMultiStore + kv sdk.KVStore +} + +func (m *oneStoreMS) GetKVStore(sdk.StoreKey) sdk.KVStore { return m.kv } + +type slowIterTraceBackend struct { + block *ethtypes.Block + tx *ethtypes.Transaction + store *slowSkipStore +} + +var _ tracers.Backend = (*slowIterTraceBackend)(nil) + +func newSlowIterTraceBackend(store *slowSkipStore) *slowIterTraceBackend { + tx := ethtypes.NewTx(ðtypes.LegacyTx{ + Nonce: 0, + GasPrice: big.NewInt(1), + Gas: 21000, + To: &common.Address{}, + }) + header := ðtypes.Header{ + Number: big.NewInt(8), + Time: 1, + Difficulty: big.NewInt(0), + } + block := ethtypes.NewBlock(header, ðtypes.Body{Transactions: ethtypes.Transactions{tx}}, nil, trie.NewStackTrie(nil)) + return &slowIterTraceBackend{block: block, tx: tx, store: store} +} + +func (b *slowIterTraceBackend) HeaderByHash(context.Context, common.Hash) (*ethtypes.Header, error) { + return b.block.Header(), nil +} + +func (b *slowIterTraceBackend) HeaderByNumber(context.Context, rpc.BlockNumber) (*ethtypes.Header, error) { + return b.block.Header(), nil +} + +func (b *slowIterTraceBackend) BlockByHash(context.Context, common.Hash) (*ethtypes.Block, []tracersutils.TraceBlockMetadata, error) { + return b.block, nil, nil +} + +func (b *slowIterTraceBackend) BlockByNumber(context.Context, rpc.BlockNumber) (*ethtypes.Block, []tracersutils.TraceBlockMetadata, error) { + return b.block, nil, nil +} + +func (b *slowIterTraceBackend) GetTransaction(context.Context, common.Hash) (bool, *ethtypes.Transaction, common.Hash, uint64, uint64, error) { + return true, b.tx, b.block.Hash(), b.block.NumberU64(), 0, nil +} + +func (b *slowIterTraceBackend) RPCGasCap() uint64 { return 0 } + +func (b *slowIterTraceBackend) ChainConfig() *params.ChainConfig { return params.TestChainConfig } + +func (b *slowIterTraceBackend) ChainConfigAtHeight(int64) *params.ChainConfig { + return params.TestChainConfig +} + +func (b *slowIterTraceBackend) Engine() consensus.Engine { return nil } + +func (b *slowIterTraceBackend) ChainDb() ethdb.Database { return nil } + +func (b *slowIterTraceBackend) StateAtBlock(context.Context, *ethtypes.Block, uint64, vm.StateDB, bool, bool) (vm.StateDB, tracers.StateReleaseFunc, error) { + return nil, func() {}, errors.New("unused") +} + +func (b *slowIterTraceBackend) GetCustomPrecompiles(int64) map[common.Address]vm.PrecompiledContract { + return nil +} + +func (b *slowIterTraceBackend) PrepareTx(vm.StateDB, *ethtypes.Transaction) error { return nil } + +func (b *slowIterTraceBackend) GetBlockContext(context.Context, *ethtypes.Block, vm.StateDB, export.ChainContextBackend) (vm.BlockContext, error) { + return vm.BlockContext{}, errors.New("unused") +} + +func (b *slowIterTraceBackend) StateAtTransaction(ctx context.Context, _ *ethtypes.Block, _ int, _ uint64) (*ethtypes.Transaction, vm.BlockContext, vm.StateDB, tracers.StateReleaseFunc, error) { + key := sdk.NewKVStoreKey("evm") + sdkCtx := sdk.NewContext(&oneStoreMS{kv: b.store}, tmproto.Header{}, false).WithContext(ctx) + iter := sdkCtx.KVStore(key).Iterator(nil, nil) + defer func() { _ = iter.Close() }() + for ; iter.Valid(); iter.Next() { + } + if err := ctx.Err(); err != nil { + return nil, vm.BlockContext{}, nil, func() {}, err + } + return nil, vm.BlockContext{}, nil, func() {}, errors.New("slow iteration returned without deadline") +} diff --git a/evmrpc/tracers.go b/evmrpc/tracers.go index 9afaf9ccfc..c5c6bf5b5b 100644 --- a/evmrpc/tracers.go +++ b/evmrpc/tracers.go @@ -101,6 +101,16 @@ func (api *DebugAPI) prepareTraceContext(ctx context.Context) (context.Context, }, nil } +// resultUnlessExpired reports the deadline instead of a synthesized error +// trace. geth recovers a panic from a cancelled SS skip as a fake trace +// result; that must not be returned or written to the trace cache. +func resultUnlessExpired(ctx context.Context, result interface{}, err error) (interface{}, error) { + if expired := ctx.Err(); expired != nil { + return nil, expired + } + return result, err +} + func (api *DebugAPI) guardHistoricalDebugTraceByTxHash(ctx context.Context, endpoint string, hash common.Hash) error { if api.keeper == nil { return nil @@ -349,7 +359,8 @@ func (api *DebugAPI) TraceTransaction(ctx context.Context, hash common.Hash, con config = &tracers.TraceConfig{} } api.clampDefaultStructLogLimit(config) - return api.tracersAPI.TraceTransaction(ctx, hash, config) + traced, err := api.tracersAPI.TraceTransaction(ctx, hash, config) + return resultUnlessExpired(ctx, traced, err) } func (api *DebugAPI) tryTraceCache(hash common.Hash, config *tracers.TraceConfig) (interface{}, bool) { @@ -567,6 +578,7 @@ func (api *DebugAPI) TraceBlockByNumber(ctx context.Context, number rpc.BlockNum } else { result, returnErr = api.tracersAPI.TraceBlockByNumber(ctx, number, config) } + result, returnErr = resultUnlessExpired(ctx, result, returnErr) return } @@ -603,6 +615,7 @@ func (api *DebugAPI) TraceBlockByHash(ctx context.Context, hash common.Hash, con } else { result, returnErr = api.tracersAPI.TraceBlockByHash(ctx, hash, config) } + result, returnErr = resultUnlessExpired(ctx, result, returnErr) return } @@ -634,6 +647,7 @@ func (api *DebugAPI) TraceCall(ctx context.Context, args export.TransactionArgs, } api.clampDefaultStructLogLimit(&config.TraceConfig) result, returnErr = api.tracersAPI.TraceCall(ctx, args, blockNrOrHash, config) + result, returnErr = resultUnlessExpired(ctx, result, returnErr) return } diff --git a/sei-cosmos/baseapp/baseapp.go b/sei-cosmos/baseapp/baseapp.go index 43d8e41cf1..b447c41ce5 100644 --- a/sei-cosmos/baseapp/baseapp.go +++ b/sei-cosmos/baseapp/baseapp.go @@ -904,6 +904,7 @@ func (app *BaseApp) runTx(ctx sdk.Context, mode runTxMode, tx sdk.Tx, checksum [ if r := recover(); r != nil { recoveryMW := newOutOfGasRecoveryMiddleware(gasWanted, ctx, app.runTxRecoveryMiddleware) recoveryMW = newOCCAbortRecoveryMiddleware(recoveryMW) // TODO: do we have to wrap with occ enabled check? + recoveryMW = newContextCancelledRecoveryMiddleware(recoveryMW) err, runTxRes.result = processRecovery(r, recoveryMW), nil } if ctx.GasMeter() == blockGasMeter { diff --git a/sei-cosmos/baseapp/recovery.go b/sei-cosmos/baseapp/recovery.go index 0c49a75b9e..381e489160 100644 --- a/sei-cosmos/baseapp/recovery.go +++ b/sei-cosmos/baseapp/recovery.go @@ -1,6 +1,8 @@ package baseapp import ( + "context" + "errors" "fmt" "runtime/debug" @@ -83,6 +85,23 @@ func newOCCAbortRecoveryMiddleware(next recoveryMiddleware) recoveryMiddleware { return newRecoveryMiddleware(handler, next) } +// newContextCancelledRecoveryMiddleware recovers a store access that panicked +// because the caller's context was cancelled or exceeded its deadline. +func newContextCancelledRecoveryMiddleware(next recoveryMiddleware) recoveryMiddleware { + handler := func(recoveryObj interface{}) error { + err, ok := recoveryObj.(error) + if !ok { + return nil + } + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + return err + } + return nil + } + + return newRecoveryMiddleware(handler, next) +} + // newDefaultRecoveryMiddleware creates a default (last in chain) recovery middleware for app.runTx method. func newDefaultRecoveryMiddleware() recoveryMiddleware { handler := func(recoveryObj interface{}) error { diff --git a/sei-cosmos/store/cachekv/store.go b/sei-cosmos/store/cachekv/store.go index da09b01478..7c25195d30 100644 --- a/sei-cosmos/store/cachekv/store.go +++ b/sei-cosmos/store/cachekv/store.go @@ -2,6 +2,7 @@ package cachekv import ( "bytes" + "context" "io" "sort" "sync" @@ -14,6 +15,8 @@ import ( dbm "github.com/tendermint/tm-db" ) +var _ types.ContextIterator = (*Store)(nil) + // Store wraps an in-memory cache around an underlying types.KVStore. type Store struct { mtx sync.RWMutex @@ -208,12 +211,22 @@ func (store *Store) CacheWrapWithTrace(storeKey types.StoreKey, w io.Writer, tc // Iterator implements types.KVStore. func (store *Store) Iterator(start, end []byte) types.Iterator { - return store.iterator(start, end, true) + return store.iterator(context.Background(), start, end, true) +} + +// IteratorWithContext implements types.ContextIterator. +func (store *Store) IteratorWithContext(ctx context.Context, start, end []byte) types.Iterator { + return store.iterator(ctx, start, end, true) } // ReverseIterator implements types.KVStore. func (store *Store) ReverseIterator(start, end []byte) types.Iterator { - return store.iterator(start, end, false) + return store.iterator(context.Background(), start, end, false) +} + +// ReverseIteratorWithContext implements types.ContextIterator. +func (store *Store) ReverseIteratorWithContext(ctx context.Context, start, end []byte) types.Iterator { + return store.iterator(ctx, start, end, false) } func (store *Store) getOrInitSortedCache() *dbm.MemDB { @@ -223,7 +236,7 @@ func (store *Store) getOrInitSortedCache() *dbm.MemDB { return store.sortedCache } -func (store *Store) iterator(start, end []byte, ascending bool) types.Iterator { +func (store *Store) iterator(ctx context.Context, start, end []byte, ascending bool) types.Iterator { store.mtx.Lock() defer store.mtx.Unlock() // TODO: (occ) Note that for iterators, we'll need to have special handling (discussed in RFC) to ensure proper validation @@ -235,11 +248,7 @@ func (store *Store) iterator(start, end []byte, ascending bool) types.Iterator { // nothing to iteration; skipping it avoids building an O(depth) chain of // cacheMergeIterators over a deep snapshot stack. parentStore := store.readThroughParent() - if ascending { - parent = parentStore.Iterator(start, end) - } else { - parent = parentStore.ReverseIterator(start, end) - } + parent = types.IteratorOn(parentStore, ctx, start, end, ascending) defer func() { if err := recover(); err != nil { // close out parent iterator, then reraise panic diff --git a/sei-cosmos/store/ctxkv/store.go b/sei-cosmos/store/ctxkv/store.go new file mode 100644 index 0000000000..7f913a8fc2 --- /dev/null +++ b/sei-cosmos/store/ctxkv/store.go @@ -0,0 +1,56 @@ +package ctxkv + +import ( + "context" + "io" + + "github.com/sei-protocol/sei-chain/sei-cosmos/store/types" +) + +var _ types.KVStore = (*Store)(nil) + +// Store forwards KVStore calls to parent, passing ctx into iteration so a +// historical SS MVCC skip loop can observe the caller's deadline. +type Store struct { + parent types.KVStore + ctx context.Context +} + +// Wrap returns parent unchanged when ctx cannot be cancelled. Otherwise it +// wraps parent so Iterator/ReverseIterator carry ctx to a ContextIterator. +func Wrap(parent types.KVStore, ctx context.Context) types.KVStore { + if parent == nil || ctx == nil || ctx.Done() == nil { + return parent + } + return &Store{parent: parent, ctx: ctx} +} + +func (s *Store) GetStoreType() types.StoreType { return s.parent.GetStoreType() } +func (s *Store) GetWorkingHash() ([]byte, error) { + return s.parent.GetWorkingHash() +} +func (s *Store) Get(key []byte) []byte { return s.parent.Get(key) } +func (s *Store) Has(key []byte) bool { return s.parent.Has(key) } +func (s *Store) Set(key, value []byte) { s.parent.Set(key, value) } +func (s *Store) Delete(key []byte) { s.parent.Delete(key) } +func (s *Store) CacheWrap(storeKey types.StoreKey) types.CacheWrap { + return s.parent.CacheWrap(storeKey) +} +func (s *Store) CacheWrapWithTrace(storeKey types.StoreKey, w io.Writer, tc types.TraceContext) types.CacheWrap { + return s.parent.CacheWrapWithTrace(storeKey, w, tc) +} +func (s *Store) VersionExists(version int64) bool { return s.parent.VersionExists(version) } +func (s *Store) DeleteAll(start, end []byte) error { + return s.parent.DeleteAll(start, end) +} +func (s *Store) GetAllKeyStrsInRange(start, end []byte) []string { + return s.parent.GetAllKeyStrsInRange(start, end) +} + +func (s *Store) Iterator(start, end []byte) types.Iterator { + return types.IteratorOn(s.parent, s.ctx, start, end, true) +} + +func (s *Store) ReverseIterator(start, end []byte) types.Iterator { + return types.IteratorOn(s.parent, s.ctx, start, end, false) +} diff --git a/sei-cosmos/store/ctxkv/store_test.go b/sei-cosmos/store/ctxkv/store_test.go new file mode 100644 index 0000000000..db49ace97d --- /dev/null +++ b/sei-cosmos/store/ctxkv/store_test.go @@ -0,0 +1,87 @@ +package ctxkv_test + +import ( + "context" + "io" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/sei-protocol/sei-chain/sei-cosmos/store/ctxkv" + "github.com/sei-protocol/sei-chain/sei-cosmos/store/types" +) + +type recordingStore struct { + stubIterStore + gotCtx context.Context +} + +func (s *recordingStore) IteratorWithContext(ctx context.Context, start, end []byte) types.Iterator { + s.gotCtx = ctx + return s.Iterator(start, end) +} + +func (s *recordingStore) ReverseIteratorWithContext(ctx context.Context, start, end []byte) types.Iterator { + s.gotCtx = ctx + return s.ReverseIterator(start, end) +} + +func TestWrapIsNoopWhenContextCannotCancel(t *testing.T) { + parent := &recordingStore{} + require.Equal(t, parent, ctxkv.Wrap(parent, nil)) + require.Equal(t, parent, ctxkv.Wrap(parent, context.Background())) +} + +func TestWrapForwardsDeadlineToContextIterator(t *testing.T) { + parent := &recordingStore{} + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + wrapped := ctxkv.Wrap(parent, ctx) + require.NotEqual(t, parent, wrapped) + + _ = wrapped.Iterator(nil, nil) + require.Equal(t, ctx, parent.gotCtx) + + parent.gotCtx = nil + _ = wrapped.ReverseIterator(nil, nil) + require.Equal(t, ctx, parent.gotCtx) +} + +type stubIterStore struct{} + +func (stubIterStore) GetStoreType() types.StoreType { return types.StoreTypeDB } +func (stubIterStore) CacheWrap(types.StoreKey) types.CacheWrap { + panic("not implemented") +} +func (stubIterStore) CacheWrapWithTrace(types.StoreKey, io.Writer, types.TraceContext) types.CacheWrap { + panic("not implemented") +} +func (stubIterStore) Get([]byte) []byte { return nil } +func (stubIterStore) Has([]byte) bool { return false } +func (stubIterStore) Set([]byte, []byte) {} +func (stubIterStore) Delete([]byte) {} +func (stubIterStore) Iterator(start, end []byte) types.Iterator { + return &emptyIter{start: start, end: end} +} +func (stubIterStore) ReverseIterator(start, end []byte) types.Iterator { + return &emptyIter{start: start, end: end} +} +func (stubIterStore) GetWorkingHash() ([]byte, error) { return nil, nil } +func (stubIterStore) VersionExists(int64) bool { return true } +func (stubIterStore) DeleteAll([]byte, []byte) error { return nil } +func (stubIterStore) GetAllKeyStrsInRange([]byte, []byte) []string { + return nil +} + +type emptyIter struct { + start, end []byte +} + +func (e *emptyIter) Domain() ([]byte, []byte) { return e.start, e.end } +func (e *emptyIter) Valid() bool { return false } +func (e *emptyIter) Next() {} +func (e *emptyIter) Key() []byte { return nil } +func (e *emptyIter) Value() []byte { return nil } +func (e *emptyIter) Error() error { return nil } +func (e *emptyIter) Close() error { return nil } diff --git a/sei-cosmos/store/types/store.go b/sei-cosmos/store/types/store.go index 974be43bb3..e247214a68 100644 --- a/sei-cosmos/store/types/store.go +++ b/sei-cosmos/store/types/store.go @@ -251,6 +251,30 @@ type KVStore interface { GetAllKeyStrsInRange(start, end []byte) []string } +// ContextIterator is implemented by KVStores whose Iterator/Next may block in +// the historical SS MVCC skip loops. A deadline on ctx is observed inside those +// loops; stores that do not implement this are iterated via Iterator as today. +type ContextIterator interface { + IteratorWithContext(ctx context.Context, start, end []byte) Iterator + ReverseIteratorWithContext(ctx context.Context, start, end []byte) Iterator +} + +// IteratorOn iterates store, using ctx when store implements ContextIterator. +func IteratorOn(store KVStore, ctx context.Context, start, end []byte, ascending bool) Iterator { + if ctx != nil { + if ci, ok := store.(ContextIterator); ok { + if ascending { + return ci.IteratorWithContext(ctx, start, end) + } + return ci.ReverseIteratorWithContext(ctx, start, end) + } + } + if ascending { + return store.Iterator(start, end) + } + return store.ReverseIterator(start, end) +} + // Iterator is an alias db's Iterator for convenience. type Iterator = dbm.Iterator diff --git a/sei-cosmos/storev2/state/store.go b/sei-cosmos/storev2/state/store.go index 4750e4d7bf..0a73d94b34 100644 --- a/sei-cosmos/storev2/state/store.go +++ b/sei-cosmos/storev2/state/store.go @@ -19,8 +19,9 @@ import ( const StoreTypeSSStore = 100 var ( - _ types.KVStore = (*Store)(nil) - _ types.Queryable = (*Store)(nil) + _ types.KVStore = (*Store)(nil) + _ types.Queryable = (*Store)(nil) + _ types.ContextIterator = (*Store)(nil) ) // Store wraps a SS store and implements a cosmos KVStore @@ -71,15 +72,23 @@ func (st *Store) Delete(_ []byte) { } func (st *Store) Iterator(start, end []byte) types.Iterator { - itr, err := st.store.Iterator(st.storeKey.Name(), st.version, start, end) - if err != nil { - panic(err) - } - return itr + return st.iterator(context.Background(), start, end, true) +} + +func (st *Store) IteratorWithContext(ctx context.Context, start, end []byte) types.Iterator { + return st.iterator(ctx, start, end, true) } func (st *Store) ReverseIterator(start, end []byte) types.Iterator { - itr, err := st.store.ReverseIterator(st.storeKey.Name(), st.version, start, end) + return st.iterator(context.Background(), start, end, false) +} + +func (st *Store) ReverseIteratorWithContext(ctx context.Context, start, end []byte) types.Iterator { + return st.iterator(ctx, start, end, false) +} + +func (st *Store) iterator(ctx context.Context, start, end []byte, ascending bool) types.Iterator { + itr, err := seidbtypes.IterateWithContext(st.store, ctx, st.storeKey.Name(), st.version, start, end, !ascending) if err != nil { panic(err) } diff --git a/sei-cosmos/types/context.go b/sei-cosmos/types/context.go index 226339bc18..7f2d881b96 100644 --- a/sei-cosmos/types/context.go +++ b/sei-cosmos/types/context.go @@ -12,6 +12,7 @@ import ( tmbytes "github.com/sei-protocol/sei-chain/sei-tendermint/libs/bytes" tmproto "github.com/sei-protocol/sei-chain/sei-tendermint/proto/tendermint/types" + "github.com/sei-protocol/sei-chain/sei-cosmos/store/ctxkv" "github.com/sei-protocol/sei-chain/sei-cosmos/store/gaskv" stypes "github.com/sei-protocol/sei-chain/sei-cosmos/store/types" ) @@ -565,10 +566,10 @@ func (c Context) Value(key interface{}) interface{} { func (c Context) KVStore(key StoreKey) KVStore { if c.isTracing { if _, ok := c.nextStoreKeys[key.Name()]; ok { - return gaskv.NewStore(c.nextMs.GetKVStore(key), c.GasMeter(), stypes.KVGasConfig(), key.Name(), c.StoreTracer()) + return gaskv.NewStore(ctxkv.Wrap(c.nextMs.GetKVStore(key), c.ctx), c.GasMeter(), stypes.KVGasConfig(), key.Name(), c.StoreTracer()) } } - return gaskv.NewStore(c.MultiStore().GetKVStore(key), c.GasMeter(), stypes.KVGasConfig(), key.Name(), c.StoreTracer()) + return gaskv.NewStore(ctxkv.Wrap(c.MultiStore().GetKVStore(key), c.ctx), c.GasMeter(), stypes.KVGasConfig(), key.Name(), c.StoreTracer()) } func (c Context) GigaKVStore(key StoreKey) KVStore { @@ -579,10 +580,10 @@ func (c Context) GigaKVStore(key StoreKey) KVStore { func (c Context) TransientStore(key StoreKey) KVStore { if c.isTracing { if _, ok := c.nextStoreKeys[key.Name()]; ok { - return gaskv.NewStore(c.nextMs.GetKVStore(key), c.GasMeter(), stypes.TransientGasConfig(), key.Name(), c.StoreTracer()) + return gaskv.NewStore(ctxkv.Wrap(c.nextMs.GetKVStore(key), c.ctx), c.GasMeter(), stypes.TransientGasConfig(), key.Name(), c.StoreTracer()) } } - return gaskv.NewStore(c.MultiStore().GetKVStore(key), c.GasMeter(), stypes.TransientGasConfig(), key.Name(), c.StoreTracer()) + return gaskv.NewStore(ctxkv.Wrap(c.MultiStore().GetKVStore(key), c.ctx), c.GasMeter(), stypes.TransientGasConfig(), key.Name(), c.StoreTracer()) } // CacheContext returns a new Context with the multi-store cached and a new diff --git a/sei-db/db_engine/pebbledb/mvcc/db.go b/sei-db/db_engine/pebbledb/mvcc/db.go index 114659be57..5fe3cd72d1 100644 --- a/sei-db/db_engine/pebbledb/mvcc/db.go +++ b/sei-db/db_engine/pebbledb/mvcc/db.go @@ -31,6 +31,8 @@ import ( "github.com/sei-protocol/sei-chain/sei-db/wal" ) +var _ types.ContextIteratorStore = (*Database)(nil) + const ( VersionSize = 8 @@ -540,19 +542,25 @@ func (db *Database) compactPrunedRange(first, last []byte) error { // Iterator dispatches between descending- and ascending-mode implementations // depending on the on-disk encoding detected at open time. func (db *Database) Iterator(storeKey string, version int64, start, end []byte) (dbm.Iterator, error) { + return db.IteratorWithContext(context.Background(), storeKey, version, start, end) +} + +func (db *Database) IteratorWithContext(ctx context.Context, storeKey string, version int64, start, end []byte) (dbm.Iterator, error) { if db.descending { - return db.iteratorDescending(storeKey, version, start, end) + return db.iteratorDescending(ctx, storeKey, version, start, end) } - return db.iteratorAscending(storeKey, version, start, end) + return db.iteratorAscending(ctx, storeKey, version, start, end) } -// ReverseIterator dispatches between descending- and ascending-mode -// implementations depending on the on-disk encoding detected at open time. func (db *Database) ReverseIterator(storeKey string, version int64, start, end []byte) (dbm.Iterator, error) { + return db.ReverseIteratorWithContext(context.Background(), storeKey, version, start, end) +} + +func (db *Database) ReverseIteratorWithContext(ctx context.Context, storeKey string, version int64, start, end []byte) (dbm.Iterator, error) { if db.descending { - return db.reverseIteratorDescending(storeKey, version, start, end) + return db.reverseIteratorDescending(ctx, storeKey, version, start, end) } - return db.reverseIteratorAscending(storeKey, version, start, end) + return db.reverseIteratorAscending(ctx, storeKey, version, start, end) } // --------------------------------------------------------------------------- @@ -754,7 +762,7 @@ func (db *Database) pruneDescending(version int64) (_err error) { return db.compactPrunedRange(firstDeletedKey, lastDeletedKey) } -func (db *Database) iteratorDescending(storeKey string, version int64, start, end []byte) (dbm.Iterator, error) { +func (db *Database) iteratorDescending(ctx context.Context, storeKey string, version int64, start, end []byte) (dbm.Iterator, error) { if (start != nil && len(start) == 0) || (end != nil && len(end) == 0) { return nil, errorutils.ErrKeyEmpty } @@ -777,10 +785,10 @@ func (db *Database) iteratorDescending(storeKey string, version int64, start, en return nil, fmt.Errorf("failed to create PebbleDB iterator: %w", err) } - return newPebbleDBIterator(itr, storePrefix(storeKey), start, end, version, db.GetEarliestVersion(), false, db.config.UseDefaultComparer, storeKey, db.operationMetrics), nil + return finishMVCCIterator(newPebbleDBIterator(ctx, itr, storePrefix(storeKey), start, end, version, db.GetEarliestVersion(), false, db.config.UseDefaultComparer, storeKey, db.operationMetrics)) } -func (db *Database) reverseIteratorDescending(storeKey string, version int64, start, end []byte) (dbm.Iterator, error) { +func (db *Database) reverseIteratorDescending(ctx context.Context, storeKey string, version int64, start, end []byte) (dbm.Iterator, error) { if (start != nil && len(start) == 0) || (end != nil && len(end) == 0) { return nil, errorutils.ErrKeyEmpty } @@ -803,7 +811,7 @@ func (db *Database) reverseIteratorDescending(storeKey string, version int64, st return nil, fmt.Errorf("failed to create PebbleDB iterator: %w", err) } - return newPebbleDBIterator(itr, storePrefix(storeKey), start, end, version, db.GetEarliestVersion(), true, db.config.UseDefaultComparer, storeKey, db.operationMetrics), nil + return finishMVCCIterator(newPebbleDBIterator(ctx, itr, storePrefix(storeKey), start, end, version, db.GetEarliestVersion(), true, db.config.UseDefaultComparer, storeKey, db.operationMetrics)) } func getMVCCSliceDescending(db *pebble.DB, storeKey string, key []byte, version int64) (_ []byte, err error) { diff --git a/sei-db/db_engine/pebbledb/mvcc/db_ascending.go b/sei-db/db_engine/pebbledb/mvcc/db_ascending.go index 4075f9eea1..951123ec63 100644 --- a/sei-db/db_engine/pebbledb/mvcc/db_ascending.go +++ b/sei-db/db_engine/pebbledb/mvcc/db_ascending.go @@ -230,7 +230,7 @@ func (db *Database) pruneAscending(version int64) (_err error) { return db.compactPrunedRange(firstDeletedKey, lastDeletedKey) } -func (db *Database) iteratorAscending(storeKey string, version int64, start, end []byte) (dbm.Iterator, error) { +func (db *Database) iteratorAscending(ctx context.Context, storeKey string, version int64, start, end []byte) (dbm.Iterator, error) { if (start != nil && len(start) == 0) || (end != nil && len(end) == 0) { return nil, errorutils.ErrKeyEmpty } @@ -251,10 +251,10 @@ func (db *Database) iteratorAscending(storeKey string, version int64, start, end return nil, fmt.Errorf("failed to create PebbleDB iterator: %w", err) } - return newAscendingIterator(itr, storePrefix(storeKey), start, end, version, db.GetEarliestVersion(), false, storeKey, db.operationMetrics), nil + return finishMVCCIterator(newAscendingIterator(ctx, itr, storePrefix(storeKey), start, end, version, db.GetEarliestVersion(), false, storeKey, db.operationMetrics)) } -func (db *Database) reverseIteratorAscending(storeKey string, version int64, start, end []byte) (dbm.Iterator, error) { +func (db *Database) reverseIteratorAscending(ctx context.Context, storeKey string, version int64, start, end []byte) (dbm.Iterator, error) { if (start != nil && len(start) == 0) || (end != nil && len(end) == 0) { return nil, errorutils.ErrKeyEmpty } @@ -277,7 +277,7 @@ func (db *Database) reverseIteratorAscending(storeKey string, version int64, sta return nil, fmt.Errorf("failed to create PebbleDB iterator: %w", err) } - return newAscendingIterator(itr, storePrefix(storeKey), start, end, version, db.GetEarliestVersion(), true, storeKey, db.operationMetrics), nil + return finishMVCCIterator(newAscendingIterator(ctx, itr, storePrefix(storeKey), start, end, version, db.GetEarliestVersion(), true, storeKey, db.operationMetrics)) } func getMVCCSliceAscending(db *pebble.DB, storeKey string, key []byte, version int64) ([]byte, error) { diff --git a/sei-db/db_engine/pebbledb/mvcc/iterator.go b/sei-db/db_engine/pebbledb/mvcc/iterator.go index c9137c01df..dd2a76962a 100644 --- a/sei-db/db_engine/pebbledb/mvcc/iterator.go +++ b/sei-db/db_engine/pebbledb/mvcc/iterator.go @@ -36,11 +36,28 @@ type iterator struct { readCount int64 storeKey string operationMetrics *pebbledbmetrics.OperationMetrics + ctx context.Context + err error closeSync sync.Once } -func newPebbleDBIterator(src *pebble.Iterator, prefix, mvccStart, mvccEnd []byte, version int64, earliestVersion int64, reverse bool, useDefaultComparer bool, storeKey string, operationMetrics *pebbledbmetrics.OperationMetrics) *iterator { +func abortIfCancelled(ctx context.Context) error { + if ctx == nil { + return nil + } + return ctx.Err() +} + +func finishMVCCIterator(itr dbm.Iterator) (dbm.Iterator, error) { + if err := itr.Error(); err != nil { + _ = itr.Close() + return nil, err + } + return itr, nil +} + +func newPebbleDBIterator(ctx context.Context, src *pebble.Iterator, prefix, mvccStart, mvccEnd []byte, version int64, earliestVersion int64, reverse bool, useDefaultComparer bool, storeKey string, operationMetrics *pebbledbmetrics.OperationMetrics) *iterator { // Return invalid iterator if requested iterator height is lower than earliest version after pruning if version < earliestVersion { return &iterator{ @@ -54,6 +71,7 @@ func newPebbleDBIterator(src *pebble.Iterator, prefix, mvccStart, mvccEnd []byte useDefaultComparer: useDefaultComparer, storeKey: storeKey, operationMetrics: operationMetrics, + ctx: ctx, } } @@ -76,6 +94,7 @@ func newPebbleDBIterator(src *pebble.Iterator, prefix, mvccStart, mvccEnd []byte useDefaultComparer: useDefaultComparer, storeKey: storeKey, operationMetrics: operationMetrics, + ctx: ctx, } if valid { @@ -146,6 +165,11 @@ func (itr *iterator) nextLogicalKey(currKey []byte) ([]byte, bool) { func (itr *iterator) nextLogicalKeyByScan(currKey []byte) ([]byte, bool) { for valid := itr.source.Next(); valid; valid = itr.source.Next() { + if err := abortIfCancelled(itr.ctx); err != nil { + itr.err = err + itr.valid = false + return nil, false + } nextKey, _, ok := SplitMVCCKey(itr.source.Key()) if !ok || !bytes.HasPrefix(nextKey, itr.prefix) { return nil, false @@ -173,6 +197,11 @@ func (itr *iterator) prevLogicalKey(currKey []byte) ([]byte, bool) { func (itr *iterator) positionAtOrAfterKey(startKey []byte) { currentKey := startKey for { + if err := abortIfCancelled(itr.ctx); err != nil { + itr.err = err + itr.valid = false + return + } itr.valid = itr.seekVisibleVersionForKey(currentKey) if itr.valid && !itr.cursorTombstoned() { return @@ -189,6 +218,11 @@ func (itr *iterator) positionAtOrAfterKey(startKey []byte) { func (itr *iterator) positionAtOrBeforeKey(startKey []byte) { currentKey := startKey for { + if err := abortIfCancelled(itr.ctx); err != nil { + itr.err = err + itr.valid = false + return + } itr.valid = itr.seekVisibleVersionForKey(currentKey) if itr.valid && !itr.cursorTombstoned() { return @@ -284,14 +318,21 @@ func (itr *iterator) Next() { } else { itr.nextForward() } + if itr.err != nil { + panic(itr.err) + } if itr.Valid() { itr.readCount++ } } func (itr *iterator) Valid() bool { + if itr.err != nil { + itr.valid = false + return false + } // once invalid, forever invalid - if !itr.valid || !itr.source.Valid() { + if !itr.valid || itr.source == nil || !itr.source.Valid() { itr.valid = false return itr.valid } @@ -314,6 +355,12 @@ func (itr *iterator) Valid() bool { } func (itr *iterator) Error() error { + if itr.err != nil { + return itr.err + } + if itr.source == nil { + return nil + } return itr.source.Error() } diff --git a/sei-db/db_engine/pebbledb/mvcc/iterator_ascending.go b/sei-db/db_engine/pebbledb/mvcc/iterator_ascending.go index 065cb57698..a841cbfdc7 100644 --- a/sei-db/db_engine/pebbledb/mvcc/iterator_ascending.go +++ b/sei-db/db_engine/pebbledb/mvcc/iterator_ascending.go @@ -41,11 +41,13 @@ type ascendingIterator struct { readCount int64 storeKey string operationMetrics *pebbledbmetrics.OperationMetrics + ctx context.Context + err error closeSync sync.Once } -func newAscendingIterator(src *pebble.Iterator, prefix, mvccStart, mvccEnd []byte, version int64, earliestVersion int64, reverse bool, storeKey string, operationMetrics *pebbledbmetrics.OperationMetrics) *ascendingIterator { +func newAscendingIterator(ctx context.Context, src *pebble.Iterator, prefix, mvccStart, mvccEnd []byte, version int64, earliestVersion int64, reverse bool, storeKey string, operationMetrics *pebbledbmetrics.OperationMetrics) *ascendingIterator { // Return invalid iterator if requested iterator height is lower than earliest version after pruning if version < earliestVersion { return &ascendingIterator{ @@ -58,6 +60,7 @@ func newAscendingIterator(src *pebble.Iterator, prefix, mvccStart, mvccEnd []byt reverse: reverse, storeKey: storeKey, operationMetrics: operationMetrics, + ctx: ctx, } } @@ -79,6 +82,7 @@ func newAscendingIterator(src *pebble.Iterator, prefix, mvccStart, mvccEnd []byt reverse: reverse, storeKey: storeKey, operationMetrics: operationMetrics, + ctx: ctx, } if valid { @@ -154,6 +158,11 @@ func (itr *ascendingIterator) seekVisibleVersionForKey(targetKey []byte) bool { func (itr *ascendingIterator) nextLogicalKey(currKey []byte) ([]byte, bool) { seekKey := MVCCEncodeAscending(currKey, math.MaxInt64) for valid := itr.source.SeekGE(seekKey); valid; valid = itr.source.Next() { + if err := abortIfCancelled(itr.ctx); err != nil { + itr.err = err + itr.valid = false + return nil, false + } nextKey, _, ok := SplitMVCCKey(itr.source.Key()) if !ok || !bytes.HasPrefix(nextKey, itr.prefix) { return nil, false @@ -186,6 +195,11 @@ func (itr *ascendingIterator) prevLogicalKey(currKey []byte) ([]byte, bool) { func (itr *ascendingIterator) positionAtOrAfterKey(startKey []byte) { currentKey := startKey for { + if err := abortIfCancelled(itr.ctx); err != nil { + itr.err = err + itr.valid = false + return + } itr.valid = itr.seekVisibleVersionForKey(currentKey) if itr.valid && !itr.cursorTombstoned() { return @@ -206,6 +220,11 @@ func (itr *ascendingIterator) positionAtOrAfterKey(startKey []byte) { func (itr *ascendingIterator) positionAtOrBeforeKey(startKey []byte) { currentKey := startKey for { + if err := abortIfCancelled(itr.ctx); err != nil { + itr.err = err + itr.valid = false + return + } itr.valid = itr.seekVisibleVersionForKey(currentKey) if itr.valid && !itr.cursorTombstoned() { return @@ -302,14 +321,21 @@ func (itr *ascendingIterator) Next() { } else { itr.nextForward() } + if itr.err != nil { + panic(itr.err) + } if itr.Valid() { itr.readCount++ } } func (itr *ascendingIterator) Valid() bool { + if itr.err != nil { + itr.valid = false + return false + } // once invalid, forever invalid - if !itr.valid || !itr.source.Valid() { + if !itr.valid || itr.source == nil || !itr.source.Valid() { itr.valid = false return itr.valid } @@ -332,6 +358,12 @@ func (itr *ascendingIterator) Valid() bool { } func (itr *ascendingIterator) Error() error { + if itr.err != nil { + return itr.err + } + if itr.source == nil { + return nil + } return itr.source.Error() } diff --git a/sei-db/db_engine/pebbledb/mvcc/iterator_context_test.go b/sei-db/db_engine/pebbledb/mvcc/iterator_context_test.go new file mode 100644 index 0000000000..7cade177ea --- /dev/null +++ b/sei-db/db_engine/pebbledb/mvcc/iterator_context_test.go @@ -0,0 +1,64 @@ +package mvcc + +import ( + "context" + "fmt" + "testing" + + "github.com/stretchr/testify/require" +) + +const ctxIterStore = "store1" + +func TestIteratorWithCancelledContextAbortsSkip(t *testing.T) { + db := newTestDB(t, true) + require.True(t, db.descending) + + applyVersion(t, db, ctxIterStore, 10, []byte("aaa"), []byte("visible")) + for i := 0; i < 1000; i++ { + applyVersion(t, db, ctxIterStore, 1000, []byte(fmt.Sprintf("zzz%04d", i)), []byte("new")) + } + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + itr, err := db.IteratorWithContext(ctx, ctxIterStore, 10, nil, nil) + require.ErrorIs(t, err, context.Canceled) + require.Nil(t, itr) + + itr, err = db.ReverseIteratorWithContext(ctx, ctxIterStore, 10, nil, nil) + require.ErrorIs(t, err, context.Canceled) + require.Nil(t, itr) +} + +func TestIteratorNextPanicsAfterCancel(t *testing.T) { + db := newTestDB(t, true) + + applyVersion(t, db, ctxIterStore, 1, []byte("a"), []byte("va")) + applyVersion(t, db, ctxIterStore, 1, []byte("b"), []byte("vb")) + + ctx, cancel := context.WithCancel(context.Background()) + itr, err := db.IteratorWithContext(ctx, ctxIterStore, 1, nil, nil) + require.NoError(t, err) + defer func() { _ = itr.Close() }() + require.True(t, itr.Valid()) + + cancel() + require.Panics(t, func() { itr.Next() }) +} + +func TestAscendingIteratorWithCancelledContextAbortsSkip(t *testing.T) { + db := newAscendingIterTestDB(t) + + applyVersion(t, db, ctxIterStore, 10, []byte("aaa"), []byte("visible")) + for i := 0; i < 1000; i++ { + applyVersion(t, db, ctxIterStore, 1000, []byte(fmt.Sprintf("zzz%04d", i)), []byte("new")) + } + + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + itr, err := db.ReverseIteratorWithContext(ctx, ctxIterStore, 10, nil, nil) + require.ErrorIs(t, err, context.Canceled) + require.Nil(t, itr) +} diff --git a/sei-db/db_engine/rocksdb/mvcc/db.go b/sei-db/db_engine/rocksdb/mvcc/db.go index 1b8c45b590..6d109a28f9 100644 --- a/sei-db/db_engine/rocksdb/mvcc/db.go +++ b/sei-db/db_engine/rocksdb/mvcc/db.go @@ -5,6 +5,7 @@ package mvcc import ( "bytes" + "context" "encoding/binary" "fmt" "math" @@ -349,6 +350,14 @@ func (db *Database) ReverseIterator(storeKey string, version int64, start, end [ return NewRocksDBIterator(itr, readOpts, prefix, start, end, version, db.earliestVersion, true), nil } +func (db *Database) IteratorWithContext(_ context.Context, storeKey string, version int64, start, end []byte) (dbm.Iterator, error) { + return db.Iterator(storeKey, version, start, end) +} + +func (db *Database) ReverseIteratorWithContext(_ context.Context, storeKey string, version int64, start, end []byte) (dbm.Iterator, error) { + return db.ReverseIterator(storeKey, version, start, end) +} + // Import loads the initial version of the state in parallel with numWorkers goroutines // TODO: Potentially add retries instead of panics func (db *Database) Import(version int64, ch <-chan types.SnapshotNode) error { diff --git a/sei-db/db_engine/types/types.go b/sei-db/db_engine/types/types.go index 00096bf691..32d5192383 100644 --- a/sei-db/db_engine/types/types.go +++ b/sei-db/db_engine/types/types.go @@ -1,6 +1,7 @@ package types import ( + "context" "io" "github.com/sei-protocol/sei-chain/sei-db/proto" @@ -144,6 +145,28 @@ type StateStore interface { io.Closer } +// ContextIteratorStore is implemented by StateStores whose iterators can observe +// a deadline while skipping MVCC versions. Historical traces attach the RPC +// timeout here so a skip loop does not run for minutes after the caller gave up. +type ContextIteratorStore interface { + IteratorWithContext(ctx context.Context, storeKey string, version int64, start, end []byte) (dbm.Iterator, error) + ReverseIteratorWithContext(ctx context.Context, storeKey string, version int64, start, end []byte) (dbm.Iterator, error) +} + +// IterateWithContext prefers ContextIteratorStore when the store implements it. +func IterateWithContext(store StateStore, ctx context.Context, storeKey string, version int64, start, end []byte, reverse bool) (dbm.Iterator, error) { + if c, ok := store.(ContextIteratorStore); ok { + if reverse { + return c.ReverseIteratorWithContext(ctx, storeKey, version, start, end) + } + return c.IteratorWithContext(ctx, storeKey, version, start, end) + } + if reverse { + return store.ReverseIterator(storeKey, version, start, end) + } + return store.Iterator(storeKey, version, start, end) +} + type SnapshotNode struct { StoreKey string Key []byte diff --git a/sei-db/state_db/ss/composite/store.go b/sei-db/state_db/ss/composite/store.go index d88c77647e..cb8f4e8584 100644 --- a/sei-db/state_db/ss/composite/store.go +++ b/sei-db/state_db/ss/composite/store.go @@ -1,6 +1,7 @@ package composite import ( + "context" "encoding/binary" "fmt" "os" @@ -27,6 +28,7 @@ var logger = seilog.NewLogger("db", "state-db", "ss", "composite") // Compile-time check. var _ types.StateStore = (*CompositeStateStore)(nil) +var _ types.ContextIteratorStore = (*CompositeStateStore)(nil) // CompositeStateStore routes operations between Cosmos_SS and EVM_SS. // Both are db_engine.StateStore; the composite itself also implements db_engine.StateStore. @@ -207,6 +209,20 @@ func (s *CompositeStateStore) ReverseIterator(storeKey string, version int64, st return s.cosmosStore.ReverseIterator(storeKey, version, start, end) } +func (s *CompositeStateStore) IteratorWithContext(ctx context.Context, storeKey string, version int64, start, end []byte) (dbm.Iterator, error) { + if s.evmRouted(storeKey) { + return types.IterateWithContext(s.evmStore, ctx, storeKey, version, start, end, false) + } + return types.IterateWithContext(s.cosmosStore, ctx, storeKey, version, start, end, false) +} + +func (s *CompositeStateStore) ReverseIteratorWithContext(ctx context.Context, storeKey string, version int64, start, end []byte) (dbm.Iterator, error) { + if s.evmRouted(storeKey) { + return types.IterateWithContext(s.evmStore, ctx, storeKey, version, start, end, true) + } + return types.IterateWithContext(s.cosmosStore, ctx, storeKey, version, start, end, true) +} + func (s *CompositeStateStore) RawIterate(storeKey string, fn func([]byte, []byte, int64) bool) (bool, error) { return s.cosmosStore.RawIterate(storeKey, fn) } diff --git a/sei-db/state_db/ss/cosmos/store.go b/sei-db/state_db/ss/cosmos/store.go index 5b02d8ed15..5def9f960d 100644 --- a/sei-db/state_db/ss/cosmos/store.go +++ b/sei-db/state_db/ss/cosmos/store.go @@ -1,6 +1,8 @@ package cosmos import ( + "context" + dbm "github.com/tendermint/tm-db" "github.com/sei-protocol/sei-chain/sei-db/db_engine/types" @@ -9,6 +11,7 @@ import ( // Compile-time check: CosmosStateStore implements db_engine.StateStore. var _ types.StateStore = (*CosmosStateStore)(nil) +var _ types.ContextIteratorStore = (*CosmosStateStore)(nil) // CosmosStateStore wraps a single StateStore (MVCC DB) and satisfies db_engine.StateStore. // It is the SS-layer adapter for the main Cosmos state (all non-EVM modules). @@ -37,6 +40,14 @@ func (s *CosmosStateStore) ReverseIterator(storeKey string, version int64, start return s.db.ReverseIterator(storeKey, version, start, end) } +func (s *CosmosStateStore) IteratorWithContext(ctx context.Context, storeKey string, version int64, start, end []byte) (dbm.Iterator, error) { + return types.IterateWithContext(s.db, ctx, storeKey, version, start, end, false) +} + +func (s *CosmosStateStore) ReverseIteratorWithContext(ctx context.Context, storeKey string, version int64, start, end []byte) (dbm.Iterator, error) { + return types.IterateWithContext(s.db, ctx, storeKey, version, start, end, true) +} + func (s *CosmosStateStore) RawIterate(storeKey string, fn func([]byte, []byte, int64) bool) (bool, error) { return s.db.RawIterate(storeKey, fn) } diff --git a/sei-db/state_db/ss/evm/store.go b/sei-db/state_db/ss/evm/store.go index 01337940f3..64bd37a710 100644 --- a/sei-db/state_db/ss/evm/store.go +++ b/sei-db/state_db/ss/evm/store.go @@ -1,6 +1,7 @@ package evm import ( + "context" "fmt" "path/filepath" "sync" @@ -15,6 +16,7 @@ import ( ) var _ types.StateStore = (*EVMStateStore)(nil) +var _ types.ContextIteratorStore = (*EVMStateStore)(nil) // EVMStateStore manages either a single MVCC DB for all EVM data or one DB per // EVM sub-type, depending on config. In both modes, the logical store key and @@ -125,6 +127,28 @@ func (s *EVMStateStore) ReverseIterator(_ string, version int64, start, end []by return db.ReverseIterator(EVMStoreKey, version, start, end) } +func (s *EVMStateStore) IteratorWithContext(ctx context.Context, _ string, version int64, start, end []byte) (dbm.Iterator, error) { + if !s.separateDBs { + return types.IterateWithContext(s.primaryDB(), ctx, EVMStoreKey, version, start, end, false) + } + db := s.routeKey(start) + if db == nil { + return nil, fmt.Errorf("EVMStateStore: cannot route iteration for key") + } + return types.IterateWithContext(db, ctx, EVMStoreKey, version, start, end, false) +} + +func (s *EVMStateStore) ReverseIteratorWithContext(ctx context.Context, _ string, version int64, start, end []byte) (dbm.Iterator, error) { + if !s.separateDBs { + return types.IterateWithContext(s.primaryDB(), ctx, EVMStoreKey, version, start, end, true) + } + db := s.routeKey(start) + if db == nil { + return nil, fmt.Errorf("EVMStateStore: cannot route reverse iteration for key") + } + return types.IterateWithContext(db, ctx, EVMStoreKey, version, start, end, true) +} + func (s *EVMStateStore) RawIterate(_ string, _ func([]byte, []byte, int64) bool) (bool, error) { return false, fmt.Errorf("EVMStateStore: RawIterate not supported") } From bee64c411e19094a2f97ee1ee36cdf4c0144dc20 Mon Sep 17 00:00:00 2001 From: YimingZang Date: Mon, 17 Aug 2026 12:51:04 -0700 Subject: [PATCH 2/6] Address AI comment --- evmrpc/trace_baker.go | 2 +- evmrpc/trace_baker_test.go | 50 ++++++++++++++++++++++++++++++++------ 2 files changed, 44 insertions(+), 8 deletions(-) diff --git a/evmrpc/trace_baker.go b/evmrpc/trace_baker.go index 4cfe03a0b4..ab698d6807 100644 --- a/evmrpc/trace_baker.go +++ b/evmrpc/trace_baker.go @@ -191,7 +191,7 @@ func (b *TraceBaker) bakeBlockOneTracer(height int64, tracer string) bool { tracerName := tracer results, err := b.tracersAPI.TraceBlockByNumber(ctx, rpc.BlockNumber(height), &gethtracers.TraceConfig{Tracer: &tracerName}) - if err != nil { + if _, err = resultUnlessExpired(ctx, results, err); err != nil { atomic.AddUint64(&b.failed, 1) bakerLogger.Debug("trace baker block trace failed", "height", height, "tracer", tracer, "err", err) return false diff --git a/evmrpc/trace_baker_test.go b/evmrpc/trace_baker_test.go index 6e8434ef1b..687f7b1ab9 100644 --- a/evmrpc/trace_baker_test.go +++ b/evmrpc/trace_baker_test.go @@ -20,16 +20,20 @@ import ( // fakeTracerAPI drives the baker with controllable per-call results. type fakeTracerAPI struct { - mu sync.Mutex - calls int32 - results map[int64][]*gethtracers.TxTraceResult // keyed by height - errs map[int64]error - gate chan struct{} // when set, blocks each call until released - gates map[int64]chan struct{} // optional per-height gates + mu sync.Mutex + calls int32 + results map[int64][]*gethtracers.TxTraceResult // keyed by height + errs map[int64]error + gate chan struct{} // when set, blocks each call until released + gates map[int64]chan struct{} // optional per-height gates + synthesizeOnCancel bool // wait for ctx.Done, then return results with a nil error } -func (f *fakeTracerAPI) TraceBlockByNumber(_ context.Context, number rpc.BlockNumber, _ *gethtracers.TraceConfig) ([]*gethtracers.TxTraceResult, error) { +func (f *fakeTracerAPI) TraceBlockByNumber(ctx context.Context, number rpc.BlockNumber, _ *gethtracers.TraceConfig) ([]*gethtracers.TxTraceResult, error) { atomic.AddInt32(&f.calls, 1) + if f.synthesizeOnCancel { + <-ctx.Done() + } f.mu.Lock() gate := f.gate if f.gates != nil { @@ -151,6 +155,38 @@ func TestTraceBakerErrorBecomesFailedCount(t *testing.T) { require.Equal(t, uint64(0), b.BakedCount(), "errors should not count as baked") } +func TestTraceBakerDoesNotCacheExpiredSynthesizedTrace(t *testing.T) { + cache, err := keeper.NewTraceDB(t.TempDir()) + require.NoError(t, err) + defer cache.Close() + + tx := common.HexToHash("0xee") + api := &fakeTracerAPI{ + synthesizeOnCancel: true, + results: map[int64][]*gethtracers.TxTraceResult{ + 11: {{TxHash: tx, Result: map[string]string{"error": "store access cancelled"}}}, + }, + } + b := NewTraceBaker(nil, cache, TraceBakerConfig{ + Workers: 1, + QueueSize: 8, + BakeTimeout: time.Millisecond, + }) + b.tracersAPI = api + + require.False(t, b.bakeBlockOneTracer(11, "callTracer")) + require.Equal(t, uint64(1), b.FailedCount()) + require.Equal(t, uint64(0), b.BakedCount()) + + _, ok, err := cache.Get(11, "callTracer", tx) + require.NoError(t, err) + require.False(t, ok, "expired bake must not persist a per-tx row") + + _, ok, err = cache.GetBlock(11, "callTracer") + require.NoError(t, err) + require.False(t, ok, "expired bake must not persist a block row") +} + func TestTraceBakerSkipsNilOrErroredTxResults(t *testing.T) { // Per-tx errors come back as {Result: nil}; baker must skip them. cache, err := keeper.NewTraceDB(t.TempDir()) From 8b6955d060e8bb1560e01d52fdb8d224a546c733 Mon Sep 17 00:00:00 2001 From: YimingZang Date: Mon, 17 Aug 2026 15:22:04 -0700 Subject: [PATCH 3/6] Address comments --- sei-cosmos/baseapp/baseapp.go | 4 +-- sei-cosmos/baseapp/recovery.go | 18 +++++++--- sei-cosmos/baseapp/recovery_test.go | 29 +++++++++++++++ sei-cosmos/store/ctxkv/store_test.go | 13 +++++++ sei-cosmos/store/types/store.go | 5 +-- sei-db/db_engine/pebbledb/mvcc/iterator.go | 8 +++-- .../pebbledb/mvcc/iterator_context_test.go | 35 +++++++++++++++++++ 7 files changed, 101 insertions(+), 11 deletions(-) diff --git a/sei-cosmos/baseapp/baseapp.go b/sei-cosmos/baseapp/baseapp.go index b447c41ce5..34043eb30b 100644 --- a/sei-cosmos/baseapp/baseapp.go +++ b/sei-cosmos/baseapp/baseapp.go @@ -902,9 +902,9 @@ func (app *BaseApp) runTx(ctx sdk.Context, mode runTxMode, tx sdk.Tx, checksum [ blockGasMeter := ctx.GasMeter() defer func() { if r := recover(); r != nil { - recoveryMW := newOutOfGasRecoveryMiddleware(gasWanted, ctx, app.runTxRecoveryMiddleware) + recoveryMW := newContextCancelledRecoveryMiddleware(ctx, app.runTxRecoveryMiddleware) + recoveryMW = newOutOfGasRecoveryMiddleware(gasWanted, ctx, recoveryMW) recoveryMW = newOCCAbortRecoveryMiddleware(recoveryMW) // TODO: do we have to wrap with occ enabled check? - recoveryMW = newContextCancelledRecoveryMiddleware(recoveryMW) err, runTxRes.result = processRecovery(r, recoveryMW), nil } if ctx.GasMeter() == blockGasMeter { diff --git a/sei-cosmos/baseapp/recovery.go b/sei-cosmos/baseapp/recovery.go index 381e489160..1fa995bf1b 100644 --- a/sei-cosmos/baseapp/recovery.go +++ b/sei-cosmos/baseapp/recovery.go @@ -1,7 +1,6 @@ package baseapp import ( - "context" "errors" "fmt" "runtime/debug" @@ -85,15 +84,24 @@ func newOCCAbortRecoveryMiddleware(next recoveryMiddleware) recoveryMiddleware { return newRecoveryMiddleware(handler, next) } -// newContextCancelledRecoveryMiddleware recovers a store access that panicked -// because the caller's context was cancelled or exceeded its deadline. -func newContextCancelledRecoveryMiddleware(next recoveryMiddleware) recoveryMiddleware { +// newContextCancelledRecoveryMiddleware recovers a panic that is this +// context's already-expired cancel or deadline. Live contexts and other +// panic values fall through. +func newContextCancelledRecoveryMiddleware(ctx sdk.Context, next recoveryMiddleware) recoveryMiddleware { handler := func(recoveryObj interface{}) error { + stdCtx := ctx.Context() + if stdCtx == nil { + return nil + } + expired := stdCtx.Err() + if expired == nil { + return nil + } err, ok := recoveryObj.(error) if !ok { return nil } - if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { + if errors.Is(err, expired) { return err } return nil diff --git a/sei-cosmos/baseapp/recovery_test.go b/sei-cosmos/baseapp/recovery_test.go index b75892c638..5a71650749 100644 --- a/sei-cosmos/baseapp/recovery_test.go +++ b/sei-cosmos/baseapp/recovery_test.go @@ -1,9 +1,13 @@ package baseapp import ( + "context" + "errors" "fmt" "testing" + sdk "github.com/sei-protocol/sei-chain/sei-cosmos/types" + sdkerrors "github.com/sei-protocol/sei-chain/sei-cosmos/types/errors" "github.com/stretchr/testify/require" ) @@ -62,3 +66,28 @@ func TestRecoveryChain(t *testing.T) { require.Nil(t, receivedErr) } } + +func TestContextCancelledRecoveryOnlyWhenContextExpired(t *testing.T) { + defaultMW := newDefaultRecoveryMiddleware() + + live := sdk.Context{}.WithContext(context.Background()) + liveMW := newContextCancelledRecoveryMiddleware(live, defaultMW) + err := processRecovery(context.Canceled, liveMW) + require.ErrorIs(t, err, sdkerrors.ErrPanic) + err = processRecovery(fmt.Errorf("db: %w", context.DeadlineExceeded), liveMW) + require.ErrorIs(t, err, sdkerrors.ErrPanic) + + zeroMW := newContextCancelledRecoveryMiddleware(sdk.Context{}, defaultMW) + err = processRecovery(context.Canceled, zeroMW) + require.ErrorIs(t, err, sdkerrors.ErrPanic) + + expired, cancel := context.WithCancel(context.Background()) + cancel() + expiredMW := newContextCancelledRecoveryMiddleware(sdk.Context{}.WithContext(expired), defaultMW) + err = processRecovery(context.Canceled, expiredMW) + require.ErrorIs(t, err, context.Canceled) + require.False(t, errors.Is(err, sdkerrors.ErrPanic)) + + err = processRecovery(errors.New("unrelated panic"), expiredMW) + require.ErrorIs(t, err, sdkerrors.ErrPanic) +} diff --git a/sei-cosmos/store/ctxkv/store_test.go b/sei-cosmos/store/ctxkv/store_test.go index db49ace97d..3e6e6a8457 100644 --- a/sei-cosmos/store/ctxkv/store_test.go +++ b/sei-cosmos/store/ctxkv/store_test.go @@ -32,6 +32,19 @@ func TestWrapIsNoopWhenContextCannotCancel(t *testing.T) { require.Equal(t, parent, ctxkv.Wrap(parent, context.Background())) } +func TestIteratorOnUsesContextIteratorOnlyWhenCancellable(t *testing.T) { + parent := &recordingStore{} + _ = types.IteratorOn(parent, nil, nil, nil, true) + require.Nil(t, parent.gotCtx) + _ = types.IteratorOn(parent, context.Background(), nil, nil, true) + require.Nil(t, parent.gotCtx) + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + _ = types.IteratorOn(parent, ctx, nil, nil, true) + require.Equal(t, ctx, parent.gotCtx) +} + func TestWrapForwardsDeadlineToContextIterator(t *testing.T) { parent := &recordingStore{} ctx, cancel := context.WithCancel(context.Background()) diff --git a/sei-cosmos/store/types/store.go b/sei-cosmos/store/types/store.go index e247214a68..fb3b455cc1 100644 --- a/sei-cosmos/store/types/store.go +++ b/sei-cosmos/store/types/store.go @@ -259,9 +259,10 @@ type ContextIterator interface { ReverseIteratorWithContext(ctx context.Context, start, end []byte) Iterator } -// IteratorOn iterates store, using ctx when store implements ContextIterator. +// IteratorOn iterates store, using ctx when store implements ContextIterator +// and ctx can be cancelled. func IteratorOn(store KVStore, ctx context.Context, start, end []byte, ascending bool) Iterator { - if ctx != nil { + if ctx != nil && ctx.Done() != nil { if ci, ok := store.(ContextIterator); ok { if ascending { return ci.IteratorWithContext(ctx, start, end) diff --git a/sei-db/db_engine/pebbledb/mvcc/iterator.go b/sei-db/db_engine/pebbledb/mvcc/iterator.go index dd2a76962a..76ebc0d5e5 100644 --- a/sei-db/db_engine/pebbledb/mvcc/iterator.go +++ b/sei-db/db_engine/pebbledb/mvcc/iterator.go @@ -3,6 +3,7 @@ package mvcc import ( "bytes" "context" + "errors" "fmt" "math" "sync" @@ -43,14 +44,17 @@ type iterator struct { } func abortIfCancelled(ctx context.Context) error { - if ctx == nil { + if ctx == nil || ctx.Done() == nil { return nil } return ctx.Err() } +// finishMVCCIterator returns a construction-time cancel or deadline as an +// error so the caller can abort. Other iterator errors stay on the iterator. func finishMVCCIterator(itr dbm.Iterator) (dbm.Iterator, error) { - if err := itr.Error(); err != nil { + err := itr.Error() + if errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded) { _ = itr.Close() return nil, err } diff --git a/sei-db/db_engine/pebbledb/mvcc/iterator_context_test.go b/sei-db/db_engine/pebbledb/mvcc/iterator_context_test.go index 7cade177ea..c10cbd05b5 100644 --- a/sei-db/db_engine/pebbledb/mvcc/iterator_context_test.go +++ b/sei-db/db_engine/pebbledb/mvcc/iterator_context_test.go @@ -2,6 +2,7 @@ package mvcc import ( "context" + "errors" "fmt" "testing" @@ -62,3 +63,37 @@ func TestAscendingIteratorWithCancelledContextAbortsSkip(t *testing.T) { require.ErrorIs(t, err, context.Canceled) require.Nil(t, itr) } + +func TestFinishMVCCIteratorReturnsOnlyCancelErrors(t *testing.T) { + pebbleErr := errors.New("pebble: seek failed") + readErr := &finishIterStub{err: pebbleErr} + got, err := finishMVCCIterator(readErr) + require.NoError(t, err) + require.Equal(t, readErr, got) + require.False(t, readErr.closed) + + canceled := &finishIterStub{err: context.Canceled} + got, err = finishMVCCIterator(canceled) + require.ErrorIs(t, err, context.Canceled) + require.Nil(t, got) + require.True(t, canceled.closed) + + deadline := &finishIterStub{err: context.DeadlineExceeded} + got, err = finishMVCCIterator(deadline) + require.ErrorIs(t, err, context.DeadlineExceeded) + require.Nil(t, got) + require.True(t, deadline.closed) +} + +type finishIterStub struct { + err error + closed bool +} + +func (s *finishIterStub) Domain() ([]byte, []byte) { return nil, nil } +func (s *finishIterStub) Valid() bool { return false } +func (s *finishIterStub) Next() {} +func (s *finishIterStub) Key() []byte { return nil } +func (s *finishIterStub) Value() []byte { return nil } +func (s *finishIterStub) Error() error { return s.err } +func (s *finishIterStub) Close() error { s.closed = true; return nil } From 9c37151c8b350579d2e52ca038020bd9eb40c6fc Mon Sep 17 00:00:00 2001 From: YimingZang Date: Mon, 17 Aug 2026 16:08:47 -0700 Subject: [PATCH 4/6] Address blocker comments for beginblock --- evmrpc/initialize_block_test.go | 107 +++++++++++++++++++++++++++++++ evmrpc/simulate.go | 54 +++++++++++++--- evmrpc/watermark_manager_test.go | 4 ++ 3 files changed, 156 insertions(+), 9 deletions(-) create mode 100644 evmrpc/initialize_block_test.go diff --git a/evmrpc/initialize_block_test.go b/evmrpc/initialize_block_test.go new file mode 100644 index 0000000000..62657f9046 --- /dev/null +++ b/evmrpc/initialize_block_test.go @@ -0,0 +1,107 @@ +package evmrpc + +import ( + "context" + "errors" + "fmt" + "math/big" + "testing" + + ethtypes "github.com/ethereum/go-ethereum/core/types" + "github.com/ethereum/go-ethereum/trie" + "github.com/stretchr/testify/require" + + "github.com/sei-protocol/sei-chain/app/legacyabci" + sdk "github.com/sei-protocol/sei-chain/sei-cosmos/types" + abci "github.com/sei-protocol/sei-chain/sei-tendermint/abci/types" + "github.com/sei-protocol/sei-chain/sei-tendermint/rpc/coretypes" + tmtypes "github.com/sei-protocol/sei-chain/sei-tendermint/types" +) + +func TestReleaseOnContextPanic(t *testing.T) { + t.Parallel() + + var released int + release := func() { released++ } + + err := releaseOnContextPanic(release, context.Canceled) + require.ErrorIs(t, err, context.Canceled) + require.Equal(t, 1, released) + + err = releaseOnContextPanic(release, context.DeadlineExceeded) + require.ErrorIs(t, err, context.DeadlineExceeded) + require.Equal(t, 2, released) + + err = releaseOnContextPanic(release, fmt.Errorf("skip: %w", context.DeadlineExceeded)) + require.ErrorIs(t, err, context.DeadlineExceeded) + require.Equal(t, 3, released) + + require.Panics(t, func() { + _ = releaseOnContextPanic(release, errors.New("pebble seek failed")) + }) + require.Equal(t, 4, released) + + require.Panics(t, func() { + _ = releaseOnContextPanic(release, "not an error") + }) + require.Equal(t, 5, released) +} + +func TestInitializeBlockReleasesLeaseOnBeginBlockDeadline(t *testing.T) { + orig := runTraceBeginBlock + t.Cleanup(func() { runTraceBeginBlock = orig }) + runTraceBeginBlock = func(sdk.Context, int64, []abci.VoteInfo, []abci.Misbehavior, legacyabci.BeginBlockKeepers) { + panic(context.DeadlineExceeded) + } + + var released int + backend, block := newInitializeBlockTestBackend(t) + _, _, release, err := backend.initializeBlock(context.Background(), block, func(int64) (sdk.Context, func()) { + return sdk.Context{}, func() { released++ } + }) + require.ErrorIs(t, err, context.DeadlineExceeded) + require.Equal(t, 1, released, "base snapshot lease must be released on BeginBlock abort") + release() + require.Equal(t, 1, released, "returned release must be a no-op after recover") +} + +func TestInitializeBlockReleasesLeaseOnUnrelatedBeginBlockPanic(t *testing.T) { + orig := runTraceBeginBlock + t.Cleanup(func() { runTraceBeginBlock = orig }) + runTraceBeginBlock = func(sdk.Context, int64, []abci.VoteInfo, []abci.Misbehavior, legacyabci.BeginBlockKeepers) { + panic("boom") + } + + var released int + backend, block := newInitializeBlockTestBackend(t) + require.Panics(t, func() { + _, _, _, _ = backend.initializeBlock(context.Background(), block, func(int64) (sdk.Context, func()) { + return sdk.Context{}, func() { released++ } + }) + }) + require.Equal(t, 1, released) +} + +func newInitializeBlockTestBackend(t *testing.T) (*Backend, *ethtypes.Block) { + t.Helper() + tm := &fakeTMClient{ + status: &coretypes.ResultStatus{SyncInfo: coretypes.SyncInfo{LatestBlockHeight: 10, EarliestBlockHeight: 1}}, + blocksByHeight: map[int64]*coretypes.ResultBlock{ + 8: { + Block: &tmtypes.Block{ + Header: tmtypes.Header{Height: 8}, + LastCommit: &tmtypes.Commit{}, + }, + }, + }, + } + return &Backend{ + tmClient: tm, + watermarks: newTestWatermarkManager(tm, 10, nil, 10), + }, ethtypes.NewBlock( + ðtypes.Header{Number: big.NewInt(8), Time: 1, Difficulty: big.NewInt(0)}, + ðtypes.Body{}, + nil, + trie.NewStackTrie(nil), + ) +} diff --git a/evmrpc/simulate.go b/evmrpc/simulate.go index 27941a87e7..475a8c6cb2 100644 --- a/evmrpc/simulate.go +++ b/evmrpc/simulate.go @@ -686,13 +686,14 @@ func (b *Backend) StateAtBlock(ctx context.Context, block *ethtypes.Block, reexe return statedb, release, nil } -func (b *Backend) initializeBlock(ctx context.Context, block *ethtypes.Block, ctxProvider TraceContextProvider) (sdk.Context, *coretypes.ResultBlock, tracers.StateReleaseFunc, error) { +func (b *Backend) initializeBlock(ctx context.Context, block *ethtypes.Block, ctxProvider TraceContextProvider) (sdkCtx sdk.Context, tmBlock *coretypes.ResultBlock, release tracers.StateReleaseFunc, err error) { emptyRelease := func() {} + release = emptyRelease // get the parent block using block.parentHash prevBlockHeight := max(block.Number().Int64()-1, 0) blockNumber := block.Number().Int64() - tmBlock, err := blockByNumberRespectingWatermarks(ctx, b.tmClient, b.watermarks, &blockNumber, 1) + tmBlock, err = blockByNumberRespectingWatermarks(ctx, b.tmClient, b.watermarks, &blockNumber, 1) if err != nil { return sdk.Context{}, nil, emptyRelease, fmt.Errorf("cannot find block %d from tendermint", blockNumber) } @@ -704,22 +705,57 @@ func (b *Backend) initializeBlock(ctx context.Context, block *ethtypes.Block, ct reqBeginBlock := tmBlock.Block.ToReqBeginBlock(res.Validators) reqBeginBlock.Simulate = true baseCtx, baseRelease := ctxProvider(prevBlockHeight) - sdkCtx := baseCtx.WithBlockHeight(blockNumber).WithBlockTime(tmBlock.Block.Time) + var nextRelease func() + var released bool + release = func() { + if released { + return + } + released = true + if nextRelease != nil { + nextRelease() + } + baseRelease() + } + defer func() { + if r := recover(); r != nil { + err = releaseOnContextPanic(release, r) + release = emptyRelease + sdkCtx = sdk.Context{} + tmBlock = nil + } + }() + sdkCtx = baseCtx.WithBlockHeight(blockNumber).WithBlockTime(tmBlock.Block.Time) if ctx != nil { // The RPC/trace deadline must be on the SDK context so KVStore // iteration can pass it into the SS MVCC skip loops. sdkCtx = sdkCtx.WithContext(ctx) } - legacyabci.BeginBlock(sdkCtx, blockNumber, reqBeginBlock.LastCommitInfo.Votes, tmBlock.Block.Evidence.ToABCI(), b.beginBlockKeepers) - nextCtx, nextRelease := ctxProvider(sdkCtx.BlockHeight()) + runTraceBeginBlock(sdkCtx, blockNumber, reqBeginBlock.LastCommitInfo.Votes, tmBlock.Block.Evidence.ToABCI(), b.beginBlockKeepers) + var nextCtx sdk.Context + nextCtx, nextRelease = ctxProvider(sdkCtx.BlockHeight()) sdkCtx = sdkCtx.WithNextMs( nextCtx.MultiStore(), []string{"oracle", "oracle_mem"}, ) - return sdkCtx, tmBlock, func() { - nextRelease() - baseRelease() - }, nil + return sdkCtx, tmBlock, release, nil +} + +// runTraceBeginBlock is the BeginBlock used when reconstructing historical +// state for traces. +var runTraceBeginBlock = legacyabci.BeginBlock + +// releaseOnContextPanic releases leased snapshots and returns a cancel or +// deadline panic as an error. Other panics are re-raised after release. +func releaseOnContextPanic(release func(), recovered any) error { + if release != nil { + release() + } + err, ok := recovered.(error) + if ok && (errors.Is(err, context.Canceled) || errors.Is(err, context.DeadlineExceeded)) { + return err + } + panic(recovered) } func (b *Backend) GetEVM(_ context.Context, msg *core.Message, stateDB vm.StateDB, h *ethtypes.Header, vmConfig *vm.Config, blockCtx *vm.BlockContext) *vm.EVM { diff --git a/evmrpc/watermark_manager_test.go b/evmrpc/watermark_manager_test.go index 60c573b21a..63c0702fa4 100644 --- a/evmrpc/watermark_manager_test.go +++ b/evmrpc/watermark_manager_test.go @@ -382,6 +382,10 @@ func (f *fakeTMClient) Genesis(context.Context) (*coretypes.ResultGenesis, error return &coretypes.ResultGenesis{Genesis: &tmtypes.GenesisDoc{InitialHeight: 1}}, nil } +func (f *fakeTMClient) Validators(context.Context, *int64, *int, *int) (*coretypes.ResultValidators, error) { + return &coretypes.ResultValidators{}, nil +} + type fakeStateStore struct { latest int64 earliest int64 From fc1a430a445f66c8d00a2fc8d5365a66329208da Mon Sep 17 00:00:00 2001 From: YimingZang Date: Tue, 18 Aug 2026 00:01:34 -0700 Subject: [PATCH 5/6] Fix for trace_profile.go --- evmrpc/trace_profile.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/evmrpc/trace_profile.go b/evmrpc/trace_profile.go index 244836c5c6..97f4ae5283 100644 --- a/evmrpc/trace_profile.go +++ b/evmrpc/trace_profile.go @@ -116,7 +116,7 @@ func (api *DebugAPI) TraceTransactionProfile(ctx context.Context, hash common.Ha } api.clampDefaultStructLogLimit(config) traceResult, err := api.profiledTraceTx(ctx, tx, msg, txctx, blockCtx, statedb, config, nil, false, &phases.traceExecutionPhaseDurations) - if err != nil { + if _, err = resultUnlessExpired(ctx, traceResult, err); err != nil { return nil, err } From d9981daf7729a9215341b3666c7c1baf083b9e7e Mon Sep 17 00:00:00 2001 From: YimingZang Date: Tue, 18 Aug 2026 10:26:39 -0700 Subject: [PATCH 6/6] Address comment using t.context --- evmrpc/initialize_block_test.go | 4 ++-- evmrpc/trace_cancellation_test.go | 5 ++--- evmrpc/trace_transaction_timeout_test.go | 6 +++--- sei-cosmos/baseapp/recovery_test.go | 4 ++-- sei-cosmos/store/ctxkv/store_test.go | 6 ++---- sei-db/db_engine/pebbledb/mvcc/iterator_context_test.go | 6 +++--- 6 files changed, 14 insertions(+), 17 deletions(-) diff --git a/evmrpc/initialize_block_test.go b/evmrpc/initialize_block_test.go index 62657f9046..a75dbe644e 100644 --- a/evmrpc/initialize_block_test.go +++ b/evmrpc/initialize_block_test.go @@ -56,7 +56,7 @@ func TestInitializeBlockReleasesLeaseOnBeginBlockDeadline(t *testing.T) { var released int backend, block := newInitializeBlockTestBackend(t) - _, _, release, err := backend.initializeBlock(context.Background(), block, func(int64) (sdk.Context, func()) { + _, _, release, err := backend.initializeBlock(t.Context(), block, func(int64) (sdk.Context, func()) { return sdk.Context{}, func() { released++ } }) require.ErrorIs(t, err, context.DeadlineExceeded) @@ -75,7 +75,7 @@ func TestInitializeBlockReleasesLeaseOnUnrelatedBeginBlockPanic(t *testing.T) { var released int backend, block := newInitializeBlockTestBackend(t) require.Panics(t, func() { - _, _, _, _ = backend.initializeBlock(context.Background(), block, func(int64) (sdk.Context, func()) { + _, _, _, _ = backend.initializeBlock(t.Context(), block, func(int64) (sdk.Context, func()) { return sdk.Context{}, func() { released++ } }) }) diff --git a/evmrpc/trace_cancellation_test.go b/evmrpc/trace_cancellation_test.go index f6e550cf75..ea126ea130 100644 --- a/evmrpc/trace_cancellation_test.go +++ b/evmrpc/trace_cancellation_test.go @@ -9,8 +9,7 @@ import ( ) func TestResultUnlessExpiredReportsTheDeadline(t *testing.T) { - live, cancel := context.WithCancel(context.Background()) - defer cancel() + live := t.Context() traced := map[string]string{"gas": "0x1"} result, err := resultUnlessExpired(live, traced, nil) @@ -22,7 +21,7 @@ func TestResultUnlessExpiredReportsTheDeadline(t *testing.T) { require.ErrorIs(t, err, underlying) require.Nil(t, result) - expired, cancelExpired := context.WithCancel(context.Background()) + expired, cancelExpired := context.WithCancel(t.Context()) cancelExpired() result, err = resultUnlessExpired(expired, map[string]string{"error": "store access cancelled"}, nil) require.ErrorIs(t, err, context.Canceled) diff --git a/evmrpc/trace_transaction_timeout_test.go b/evmrpc/trace_transaction_timeout_test.go index af892fd4b8..6398fb6cab 100644 --- a/evmrpc/trace_transaction_timeout_test.go +++ b/evmrpc/trace_transaction_timeout_test.go @@ -46,7 +46,7 @@ func TestTraceTransactionTimeoutReleasesSemaphore(t *testing.T) { firstErr := make(chan error, 1) go func() { - _, err := api.TraceTransaction(context.Background(), backend.tx.Hash(), nil) + _, err := api.TraceTransaction(t.Context(), backend.tx.Hash(), nil) firstErr <- err }() @@ -56,7 +56,7 @@ func TestTraceTransactionTimeoutReleasesSemaphore(t *testing.T) { t.Fatal("StateAtTransaction did not enter the skip loop") } - _, busyErr := api.TraceTransaction(context.Background(), backend.tx.Hash(), nil) + _, busyErr := api.TraceTransaction(t.Context(), backend.tx.Hash(), nil) require.ErrorIs(t, busyErr, errTraceConcurrencyLimit) select { @@ -73,7 +73,7 @@ func TestTraceTransactionTimeoutReleasesSemaphore(t *testing.T) { t.Fatal("expected the timed-out trace to release the semaphore") } - _, err := api.TraceTransaction(context.Background(), backend.tx.Hash(), nil) + _, err := api.TraceTransaction(t.Context(), backend.tx.Hash(), nil) require.NotErrorIs(t, err, errTraceConcurrencyLimit) require.ErrorIs(t, err, context.DeadlineExceeded) } diff --git a/sei-cosmos/baseapp/recovery_test.go b/sei-cosmos/baseapp/recovery_test.go index 5a71650749..82211b93ca 100644 --- a/sei-cosmos/baseapp/recovery_test.go +++ b/sei-cosmos/baseapp/recovery_test.go @@ -70,7 +70,7 @@ func TestRecoveryChain(t *testing.T) { func TestContextCancelledRecoveryOnlyWhenContextExpired(t *testing.T) { defaultMW := newDefaultRecoveryMiddleware() - live := sdk.Context{}.WithContext(context.Background()) + live := sdk.Context{}.WithContext(t.Context()) liveMW := newContextCancelledRecoveryMiddleware(live, defaultMW) err := processRecovery(context.Canceled, liveMW) require.ErrorIs(t, err, sdkerrors.ErrPanic) @@ -81,7 +81,7 @@ func TestContextCancelledRecoveryOnlyWhenContextExpired(t *testing.T) { err = processRecovery(context.Canceled, zeroMW) require.ErrorIs(t, err, sdkerrors.ErrPanic) - expired, cancel := context.WithCancel(context.Background()) + expired, cancel := context.WithCancel(t.Context()) cancel() expiredMW := newContextCancelledRecoveryMiddleware(sdk.Context{}.WithContext(expired), defaultMW) err = processRecovery(context.Canceled, expiredMW) diff --git a/sei-cosmos/store/ctxkv/store_test.go b/sei-cosmos/store/ctxkv/store_test.go index 3e6e6a8457..1b9b4a624c 100644 --- a/sei-cosmos/store/ctxkv/store_test.go +++ b/sei-cosmos/store/ctxkv/store_test.go @@ -39,16 +39,14 @@ func TestIteratorOnUsesContextIteratorOnlyWhenCancellable(t *testing.T) { _ = types.IteratorOn(parent, context.Background(), nil, nil, true) require.Nil(t, parent.gotCtx) - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() + ctx := t.Context() _ = types.IteratorOn(parent, ctx, nil, nil, true) require.Equal(t, ctx, parent.gotCtx) } func TestWrapForwardsDeadlineToContextIterator(t *testing.T) { parent := &recordingStore{} - ctx, cancel := context.WithCancel(context.Background()) - defer cancel() + ctx := t.Context() wrapped := ctxkv.Wrap(parent, ctx) require.NotEqual(t, parent, wrapped) diff --git a/sei-db/db_engine/pebbledb/mvcc/iterator_context_test.go b/sei-db/db_engine/pebbledb/mvcc/iterator_context_test.go index c10cbd05b5..f36d09c5f7 100644 --- a/sei-db/db_engine/pebbledb/mvcc/iterator_context_test.go +++ b/sei-db/db_engine/pebbledb/mvcc/iterator_context_test.go @@ -20,7 +20,7 @@ func TestIteratorWithCancelledContextAbortsSkip(t *testing.T) { applyVersion(t, db, ctxIterStore, 1000, []byte(fmt.Sprintf("zzz%04d", i)), []byte("new")) } - ctx, cancel := context.WithCancel(context.Background()) + ctx, cancel := context.WithCancel(t.Context()) cancel() itr, err := db.IteratorWithContext(ctx, ctxIterStore, 10, nil, nil) @@ -38,7 +38,7 @@ func TestIteratorNextPanicsAfterCancel(t *testing.T) { applyVersion(t, db, ctxIterStore, 1, []byte("a"), []byte("va")) applyVersion(t, db, ctxIterStore, 1, []byte("b"), []byte("vb")) - ctx, cancel := context.WithCancel(context.Background()) + ctx, cancel := context.WithCancel(t.Context()) itr, err := db.IteratorWithContext(ctx, ctxIterStore, 1, nil, nil) require.NoError(t, err) defer func() { _ = itr.Close() }() @@ -56,7 +56,7 @@ func TestAscendingIteratorWithCancelledContextAbortsSkip(t *testing.T) { applyVersion(t, db, ctxIterStore, 1000, []byte(fmt.Sprintf("zzz%04d", i)), []byte("new")) } - ctx, cancel := context.WithCancel(context.Background()) + ctx, cancel := context.WithCancel(t.Context()) cancel() itr, err := db.ReverseIteratorWithContext(ctx, ctxIterStore, 10, nil, nil)