diff --git a/abci/client/client.go b/abci/client/client.go index 375317e04..27bf68c1d 100644 --- a/abci/client/client.go +++ b/abci/client/client.go @@ -3,9 +3,11 @@ package abcicli import ( "context" "fmt" + "sync" "github.com/tendermint/tendermint/abci/types" "github.com/tendermint/tendermint/libs/service" + tmsync "github.com/tendermint/tendermint/libs/sync" ) const ( @@ -24,10 +26,19 @@ type Client interface { service.Service types.Application + // TODO: remove as each method now returns an error Error() error // TODO: remove as this is not implemented Flush(context.Context) error Echo(context.Context, string) (*types.ResponseEcho, error) + + // FIXME: All other operations are run synchronously and rely + // on the caller to dictate concurrency (i.e. run a go routine), + // with the exception of `CheckTxAsync` which we maintain + // for the v0 mempool. We should explore refactoring the + // mempool to remove this vestige behavior. + SetResponseCallback(Callback) + CheckTxAsync(context.Context, *types.RequestCheckTx) (*ReqRes, error) } //---------------------------------------- @@ -45,3 +56,78 @@ func NewClient(addr, transport string, mustConnect bool) (client Client, err err } return } + +type Callback func(*types.Request, *types.Response) + +type ReqRes struct { + *types.Request + *sync.WaitGroup + *types.Response // Not set atomically, so be sure to use WaitGroup. + + mtx tmsync.Mutex + + // callbackInvoked as a variable to track if the callback was already + // invoked during the regular execution of the request. This variable + // allows clients to set the callback simultaneously without potentially + // invoking the callback twice by accident, once when 'SetCallback' is + // called and once during the normal request. + callbackInvoked bool + cb func(*types.Response) // A single callback that may be set. +} + +func NewReqRes(req *types.Request) *ReqRes { + return &ReqRes{ + Request: req, + WaitGroup: waitGroup1(), + Response: nil, + + callbackInvoked: false, + cb: nil, + } +} + +// Sets sets the callback. If reqRes is already done, it will call the cb +// immediately. Note, reqRes.cb should not change if reqRes.done and only one +// callback is supported. +func (r *ReqRes) SetCallback(cb func(res *types.Response)) { + r.mtx.Lock() + + if r.callbackInvoked { + r.mtx.Unlock() + cb(r.Response) + return + } + + r.cb = cb + r.mtx.Unlock() +} + +// InvokeCallback invokes a thread-safe execution of the configured callback +// if non-nil. +func (r *ReqRes) InvokeCallback() { + r.mtx.Lock() + defer r.mtx.Unlock() + + if r.cb != nil { + r.cb(r.Response) + } + r.callbackInvoked = true +} + +// GetCallback returns the configured callback of the ReqRes object which may be +// nil. Note, it is not safe to concurrently call this in cases where it is +// marked done and SetCallback is called before calling GetCallback as that +// will invoke the callback twice and create a potential race condition. +// +// ref: https://github.com/tendermint/tendermint/issues/5439 +func (r *ReqRes) GetCallback() func(*types.Response) { + r.mtx.Lock() + defer r.mtx.Unlock() + return r.cb +} + +func waitGroup1() (wg *sync.WaitGroup) { + wg = &sync.WaitGroup{} + wg.Add(1) + return +} diff --git a/abci/client/grpc_client.go b/abci/client/grpc_client.go index 3d2f833da..05eaf4ac5 100644 --- a/abci/client/grpc_client.go +++ b/abci/client/grpc_client.go @@ -23,18 +23,27 @@ type grpcClient struct { service.BaseService mustConnect bool - client types.ABCIClient - conn *grpc.ClientConn + client types.ABCIClient + conn *grpc.ClientConn + chReqRes chan *ReqRes // dispatches "async" responses to callbacks *in order*, needed by mempool - mtx sync.Mutex - addr string - err error + mtx sync.Mutex + addr string + err error + resCb func(*types.Request, *types.Response) // listens to all callbacks } func NewGRPCClient(addr string, mustConnect bool) Client { cli := &grpcClient{ addr: addr, mustConnect: mustConnect, + // Buffering the channel is needed to make calls appear asynchronous, + // which is required when the caller makes multiple async calls before + // processing callbacks (e.g. due to holding locks). 64 means that a + // caller can make up to 64 async calls before a callback must be + // processed (otherwise it deadlocks). It also means that we can make 64 + // gRPC calls while processing a slow callback at the channel head. + chReqRes: make(chan *ReqRes, 64), } cli.BaseService = *service.NewBaseService(nil, "grpcClient", cli) return cli @@ -49,6 +58,33 @@ func (cli *grpcClient) OnStart() error { return err } + // This processes asynchronous request/response messages and dispatches + // them to callbacks. + go func() { + // Use a separate function to use defer for mutex unlocks (this handles panics) + callCb := func(reqres *ReqRes) { + cli.mtx.Lock() + defer cli.mtx.Unlock() + + reqres.Done() + + // Notify client listener if set + if cli.resCb != nil { + cli.resCb(reqres.Request, reqres.Response) + } + + // Notify reqRes listener if set + reqres.InvokeCallback() + } + for reqres := range cli.chReqRes { + if reqres != nil { + callCb(reqres) + } else { + cli.Logger.Error("Received nil reqres") + } + } + }() + RETRY_LOOP: for { conn, err := grpc.Dial(cli.addr, @@ -115,6 +151,34 @@ func (cli *grpcClient) Error() error { return cli.err } +// Set listener for all responses +// NOTE: callback may get internally generated flush responses. +func (cli *grpcClient) SetResponseCallback(resCb Callback) { + cli.mtx.Lock() + cli.resCb = resCb + cli.mtx.Unlock() +} + +//---------------------------------------- + +func (cli *grpcClient) CheckTxAsync(ctx context.Context, req *types.RequestCheckTx) (*ReqRes, error) { + res, err := cli.client.CheckTx(ctx, req, grpc.WaitForReady(true)) + if err != nil { + cli.StopForError(err) + return nil, err + } + return cli.finishAsyncCall(types.ToRequestCheckTx(req), &types.Response{Value: &types.Response_CheckTx{CheckTx: res}}), nil +} + +// finishAsyncCall creates a ReqRes for an async call, and immediately populates it +// with the response. We don't complete it until it's been ordered via the channel. +func (cli *grpcClient) finishAsyncCall(req *types.Request, res *types.Response) *ReqRes { + reqres := NewReqRes(req) + reqres.Response = res + cli.chReqRes <- reqres // use channel for async responses, since they must be ordered + return reqres +} + //---------------------------------------- func (cli *grpcClient) Flush(ctx context.Context) error { return nil } @@ -124,11 +188,11 @@ func (cli *grpcClient) Echo(ctx context.Context, msg string) (*types.ResponseEch } func (cli *grpcClient) Info(ctx context.Context, req *types.RequestInfo) (*types.ResponseInfo, error) { - return cli.client.Info(ctx, types.ToRequestInfo(req).GetInfo(), grpc.WaitForReady(true)) + return cli.client.Info(ctx, req, grpc.WaitForReady(true)) } func (cli *grpcClient) CheckTx(ctx context.Context, req *types.RequestCheckTx) (*types.ResponseCheckTx, error) { - return cli.client.CheckTx(ctx, types.ToRequestCheckTx(req).GetCheckTx(), grpc.WaitForReady(true)) + return cli.client.CheckTx(ctx, req, grpc.WaitForReady(true)) } func (cli *grpcClient) Query(ctx context.Context, req *types.RequestQuery) (*types.ResponseQuery, error) { diff --git a/abci/client/local_client.go b/abci/client/local_client.go index c3b291ed0..fe6e0cfd9 100644 --- a/abci/client/local_client.go +++ b/abci/client/local_client.go @@ -17,6 +17,7 @@ type localClient struct { mtx *tmsync.Mutex types.Application + Callback } var _ Client = (*localClient)(nil) @@ -37,6 +38,41 @@ func NewLocalClient(mtx *tmsync.Mutex, app types.Application) Client { return cli } +func (app *localClient) SetResponseCallback(cb Callback) { + app.mtx.Lock() + app.Callback = cb + app.mtx.Unlock() +} + +func (app *localClient) CheckTxAsync(ctx context.Context, req *types.RequestCheckTx) (*ReqRes, error) { + app.mtx.Lock() + defer app.mtx.Unlock() + + res, err := app.Application.CheckTx(ctx, req) + if err != nil { + return nil, err + } + return app.callback( + types.ToRequestCheckTx(req), + types.ToResponseCheckTx(res), + ), nil +} + +func (app *localClient) callback(req *types.Request, res *types.Response) *ReqRes { + app.Callback(req, res) + rr := newLocalReqRes(req, res) + rr.callbackInvoked = true + return rr +} + +func newLocalReqRes(req *types.Request, res *types.Response) *ReqRes { + reqRes := NewReqRes(req) + reqRes.Response = res + return reqRes +} + +//------------------------------------------------------- + func (app *localClient) Error() error { return nil } diff --git a/abci/client/mocks/client.go b/abci/client/mocks/client.go index 6b21c0b50..8892712ab 100644 --- a/abci/client/mocks/client.go +++ b/abci/client/mocks/client.go @@ -5,9 +5,12 @@ package mocks import ( context "context" - mock "github.com/stretchr/testify/mock" + abcicli "github.com/tendermint/tendermint/abci/client" + log "github.com/tendermint/tendermint/libs/log" + mock "github.com/stretchr/testify/mock" + types "github.com/tendermint/tendermint/abci/types" ) @@ -62,6 +65,29 @@ func (_m *Client) CheckTx(_a0 context.Context, _a1 *types.RequestCheckTx) (*type return r0, r1 } +// CheckTxAsync provides a mock function with given fields: _a0, _a1 +func (_m *Client) CheckTxAsync(_a0 context.Context, _a1 *types.RequestCheckTx) (*abcicli.ReqRes, error) { + ret := _m.Called(_a0, _a1) + + var r0 *abcicli.ReqRes + if rf, ok := ret.Get(0).(func(context.Context, *types.RequestCheckTx) *abcicli.ReqRes); ok { + r0 = rf(_a0, _a1) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*abcicli.ReqRes) + } + } + + var r1 error + if rf, ok := ret.Get(1).(func(context.Context, *types.RequestCheckTx) error); ok { + r1 = rf(_a0, _a1) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + // Commit provides a mock function with given fields: _a0, _a1 func (_m *Client) Commit(_a0 context.Context, _a1 *types.RequestCommit) (*types.ResponseCommit, error) { ret := _m.Called(_a0, _a1) @@ -425,6 +451,11 @@ func (_m *Client) SetLogger(_a0 log.Logger) { _m.Called(_a0) } +// SetResponseCallback provides a mock function with given fields: _a0 +func (_m *Client) SetResponseCallback(_a0 abcicli.Callback) { + _m.Called(_a0) +} + // Start provides a mock function with given fields: func (_m *Client) Start() error { ret := _m.Called() diff --git a/abci/client/socket_client.go b/abci/client/socket_client.go index 9f80d8953..de89c3ba9 100644 --- a/abci/client/socket_client.go +++ b/abci/client/socket_client.go @@ -16,6 +16,10 @@ import ( "github.com/tendermint/tendermint/libs/service" ) +const ( + reqQueueSize = 256 // TODO make configurable +) + // socketClient is the client side implementation of the Tendermint // Socket Protocol (TSP). It is used by an instance of Tendermint to pass // ABCI requests to an out of process application running the socketServer. @@ -31,11 +35,12 @@ type socketClient struct { mustConnect bool conn net.Conn - reqQueue chan *requestAndResponse + reqQueue chan *ReqRes mtx sync.Mutex err error - reqSent *list.List // list of requests sent, waiting for response + reqSent *list.List // list of requests sent, waiting for response + resCb func(*types.Request, *types.Response) // called on all requests, if set. } var _ Client = (*socketClient)(nil) @@ -45,10 +50,11 @@ var _ Client = (*socketClient)(nil) // if it fails to connect else it will continue to retry. func NewSocketClient(addr string, mustConnect bool) Client { cli := &socketClient{ - reqQueue: make(chan *requestAndResponse), + reqQueue: make(chan *ReqRes, reqQueueSize), mustConnect: mustConnect, addr: addr, reqSent: list.New(), + resCb: nil, } cli.BaseService = *service.NewBaseService(nil, "socketClient", cli) return cli @@ -87,7 +93,8 @@ func (cli *socketClient) OnStop() { if cli.conn != nil { cli.conn.Close() } - cli.drainQueue() + + cli.flushQueue() } // Error returns an error if the client was stopped abruptly. @@ -99,6 +106,22 @@ func (cli *socketClient) Error() error { //---------------------------------------- +// SetResponseCallback sets a callback, which will be executed for each +// non-error & non-empty response from the server. +// +// NOTE: callback may get internally generated flush responses. +func (cli *socketClient) SetResponseCallback(resCb Callback) { + cli.mtx.Lock() + cli.resCb = resCb + cli.mtx.Unlock() +} + +func (cli *socketClient) CheckTxAsync(ctx context.Context, req *types.RequestCheckTx) (*ReqRes, error) { + return cli.queueRequest(ctx, types.ToRequestCheckTx(req)) +} + +//---------------------------------------- + func (cli *socketClient) sendRequestsRoutine(conn io.Writer) { bw := bufio.NewWriter(conn) for { @@ -153,7 +176,7 @@ func (cli *socketClient) recvResponseRoutine(conn io.Reader) { } } -func (cli *socketClient) trackRequest(reqres *requestAndResponse) { +func (cli *socketClient) trackRequest(reqres *ReqRes) { // N.B. We must NOT hold the client state lock while checking this, or we // may deadlock with shutdown. if !cli.IsRunning() { @@ -175,60 +198,70 @@ func (cli *socketClient) didRecvResponse(res *types.Response) error { return fmt.Errorf("unexpected response %T when no call was made", res.Value) } - reqres := next.Value.(*requestAndResponse) + reqres := next.Value.(*ReqRes) if !resMatchesReq(reqres.Request, res) { return fmt.Errorf("unexpected response %T to the request %T", res.Value, reqres.Request.Value) } reqres.Response = res - reqres.markDone() // release waiters + reqres.Done() // release waiters cli.reqSent.Remove(next) // pop first item from linked list + // Notify client listener if set (global callback). + if cli.resCb != nil { + cli.resCb(reqres.Request, res) + } + + // Notify reqRes listener if set (request specific callback). + // + // NOTE: It is possible this callback isn't set on the reqres object. At this + // point, in which case it will be called after, when it is set. + reqres.InvokeCallback() + return nil } -func (cli *socketClient) doRequest(ctx context.Context, req *types.Request) (*types.Response, error) { - if !cli.IsRunning() { - return nil, errors.New("client has stopped") - } - - reqres := makeReqRes(req) +func (cli *socketClient) queueRequest(ctx context.Context, req *types.Request) (*ReqRes, error) { + reqres := NewReqRes(req) + // TODO: set cli.err if reqQueue times out select { case cli.reqQueue <- reqres: - case <-ctx.Done(): - return nil, fmt.Errorf("can't queue req: %w", ctx.Err()) - } - - select { - case <-reqres.signal: - if err := cli.Error(); err != nil { - return nil, err - } - - return reqres.Response, nil case <-ctx.Done(): return nil, ctx.Err() } + + return reqres, nil } -// drainQueue marks as complete and discards all remaining pending requests +// flushQueue marks as complete and discards all remaining pending requests // from the queue. -func (cli *socketClient) drainQueue() { +func (cli *socketClient) flushQueue() { cli.mtx.Lock() defer cli.mtx.Unlock() // mark all in-flight messages as resolved (they will get cli.Error()) for req := cli.reqSent.Front(); req != nil; req = req.Next() { - reqres := req.Value.(*requestAndResponse) - reqres.markDone() + reqres := req.Value.(*ReqRes) + reqres.Done() + } + + // mark all queued messages as resolved +LOOP: + for { + select { + case reqres := <-cli.reqQueue: + reqres.Done() + default: + break LOOP + } } } //---------------------------------------- func (cli *socketClient) Flush(ctx context.Context) error { - _, err := cli.doRequest(ctx, types.ToRequestFlush()) + _, err := cli.queueRequest(ctx, types.ToRequestFlush()) if err != nil { return err } @@ -236,107 +269,120 @@ func (cli *socketClient) Flush(ctx context.Context) error { } func (cli *socketClient) Echo(ctx context.Context, msg string) (*types.ResponseEcho, error) { - res, err := cli.doRequest(ctx, types.ToRequestEcho(msg)) + reqRes, err := cli.queueRequest(ctx, types.ToRequestEcho(msg)) if err != nil { return nil, err } - return res.GetEcho(), nil + reqRes.Wait() + return reqRes.Response.GetEcho(), cli.Error() } func (cli *socketClient) Info(ctx context.Context, req *types.RequestInfo) (*types.ResponseInfo, error) { - res, err := cli.doRequest(ctx, types.ToRequestInfo(req)) + reqRes, err := cli.queueRequest(ctx, types.ToRequestInfo(req)) if err != nil { return nil, err } - return res.GetInfo(), nil + reqRes.Wait() + return reqRes.Response.GetInfo(), cli.Error() } func (cli *socketClient) CheckTx(ctx context.Context, req *types.RequestCheckTx) (*types.ResponseCheckTx, error) { - res, err := cli.doRequest(ctx, types.ToRequestCheckTx(req)) + reqRes, err := cli.queueRequest(ctx, types.ToRequestCheckTx(req)) if err != nil { return nil, err } - return res.GetCheckTx(), nil + reqRes.Wait() + return reqRes.Response.GetCheckTx(), cli.Error() } func (cli *socketClient) Query(ctx context.Context, req *types.RequestQuery) (*types.ResponseQuery, error) { - res, err := cli.doRequest(ctx, types.ToRequestQuery(req)) + reqRes, err := cli.queueRequest(ctx, types.ToRequestQuery(req)) if err != nil { return nil, err } - return res.GetQuery(), nil + reqRes.Wait() + return reqRes.Response.GetQuery(), cli.Error() } func (cli *socketClient) Commit(ctx context.Context, req *types.RequestCommit) (*types.ResponseCommit, error) { - res, err := cli.doRequest(ctx, types.ToRequestCommit()) + reqRes, err := cli.queueRequest(ctx, types.ToRequestCommit()) if err != nil { return nil, err } - return res.GetCommit(), nil + reqRes.Wait() + return reqRes.Response.GetCommit(), cli.Error() } func (cli *socketClient) InitChain(ctx context.Context, req *types.RequestInitChain) (*types.ResponseInitChain, error) { - res, err := cli.doRequest(ctx, types.ToRequestInitChain(req)) + reqRes, err := cli.queueRequest(ctx, types.ToRequestInitChain(req)) if err != nil { return nil, err } - return res.GetInitChain(), nil + reqRes.Wait() + return reqRes.Response.GetInitChain(), cli.Error() } func (cli *socketClient) ListSnapshots(ctx context.Context, req *types.RequestListSnapshots) (*types.ResponseListSnapshots, error) { - res, err := cli.doRequest(ctx, types.ToRequestListSnapshots(req)) + reqRes, err := cli.queueRequest(ctx, types.ToRequestListSnapshots(req)) if err != nil { return nil, err } - return res.GetListSnapshots(), nil + reqRes.Wait() + return reqRes.Response.GetListSnapshots(), cli.Error() } func (cli *socketClient) OfferSnapshot(ctx context.Context, req *types.RequestOfferSnapshot) (*types.ResponseOfferSnapshot, error) { - res, err := cli.doRequest(ctx, types.ToRequestOfferSnapshot(req)) + reqRes, err := cli.queueRequest(ctx, types.ToRequestOfferSnapshot(req)) if err != nil { return nil, err } - return res.GetOfferSnapshot(), nil + reqRes.Wait() + return reqRes.Response.GetOfferSnapshot(), cli.Error() } func (cli *socketClient) LoadSnapshotChunk(ctx context.Context, req *types.RequestLoadSnapshotChunk) (*types.ResponseLoadSnapshotChunk, error) { - res, err := cli.doRequest(ctx, types.ToRequestLoadSnapshotChunk(req)) + reqRes, err := cli.queueRequest(ctx, types.ToRequestLoadSnapshotChunk(req)) if err != nil { return nil, err } - return res.GetLoadSnapshotChunk(), nil + reqRes.Wait() + return reqRes.Response.GetLoadSnapshotChunk(), cli.Error() } func (cli *socketClient) ApplySnapshotChunk(ctx context.Context, req *types.RequestApplySnapshotChunk) (*types.ResponseApplySnapshotChunk, error) { - res, err := cli.doRequest(ctx, types.ToRequestApplySnapshotChunk(req)) + reqRes, err := cli.queueRequest(ctx, types.ToRequestApplySnapshotChunk(req)) if err != nil { return nil, err } - return res.GetApplySnapshotChunk(), nil + reqRes.Wait() + return reqRes.Response.GetApplySnapshotChunk(), cli.Error() } func (cli *socketClient) PrepareProposal(ctx context.Context, req *types.RequestPrepareProposal) (*types.ResponsePrepareProposal, error) { - res, err := cli.doRequest(ctx, types.ToRequestPrepareProposal(req)) + reqRes, err := cli.queueRequest(ctx, types.ToRequestPrepareProposal(req)) if err != nil { return nil, err } - return res.GetPrepareProposal(), nil + reqRes.Wait() + return reqRes.Response.GetPrepareProposal(), cli.Error() } func (cli *socketClient) ProcessProposal(ctx context.Context, req *types.RequestProcessProposal) (*types.ResponseProcessProposal, error) { - res, err := cli.doRequest(ctx, types.ToRequestProcessProposal(req)) + reqRes, err := cli.queueRequest(ctx, types.ToRequestProcessProposal(req)) if err != nil { return nil, err } - return res.GetProcessProposal(), nil + reqRes.Wait() + return reqRes.Response.GetProcessProposal(), cli.Error() } func (cli *socketClient) FinalizeBlock(ctx context.Context, req *types.RequestFinalizeBlock) (*types.ResponseFinalizeBlock, error) { - res, err := cli.doRequest(ctx, types.ToRequestFinalizeBlock(req)) + reqRes, err := cli.queueRequest(ctx, types.ToRequestFinalizeBlock(req)) if err != nil { return nil, err } - return res.GetFinalizeBlock(), nil + reqRes.Wait() + return reqRes.Response.GetFinalizeBlock(), cli.Error() } //---------------------------------------- @@ -391,27 +437,3 @@ func (cli *socketClient) stopForError(err error) { cli.Logger.Error("Error stopping abci.socketClient", "err", err) } } - -type requestAndResponse struct { - *types.Request - *types.Response - - mtx sync.Mutex - signal chan struct{} -} - -func makeReqRes(req *types.Request) *requestAndResponse { - return &requestAndResponse{ - Request: req, - Response: nil, - signal: make(chan struct{}), - } -} - -// markDone marks the ReqRes object as done. -func (r *requestAndResponse) markDone() { - r.mtx.Lock() - defer r.mtx.Unlock() - - close(r.signal) -} diff --git a/abci/client/socket_client_test.go b/abci/client/socket_client_test.go index b8b44a929..cbc0c5dd6 100644 --- a/abci/client/socket_client_test.go +++ b/abci/client/socket_client_test.go @@ -3,6 +3,7 @@ package abcicli_test import ( "context" "fmt" + "sync" "testing" "time" @@ -70,3 +71,79 @@ func setupClientServer(t *testing.T, app types.Application) ( return s, c } + +// TestCallbackInvokedWhenSetLaet ensures that the callback is invoked when +// set after the client completes the call into the app. Currently this +// test relies on the callback being allowed to be invoked twice if set multiple +// times, once when set early and once when set late. +func TestCallbackInvokedWhenSetLate(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + wg := &sync.WaitGroup{} + wg.Add(1) + app := blockedABCIApplication{ + wg: wg, + } + _, c := setupClientServer(t, app) + reqRes, err := c.CheckTxAsync(ctx, &types.RequestCheckTx{}) + require.NoError(t, err) + + done := make(chan struct{}) + cb := func(_ *types.Response) { + close(done) + } + reqRes.SetCallback(cb) + app.wg.Done() + <-done + + var called bool + cb = func(_ *types.Response) { + called = true + } + reqRes.SetCallback(cb) + require.True(t, called) +} + +type blockedABCIApplication struct { + wg *sync.WaitGroup + types.BaseApplication +} + +func (b blockedABCIApplication) CheckTxAsync(ctx context.Context, r *types.RequestCheckTx) (*types.ResponseCheckTx, error) { + b.wg.Wait() + return b.BaseApplication.CheckTx(ctx, r) +} + +// TestCallbackInvokedWhenSetEarly ensures that the callback is invoked when +// set before the client completes the call into the app. +func TestCallbackInvokedWhenSetEarly(t *testing.T) { + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + wg := &sync.WaitGroup{} + wg.Add(1) + app := blockedABCIApplication{ + wg: wg, + } + _, c := setupClientServer(t, app) + reqRes, err := c.CheckTxAsync(ctx, &types.RequestCheckTx{}) + require.NoError(t, err) + + done := make(chan struct{}) + cb := func(_ *types.Response) { + close(done) + } + reqRes.SetCallback(cb) + app.wg.Done() + + called := func() bool { + select { + case <-done: + return true + default: + return false + } + } + require.Eventually(t, called, time.Second, time.Millisecond*25) +} diff --git a/abci/client/unsync_local_client.go b/abci/client/unsync_local_client.go index 0d0cb0ab8..0f37c0fda 100644 --- a/abci/client/unsync_local_client.go +++ b/abci/client/unsync_local_client.go @@ -2,6 +2,7 @@ package abcicli import ( "context" + "sync" types "github.com/tendermint/tendermint/abci/types" "github.com/tendermint/tendermint/libs/service" @@ -11,6 +12,10 @@ type unsyncLocalClient struct { service.BaseService types.Application + + // This mutex is exclusively used to protect the callback. + mtx sync.RWMutex + Callback } var _ Client = (*unsyncLocalClient)(nil) @@ -40,3 +45,31 @@ func (app *unsyncLocalClient) Flush(_ context.Context) error { func (app *unsyncLocalClient) Echo(ctx context.Context, msg string) (*types.ResponseEcho, error) { return &types.ResponseEcho{Message: msg}, nil } + +//------------------------------------------------------- + +func (app *unsyncLocalClient) SetResponseCallback(cb Callback) { + app.mtx.Lock() + defer app.mtx.Unlock() + app.Callback = cb +} + +func (app *unsyncLocalClient) CheckTxAsync(ctx context.Context, req *types.RequestCheckTx) (*ReqRes, error) { + res, err := app.Application.CheckTx(ctx, req) + if err != nil { + return nil, err + } + return app.callback( + types.ToRequestCheckTx(req), + types.ToResponseCheckTx(res), + ), nil +} + +func (app *unsyncLocalClient) callback(req *types.Request, res *types.Response) *ReqRes { + app.mtx.RLock() + defer app.mtx.RUnlock() + app.Callback(req, res) + rr := newLocalReqRes(req, res) + rr.callbackInvoked = true + return rr +} diff --git a/mempool/v0/clist_mempool.go b/mempool/v0/clist_mempool.go index 47b7a85ed..f1e323c79 100644 --- a/mempool/v0/clist_mempool.go +++ b/mempool/v0/clist_mempool.go @@ -93,6 +93,8 @@ func NewCListMempool( mp.cache = mempool.NopTxCache{} } + proxyAppConn.SetResponseCallback(mp.globalCb) + for _, option := range options { option(mp) } @@ -250,18 +252,36 @@ func (mem *CListMempool) CheckTx( return mempool.ErrTxInCache } - resp, err := mem.proxyAppConn.CheckTx(context.TODO(), &abci.RequestCheckTx{Tx: tx}) + reqRes, err := mem.proxyAppConn.CheckTxAsync(context.TODO(), &abci.RequestCheckTx{Tx: tx}) if err != nil { - mem.cache.Remove(tx) - mem.logger.Error("error from check tx", "err", err) return err } - - mem.addNewTransaction(tx, txInfo.SenderID, txInfo.SenderP2PID, cb, resp) + reqRes.SetCallback(mem.reqResCb(tx, txInfo.SenderID, txInfo.SenderP2PID, cb)) return nil } +// Global callback that will be called after every ABCI response. +// Having a single global callback avoids needing to set a callback for each request. +// However, processing the checkTx response requires the peerID (so we can track which txs we heard from who), +// and peerID is not included in the ABCI request, so we have to set request-specific callbacks that +// include this information. If we're not in the midst of a recheck, this function will just return, +// so the request specific callback can do the work. +// +// When rechecking, we don't need the peerID, so the recheck callback happens +// here. +func (mem *CListMempool) globalCb(req *abci.Request, res *abci.Response) { + if mem.recheckCursor == nil { + return + } + + mem.metrics.RecheckTimes.Add(1) + mem.resCbRecheck(req, res) + + // update metrics + mem.metrics.Size.Set(float64(mem.Size())) +} + // Request specific callback that should be set on individual reqRes objects // to incorporate local information when processing the response. // This allows us to track the peer that sent us this tx, so we can avoid sending it back to them. @@ -271,70 +291,27 @@ func (mem *CListMempool) CheckTx( // when all other response processing is complete. // // Used in CheckTx to record PeerID who sent us the tx. -func (mem *CListMempool) addNewTransaction( +func (mem *CListMempool) reqResCb( tx []byte, peerID uint16, peerP2PID p2p.ID, externalCb func(*abci.ResponseCheckTx), - resp *abci.ResponseCheckTx, -) { - if mem.recheckCursor != nil { - // this should never happen - panic("recheck cursor is not nil in reqResCb") - } - - var postCheckErr error - if mem.postCheck != nil { - postCheckErr = mem.postCheck(tx, resp) - } - if (resp.Code == abci.CodeTypeOK) && postCheckErr == nil { - // Check mempool isn't full again to reduce the chance of exceeding the - // limits. - if err := mem.isFull(len(tx)); err != nil { - // remove from cache (mempool might have a space later) - mem.cache.Remove(tx) - mem.logger.Error(err.Error()) - return +) func(res *abci.Response) { + return func(res *abci.Response) { + if mem.recheckCursor != nil { + // this should never happen + panic("recheck cursor is not nil in reqResCb") } - memTx := &mempoolTx{ - height: mem.height, - gasWanted: resp.GasWanted, - tx: tx, + mem.resCbFirstTime(tx, peerID, peerP2PID, res) + + // update metrics + mem.metrics.Size.Set(float64(mem.Size())) + + // passed in by the caller of CheckTx, eg. the RPC + if externalCb != nil { + externalCb(res.GetCheckTx()) } - memTx.senders.Store(peerID, true) - mem.addTx(memTx) - mem.logger.Debug( - "added good transaction", - "tx", types.Tx(tx).Hash(), - "resp", resp, - "height", memTx.height, - "total", mem.Size(), - ) - mem.notifyTxsAvailable() - } else { - // ignore bad transaction - mem.logger.Debug( - "rejected bad transaction", - "tx", types.Tx(tx).Hash(), - "peerID", peerP2PID, - "resp", resp, - "err", postCheckErr, - ) - mem.metrics.FailedTxs.Add(1) - - if !mem.config.KeepInvalidTxsInCache { - // remove from cache (it might be good later) - mem.cache.Remove(tx) - } - } - - // update metrics - mem.metrics.Size.Set(float64(mem.Size())) - - // passed in by the caller of CheckTx, eg. the RPC - if externalCb != nil { - externalCb(resp) } } @@ -392,68 +369,136 @@ func (mem *CListMempool) isFull(txSize int) error { return nil } +// callback, which is called after the app checked the tx for the first time. +// +// The case where the app checks the tx for the second and subsequent times is +// handled by the resCbRecheck callback. +func (mem *CListMempool) resCbFirstTime( + tx []byte, + peerID uint16, + peerP2PID p2p.ID, + res *abci.Response, +) { + switch r := res.Value.(type) { + case *abci.Response_CheckTx: + var postCheckErr error + if mem.postCheck != nil { + postCheckErr = mem.postCheck(tx, r.CheckTx) + } + if (r.CheckTx.Code == abci.CodeTypeOK) && postCheckErr == nil { + // Check mempool isn't full again to reduce the chance of exceeding the + // limits. + if err := mem.isFull(len(tx)); err != nil { + // remove from cache (mempool might have a space later) + mem.cache.Remove(tx) + mem.logger.Error(err.Error()) + return + } + + memTx := &mempoolTx{ + height: mem.height, + gasWanted: r.CheckTx.GasWanted, + tx: tx, + } + memTx.senders.Store(peerID, true) + mem.addTx(memTx) + mem.logger.Debug( + "added good transaction", + "tx", types.Tx(tx).Hash(), + "res", r, + "height", memTx.height, + "total", mem.Size(), + ) + mem.notifyTxsAvailable() + } else { + // ignore bad transaction + mem.logger.Debug( + "rejected bad transaction", + "tx", types.Tx(tx).Hash(), + "peerID", peerP2PID, + "res", r, + "err", postCheckErr, + ) + mem.metrics.FailedTxs.Add(1) + + if !mem.config.KeepInvalidTxsInCache { + // remove from cache (it might be good later) + mem.cache.Remove(tx) + } + } + + default: + // ignore other messages + } +} + // callback, which is called after the app rechecked the tx. // // The case where the app checks the tx for the first time is handled by the // resCbFirstTime callback. -func (mem *CListMempool) resCbRecheck(req *abci.RequestCheckTx, res *abci.ResponseCheckTx) { - tx := req.Tx - memTx := mem.recheckCursor.Value.(*mempoolTx) +func (mem *CListMempool) resCbRecheck(req *abci.Request, res *abci.Response) { + switch r := res.Value.(type) { + case *abci.Response_CheckTx: + tx := req.GetCheckTx().Tx + memTx := mem.recheckCursor.Value.(*mempoolTx) - // Search through the remaining list of tx to recheck for a transaction that matches - // the one we received from the ABCI application. - for { - if bytes.Equal(tx, memTx.tx) { - // We've found a tx in the recheck list that matches the tx that we - // received from the ABCI application. - // Break, and use this transaction for further checks. - break + // Search through the remaining list of tx to recheck for a transaction that matches + // the one we received from the ABCI application. + for { + if bytes.Equal(tx, memTx.tx) { + // We've found a tx in the recheck list that matches the tx that we + // received from the ABCI application. + // Break, and use this transaction for further checks. + break + } + + mem.logger.Error( + "re-CheckTx transaction mismatch", + "got", types.Tx(tx), + "expected", memTx.tx, + ) + + if mem.recheckCursor == mem.recheckEnd { + // we reached the end of the recheckTx list without finding a tx + // matching the one we received from the ABCI application. + // Return without processing any tx. + mem.recheckCursor = nil + return + } + + mem.recheckCursor = mem.recheckCursor.Next() + memTx = mem.recheckCursor.Value.(*mempoolTx) } - mem.logger.Error( - "re-CheckTx transaction mismatch", - "got", types.Tx(tx), - "expected", memTx.tx, - ) + var postCheckErr error + if mem.postCheck != nil { + postCheckErr = mem.postCheck(tx, r.CheckTx) + } + if (r.CheckTx.Code == abci.CodeTypeOK) && postCheckErr == nil { + // Good, nothing to do. + } else { + // Tx became invalidated due to newly committed block. + mem.logger.Debug("tx is no longer valid", "tx", types.Tx(tx).Hash(), "res", r, "err", postCheckErr) + // NOTE: we remove tx from the cache because it might be good later + mem.removeTx(tx, mem.recheckCursor, !mem.config.KeepInvalidTxsInCache) + } if mem.recheckCursor == mem.recheckEnd { - // we reached the end of the recheckTx list without finding a tx - // matching the one we received from the ABCI application. - // Return without processing any tx. mem.recheckCursor = nil - return + } else { + mem.recheckCursor = mem.recheckCursor.Next() } + if mem.recheckCursor == nil { + // Done! + mem.logger.Debug("done rechecking txs") - mem.recheckCursor = mem.recheckCursor.Next() - memTx = mem.recheckCursor.Value.(*mempoolTx) - } - - var postCheckErr error - if mem.postCheck != nil { - postCheckErr = mem.postCheck(tx, res) - } - - if (res.Code == abci.CodeTypeOK) && postCheckErr == nil { - // Good, nothing to do. - } else { - // Tx became invalidated due to newly committed block. - mem.logger.Debug("tx is no longer valid", "tx", types.Tx(tx).Hash(), "res", res, "err", postCheckErr) - // NOTE: we remove tx from the cache because it might be good later - mem.removeTx(tx, mem.recheckCursor, !mem.config.KeepInvalidTxsInCache) - } - if mem.recheckCursor == mem.recheckEnd { - mem.recheckCursor = nil - } else { - mem.recheckCursor = mem.recheckCursor.Next() - } - if mem.recheckCursor == nil { - // Done! - mem.logger.Debug("done rechecking txs") - - // incase the recheck removed all txs - if mem.Size() > 0 { - mem.notifyTxsAvailable() + // incase the recheck removed all txs + if mem.Size() > 0 { + mem.notifyTxsAvailable() + } } + default: + // ignore other messages } } @@ -609,21 +654,16 @@ func (mem *CListMempool) recheckTxs() { // NOTE: globalCb may be called concurrently. for e := mem.txs.Front(); e != nil; e = e.Next() { memTx := e.Value.(*mempoolTx) - req := &abci.RequestCheckTx{ + _, err := mem.proxyAppConn.CheckTxAsync(context.TODO(), &abci.RequestCheckTx{ Tx: memTx.tx, Type: abci.CheckTxType_Recheck, - } - res, err := mem.proxyAppConn.CheckTx(context.TODO(), req) + }) if err != nil { - mem.logger.Error("recheckTx app error", "err", err) + mem.logger.Error("recheckTx", err, "err") + return } - - mem.resCbRecheck(req, res) } - mem.metrics.RecheckTimes.Add(1) - mem.metrics.Size.Set(float64(mem.Size())) - if err := mem.proxyAppConn.Flush(context.TODO()); err != nil { mem.logger.Error("recheckTx flush", err, "err") } diff --git a/mempool/v0/clist_mempool_test.go b/mempool/v0/clist_mempool_test.go index ca5d1d269..069534f9e 100644 --- a/mempool/v0/clist_mempool_test.go +++ b/mempool/v0/clist_mempool_test.go @@ -243,12 +243,14 @@ func TestMempoolUpdate(t *testing.T) { } func TestMempoolUpdateDoesNotPanicWhenApplicationMissedTx(t *testing.T) { + var callback abciclient.Callback mockClient := new(abciclimocks.Client) mockClient.On("Start").Return(nil) mockClient.On("SetLogger", mock.Anything) mockClient.On("Error").Return(nil).Times(4) mockClient.On("Flush", mock.Anything).Return(nil) + mockClient.On("SetResponseCallback", mock.MatchedBy(func(cb abciclient.Callback) bool { callback = cb; return true })) app := kvstore.NewInMemoryApplication() cc := proxy.NewLocalClientCreator(app) @@ -258,14 +260,16 @@ func TestMempoolUpdateDoesNotPanicWhenApplicationMissedTx(t *testing.T) { // Add 4 transactions to the mempool by calling the mempool's `CheckTx` on each of them. txs := []types.Tx{[]byte{0x01}, []byte{0x02}, []byte{0x03}, []byte{0x04}} - for idx, tx := range txs { - mockClient.On("CheckTx", mock.Anything, &abci.RequestCheckTx{Tx: tx}).Return(&abci.ResponseCheckTx{Code: abci.CodeTypeOK}, nil) - if idx != 0 { - // for all other txs we expect them to be rechecked - mockClient.On("CheckTx", mock.Anything, &abci.RequestCheckTx{Tx: tx, Type: 1}).Return(&abci.ResponseCheckTx{Code: abci.CodeTypeOK}, nil) - } + for _, tx := range txs { + reqRes := abciclient.NewReqRes(abci.ToRequestCheckTx(&abci.RequestCheckTx{Tx: tx})) + reqRes.Response = abci.ToResponseCheckTx(&abci.ResponseCheckTx{Code: abci.CodeTypeOK}) + + mockClient.On("CheckTxAsync", mock.Anything, mock.Anything).Return(reqRes, nil) err := mp.CheckTx(tx, nil, mempool.TxInfo{}) require.NoError(t, err) + + // ensure that the callback that the mempool sets on the ReqRes is run. + reqRes.InvokeCallback() } // Calling update to remove the first transaction from the mempool. @@ -273,6 +277,20 @@ func TestMempoolUpdateDoesNotPanicWhenApplicationMissedTx(t *testing.T) { err = mp.Update(0, []types.Tx{txs[0]}, abciResponses(1, abci.CodeTypeOK), nil, nil) require.Nil(t, err) + // The mempool has now sent its requests off to the client to be rechecked + // and is waiting for the corresponding callbacks to be called. + // We now call the mempool-supplied callback on the first and third transaction. + // This simulates the client dropping the second request. + // Previous versions of this code panicked when the ABCI application missed + // a recheck-tx request. + resp := &abci.ResponseCheckTx{Code: abci.CodeTypeOK} + req := &abci.RequestCheckTx{Tx: txs[1]} + callback(abci.ToRequestCheckTx(req), abci.ToResponseCheckTx(resp)) + + req = &abci.RequestCheckTx{Tx: txs[3]} + callback(abci.ToRequestCheckTx(req), abci.ToResponseCheckTx(resp)) + mockClient.AssertExpectations(t) + mockClient.AssertExpectations(t) } diff --git a/proxy/app_conn.go b/proxy/app_conn.go index 697ab21c4..7a0efef61 100644 --- a/proxy/app_conn.go +++ b/proxy/app_conn.go @@ -25,9 +25,11 @@ type AppConnConsensus interface { } type AppConnMempool interface { + SetResponseCallback(abcicli.Callback) Error() error CheckTx(context.Context, *types.RequestCheckTx) (*types.ResponseCheckTx, error) + CheckTxAsync(context.Context, *types.RequestCheckTx) (*abcicli.ReqRes, error) Flush(context.Context) error } @@ -110,6 +112,10 @@ func NewAppConnMempool(appConn abcicli.Client, metrics *Metrics) AppConnMempool } } +func (app *appConnMempool) SetResponseCallback(cb abcicli.Callback) { + app.appConn.SetResponseCallback(cb) +} + func (app *appConnMempool) Error() error { return app.appConn.Error() } @@ -124,6 +130,11 @@ func (app *appConnMempool) CheckTx(ctx context.Context, req *types.RequestCheckT return app.appConn.CheckTx(ctx, req) } +func (app *appConnMempool) CheckTxAsync(ctx context.Context, req *types.RequestCheckTx) (*abcicli.ReqRes, error) { + defer addTimeSample(app.metrics.MethodTimingSeconds.With("method", "check_tx", "type", "async"))() + return app.appConn.CheckTxAsync(ctx, req) +} + //------------------------------------------------ // Implements AppConnQuery (subset of abcicli.Client) diff --git a/proxy/mocks/app_conn_mempool.go b/proxy/mocks/app_conn_mempool.go index 65780d340..0ce564199 100644 --- a/proxy/mocks/app_conn_mempool.go +++ b/proxy/mocks/app_conn_mempool.go @@ -5,6 +5,8 @@ package mocks import ( context "context" + abcicli "github.com/tendermint/tendermint/abci/client" + mock "github.com/stretchr/testify/mock" types "github.com/tendermint/tendermint/abci/types" @@ -38,6 +40,29 @@ func (_m *AppConnMempool) CheckTx(_a0 context.Context, _a1 *types.RequestCheckTx return r0, r1 } +// CheckTxAsync provides a mock function with given fields: _a0, _a1 +func (_m *AppConnMempool) CheckTxAsync(_a0 context.Context, _a1 *types.RequestCheckTx) (*abcicli.ReqRes, error) { + ret := _m.Called(_a0, _a1) + + var r0 *abcicli.ReqRes + if rf, ok := ret.Get(0).(func(context.Context, *types.RequestCheckTx) *abcicli.ReqRes); ok { + r0 = rf(_a0, _a1) + } else { + if ret.Get(0) != nil { + r0 = ret.Get(0).(*abcicli.ReqRes) + } + } + + var r1 error + if rf, ok := ret.Get(1).(func(context.Context, *types.RequestCheckTx) error); ok { + r1 = rf(_a0, _a1) + } else { + r1 = ret.Error(1) + } + + return r0, r1 +} + // Error provides a mock function with given fields: func (_m *AppConnMempool) Error() error { ret := _m.Called()