mirror of
https://github.com/tendermint/tendermint.git
synced 2026-08-16 12:16:11 +00:00
refactoring - work in progress
This commit is contained in:
@@ -58,7 +58,7 @@ func TestResetValidator(t *testing.T) {
|
||||
// priv val after signing is not same as empty
|
||||
assert.NotEqual(t, privVal.LastSignState, emptyState)
|
||||
|
||||
// priv val after reset is same as empty
|
||||
// priv val after connect is same as empty
|
||||
privVal.Reset()
|
||||
assert.Equal(t, privVal.LastSignState, emptyState)
|
||||
}
|
||||
|
||||
@@ -21,6 +21,8 @@ func RegisterRemoteSignerMsg(cdc *amino.Codec) {
|
||||
cdc.RegisterConcrete(&PingResponse{}, "tendermint/remotesigner/PingResponse", nil)
|
||||
}
|
||||
|
||||
// TODO: Add ChainIDRequest
|
||||
|
||||
// PubKeyRequest requests the consensus public key from the remote signer.
|
||||
type PubKeyRequest struct{}
|
||||
|
||||
|
||||
+37
-86
@@ -2,20 +2,15 @@ package privval
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
|
||||
"github.com/pkg/errors"
|
||||
|
||||
"github.com/tendermint/tendermint/crypto"
|
||||
cmn "github.com/tendermint/tendermint/libs/common"
|
||||
"github.com/tendermint/tendermint/types"
|
||||
)
|
||||
|
||||
// SignerRemote implements PrivValidator.
|
||||
// It uses a net.Conn to request signatures from an external process.
|
||||
type SignerRemote struct {
|
||||
conn net.Conn
|
||||
endpoint SignerValidatorEndpoint
|
||||
|
||||
// memoized
|
||||
consensusPubKey crypto.PubKey
|
||||
@@ -25,69 +20,61 @@ type SignerRemote struct {
|
||||
var _ types.PrivValidator = (*SignerRemote)(nil)
|
||||
|
||||
// NewSignerRemote returns an instance of SignerRemote.
|
||||
func NewSignerRemote(conn net.Conn) (*SignerRemote, error) {
|
||||
func NewSignerRemote(endpoint SignerValidatorEndpoint) (*SignerRemote, error) {
|
||||
|
||||
// retrieve and memoize the consensus public key once.
|
||||
pubKey, err := getPubKey(conn)
|
||||
if err != nil {
|
||||
return nil, cmn.ErrorWrap(err, "error while retrieving public key for remote signer")
|
||||
}
|
||||
return &SignerRemote{
|
||||
conn: conn,
|
||||
consensusPubKey: pubKey,
|
||||
}, nil
|
||||
// TODO: Fix this
|
||||
//// retrieve and memoize the consensus public key once.
|
||||
//pubKey, err := getPubKey(conn)
|
||||
//if err != nil {
|
||||
// return nil, cmn.ErrorWrap(err, "error while retrieving public key for remote signer")
|
||||
//}
|
||||
// TODO: Fix this
|
||||
|
||||
//return &SignerRemote{endpoint: endpoint, consensusPubKey: pubKey,}, nil
|
||||
return &SignerRemote{endpoint: endpoint}, nil
|
||||
}
|
||||
|
||||
// Close calls Close on the underlying net.Conn.
|
||||
func (sc *SignerRemote) Close() error {
|
||||
return sc.conn.Close()
|
||||
func (sr *SignerRemote) Close() error {
|
||||
return sr.endpoint.Close()
|
||||
}
|
||||
|
||||
//--------------------------------------------------------
|
||||
// Implement PrivValidator
|
||||
|
||||
// GetPubKey implements PrivValidator.
|
||||
func (sc *SignerRemote) GetPubKey() crypto.PubKey {
|
||||
return sc.consensusPubKey
|
||||
}
|
||||
|
||||
// not thread-safe (only called on startup).
|
||||
func getPubKey(conn net.Conn) (crypto.PubKey, error) {
|
||||
err := writeMsg(conn, &PubKeyRequest{})
|
||||
func (sr *SignerRemote) GetPubKey() crypto.PubKey {
|
||||
response, err := sr.endpoint.SendRequest(&PubKeyRequest{})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
return nil
|
||||
}
|
||||
|
||||
res, err := readMsg(conn)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
pubKeyResp, ok := res.(*PubKeyResponse)
|
||||
pubKeyResp, ok := response.(*PubKeyResponse)
|
||||
if !ok {
|
||||
return nil, errors.Wrap(ErrUnexpectedResponse, "response is not PubKeyResponse")
|
||||
sr.endpoint.Logger.Error("response is not PubKeyResponse")
|
||||
return nil
|
||||
}
|
||||
|
||||
if pubKeyResp.Error != nil {
|
||||
return nil, errors.Wrap(pubKeyResp.Error, "failed to get private validator's public key")
|
||||
sr.endpoint.Logger.Error("failed to get private validator's public key", "err", pubKeyResp.Error)
|
||||
return nil
|
||||
}
|
||||
|
||||
return pubKeyResp.PubKey, nil
|
||||
return pubKeyResp.PubKey
|
||||
}
|
||||
|
||||
// SignVote implements PrivValidator.
|
||||
func (sc *SignerRemote) SignVote(chainID string, vote *types.Vote) error {
|
||||
err := writeMsg(sc.conn, &SignVoteRequest{Vote: vote})
|
||||
func (sr *SignerRemote) SignVote(chainID string, vote *types.Vote) error {
|
||||
response, err := sr.endpoint.SendRequest(&SignVoteRequest{Vote: vote})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
res, err := readMsg(sc.conn)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
resp, ok := res.(*SignedVoteResponse)
|
||||
resp, ok := response.(*SignedVoteResponse)
|
||||
if !ok {
|
||||
return ErrUnexpectedResponse
|
||||
}
|
||||
|
||||
if resp.Error != nil {
|
||||
return resp.Error
|
||||
}
|
||||
@@ -97,17 +84,13 @@ func (sc *SignerRemote) SignVote(chainID string, vote *types.Vote) error {
|
||||
}
|
||||
|
||||
// SignProposal implements PrivValidator.
|
||||
func (sc *SignerRemote) SignProposal(chainID string, proposal *types.Proposal) error {
|
||||
err := writeMsg(sc.conn, &SignProposalRequest{Proposal: proposal})
|
||||
func (sr *SignerRemote) SignProposal(chainID string, proposal *types.Proposal) error {
|
||||
response, err := sr.endpoint.SendRequest(&SignProposalRequest{Proposal: proposal})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
res, err := readMsg(sc.conn)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
resp, ok := res.(*SignedProposalResponse)
|
||||
resp, ok := response.(*SignedProposalResponse)
|
||||
if !ok {
|
||||
return ErrUnexpectedResponse
|
||||
}
|
||||
@@ -119,43 +102,11 @@ func (sc *SignerRemote) SignProposal(chainID string, proposal *types.Proposal) e
|
||||
return nil
|
||||
}
|
||||
|
||||
// Ping is used to check connection health.
|
||||
func (sc *SignerRemote) Ping() error {
|
||||
err := writeMsg(sc.conn, &PingRequest{})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
func handleRequest(
|
||||
req RemoteSignerMsg,
|
||||
chainID string,
|
||||
privVal types.PrivValidator) (RemoteSignerMsg, error) {
|
||||
|
||||
res, err := readMsg(sc.conn)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
_, ok := res.(*PingResponse)
|
||||
if !ok {
|
||||
return ErrUnexpectedResponse
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func readMsg(r io.Reader) (msg RemoteSignerMsg, err error) {
|
||||
const maxRemoteSignerMsgSize = 1024 * 10
|
||||
_, err = cdc.UnmarshalBinaryLengthPrefixedReader(r, &msg, maxRemoteSignerMsgSize)
|
||||
if _, ok := err.(timeoutError); ok {
|
||||
err = cmn.ErrorWrap(ErrConnTimeout, err.Error())
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
func writeMsg(w io.Writer, msg interface{}) (err error) {
|
||||
_, err = cdc.MarshalBinaryLengthPrefixedWriter(w, msg)
|
||||
if _, ok := err.(timeoutError); ok {
|
||||
err = cmn.ErrorWrap(ErrConnTimeout, err.Error())
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
func handleRequest(req RemoteSignerMsg, chainID string, privVal types.PrivValidator) (RemoteSignerMsg, error) {
|
||||
var res RemoteSignerMsg
|
||||
var err error
|
||||
|
||||
|
||||
@@ -1,68 +1,30 @@
|
||||
package privval
|
||||
|
||||
import (
|
||||
"net"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"github.com/tendermint/tendermint/crypto/ed25519"
|
||||
cmn "github.com/tendermint/tendermint/libs/common"
|
||||
"github.com/tendermint/tendermint/libs/log"
|
||||
"github.com/tendermint/tendermint/libs/common"
|
||||
"github.com/tendermint/tendermint/types"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// TestSignerRemoteRetryTCPOnly will test connection retry attempts over TCP. We
|
||||
// don't need this for Unix sockets because the OS instantly knows the state of
|
||||
// both ends of the socket connection. This basically causes the
|
||||
// SignerServiceEndpoint.dialer() call inside SignerServiceEndpoint.connect() to return
|
||||
// successfully immediately, putting an instant stop to any retry attempts.
|
||||
func TestSignerRemoteRetryTCPOnly(t *testing.T) {
|
||||
var (
|
||||
attemptCh = make(chan int)
|
||||
retries = 2
|
||||
)
|
||||
func TestSignerGetPubKey(t *testing.T) {
|
||||
for _, tc := range getTestCases(t) {
|
||||
func() {
|
||||
chainID := common.RandStr(12)
|
||||
mockPV := types.NewMockPV()
|
||||
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
validatorEndpoint, serviceEndpoint := getMockEndpoints(t, chainID, mockPV, tc.addr, tc.dialer)
|
||||
|
||||
go func(ln net.Listener, attemptCh chan<- int) {
|
||||
attempts := 0
|
||||
sr, err := NewSignerRemote(*validatorEndpoint)
|
||||
assert.NoError(t, err)
|
||||
|
||||
for {
|
||||
conn, err := ln.Accept()
|
||||
require.NoError(t, err)
|
||||
defer validatorEndpoint.Stop()
|
||||
defer serviceEndpoint.Stop()
|
||||
|
||||
err = conn.Close()
|
||||
require.NoError(t, err)
|
||||
clientKey := sr.GetPubKey()
|
||||
expectedPubKey := mockPV.GetPubKey()
|
||||
|
||||
attempts++
|
||||
|
||||
if attempts == retries {
|
||||
attemptCh <- attempts
|
||||
break
|
||||
}
|
||||
}
|
||||
}(ln, attemptCh)
|
||||
|
||||
serviceEndpoint := NewSignerServiceEndpoint(
|
||||
log.TestingLogger(),
|
||||
cmn.RandStr(12),
|
||||
types.NewMockPV(),
|
||||
DialTCPFn(ln.Addr().String(), testTimeoutReadWrite, ed25519.GenPrivKey()),
|
||||
)
|
||||
defer serviceEndpoint.Stop()
|
||||
|
||||
SignerServiceEndpointTimeoutReadWrite(time.Millisecond)(serviceEndpoint)
|
||||
SignerServiceEndpointConnRetries(retries)(serviceEndpoint)
|
||||
|
||||
assert.Equal(t, serviceEndpoint.Start(), ErrDialRetryMax)
|
||||
|
||||
select {
|
||||
case attempts := <-attemptCh:
|
||||
assert.Equal(t, retries, attempts)
|
||||
case <-time.After(100 * time.Millisecond):
|
||||
t.Error("expected remote to observe connection attempts")
|
||||
assert.Equal(t, expectedPubKey, clientKey)
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
@@ -60,79 +60,107 @@ func NewSignerServiceEndpoint(
|
||||
}
|
||||
|
||||
// OnStart implements cmn.Service.
|
||||
func (se *SignerServiceEndpoint) OnStart() error {
|
||||
conn, err := se.connect()
|
||||
func (ss *SignerServiceEndpoint) OnStart() error {
|
||||
conn, err := ss.connect()
|
||||
if err != nil {
|
||||
se.Logger.Error("OnStart", "err", err)
|
||||
ss.Logger.Error("OnStart", "err", err)
|
||||
return err
|
||||
}
|
||||
|
||||
se.conn = conn
|
||||
go se.handleConnection(conn)
|
||||
ss.conn = conn
|
||||
go ss.handleConnection(conn)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// OnStop implements cmn.Service.
|
||||
func (se *SignerServiceEndpoint) OnStop() {
|
||||
if se.conn == nil {
|
||||
func (ss *SignerServiceEndpoint) OnStop() {
|
||||
if ss.conn == nil {
|
||||
return
|
||||
}
|
||||
|
||||
if err := se.conn.Close(); err != nil {
|
||||
se.Logger.Error("OnStop", "err", cmn.ErrorWrap(err, "closing listener failed"))
|
||||
if err := ss.conn.Close(); err != nil {
|
||||
ss.Logger.Error("OnStop", "err", cmn.ErrorWrap(err, "closing listener failed"))
|
||||
}
|
||||
}
|
||||
|
||||
func (se *SignerServiceEndpoint) connect() (net.Conn, error) {
|
||||
for retries := 0; retries < se.connRetries; retries++ {
|
||||
func (ss *SignerServiceEndpoint) connect() (net.Conn, error) {
|
||||
for retries := 0; retries < ss.connRetries; retries++ {
|
||||
// Don't sleep if it is the first retry.
|
||||
if retries > 0 {
|
||||
time.Sleep(se.timeoutReadWrite)
|
||||
time.Sleep(ss.timeoutReadWrite)
|
||||
}
|
||||
|
||||
conn, err := se.dialer()
|
||||
conn, err := ss.dialer()
|
||||
if err == nil {
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
se.Logger.Error("dialing", "err", err)
|
||||
ss.Logger.Error("dialing", "err", err)
|
||||
}
|
||||
|
||||
return nil, ErrDialRetryMax
|
||||
}
|
||||
|
||||
func (se *SignerServiceEndpoint) handleConnection(conn net.Conn) {
|
||||
func (ss *SignerServiceEndpoint) readMessage() (msg RemoteSignerMsg, err error) {
|
||||
// TODO: Avoid duplication
|
||||
// TODO: Check connection status
|
||||
|
||||
const maxRemoteSignerMsgSize = 1024 * 10
|
||||
_, err = cdc.UnmarshalBinaryLengthPrefixedReader(ss.conn, &msg, maxRemoteSignerMsgSize)
|
||||
if _, ok := err.(timeoutError); ok {
|
||||
err = cmn.ErrorWrap(ErrConnTimeout, err.Error())
|
||||
}
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
func (ss *SignerServiceEndpoint) writeMessage(msg RemoteSignerMsg) (err error) {
|
||||
// TODO: Avoid duplication
|
||||
// TODO: Check connection status
|
||||
|
||||
_, err = cdc.MarshalBinaryLengthPrefixedWriter(ss.conn, msg)
|
||||
if _, ok := err.(timeoutError); ok {
|
||||
err = cmn.ErrorWrap(ErrConnTimeout, err.Error())
|
||||
}
|
||||
|
||||
// TODO: Probably can assert that is a response type and check for error here
|
||||
// Check the impact of KMS/Rust
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
func (ss *SignerServiceEndpoint) handleConnection(conn net.Conn) {
|
||||
for {
|
||||
if !se.IsRunning() {
|
||||
if !ss.IsRunning() {
|
||||
return // Ignore error from listener closing.
|
||||
}
|
||||
|
||||
// Reset the connection deadline
|
||||
deadline := time.Now().Add(se.timeoutReadWrite)
|
||||
deadline := time.Now().Add(ss.timeoutReadWrite)
|
||||
err := conn.SetDeadline(deadline)
|
||||
if err != nil {
|
||||
return
|
||||
}
|
||||
|
||||
req, err := readMsg(conn)
|
||||
req, err := ss.readMessage()
|
||||
if err != nil {
|
||||
if err != io.EOF {
|
||||
se.Logger.Error("handleConnection readMsg", "err", err)
|
||||
ss.Logger.Error("handleConnection readMessage", "err", err)
|
||||
}
|
||||
return
|
||||
}
|
||||
|
||||
res, err := handleRequest(req, se.chainID, se.privVal)
|
||||
res, err := handleRequest(req, ss.chainID, ss.privVal)
|
||||
|
||||
if err != nil {
|
||||
// only log the error; we'll reply with an error in res
|
||||
se.Logger.Error("handleConnection handleRequest", "err", err)
|
||||
ss.Logger.Error("handleConnection handleRequest", "err", err)
|
||||
}
|
||||
|
||||
err = writeMsg(conn, res)
|
||||
err = ss.writeMessage(res)
|
||||
if err != nil {
|
||||
se.Logger.Error("handleConnection writeMsg", "err", err)
|
||||
ss.Logger.Error("handleConnection writeMessage", "err", err)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
@@ -6,10 +6,8 @@ import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/tendermint/tendermint/crypto"
|
||||
cmn "github.com/tendermint/tendermint/libs/common"
|
||||
"github.com/tendermint/tendermint/libs/log"
|
||||
"github.com/tendermint/tendermint/types"
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -30,33 +28,31 @@ func SignerValidatorEndpointSetHeartbeat(period time.Duration) SignerValidatorEn
|
||||
return func(sc *SignerValidatorEndpoint) { sc.heartbeatPeriod = period }
|
||||
}
|
||||
|
||||
|
||||
// TODO: Add a type for SignerEndpoints
|
||||
// getConnection
|
||||
// connect
|
||||
// read
|
||||
// write
|
||||
// close
|
||||
|
||||
// TODO: Fix comments
|
||||
// SocketVal implements PrivValidator.
|
||||
// It listens for an external process to dial in and uses
|
||||
// the socket to request signatures.
|
||||
type SignerValidatorEndpoint struct {
|
||||
cmn.BaseService
|
||||
|
||||
mtx sync.Mutex
|
||||
listener net.Listener
|
||||
conn net.Conn
|
||||
|
||||
// ping
|
||||
cancelPingCh chan struct{}
|
||||
pingTicker *time.Ticker
|
||||
heartbeatPeriod time.Duration
|
||||
|
||||
// signer is mutable since it can be reset if the connection fails.
|
||||
// failures are detected by a background ping routine.
|
||||
// All messages are request/response, so we hold the mutex
|
||||
// so only one request/response pair can happen at a time.
|
||||
// Methods on the underlying net.Conn itself are already goroutine safe.
|
||||
mtx sync.Mutex
|
||||
|
||||
// TODO: Signer should encapsulate and hide the endpoint completely. Invert the relation
|
||||
signer *SignerRemote
|
||||
}
|
||||
|
||||
// Check that SignerValidatorEndpoint implements PrivValidator.
|
||||
var _ types.PrivValidator = (*SignerValidatorEndpoint)(nil)
|
||||
|
||||
// NewSignerValidatorEndpoint returns an instance of SignerValidatorEndpoint.
|
||||
func NewSignerValidatorEndpoint(logger log.Logger, listener net.Listener) *SignerValidatorEndpoint {
|
||||
sc := &SignerValidatorEndpoint{
|
||||
@@ -69,84 +65,37 @@ func NewSignerValidatorEndpoint(logger log.Logger, listener net.Listener) *Signe
|
||||
return sc
|
||||
}
|
||||
|
||||
//--------------------------------------------------------
|
||||
// Implement PrivValidator
|
||||
|
||||
// GetPubKey implements PrivValidator.
|
||||
func (ve *SignerValidatorEndpoint) GetPubKey() crypto.PubKey {
|
||||
ve.mtx.Lock()
|
||||
defer ve.mtx.Unlock()
|
||||
return ve.signer.GetPubKey()
|
||||
}
|
||||
|
||||
// SignVote implements PrivValidator.
|
||||
func (ve *SignerValidatorEndpoint) SignVote(chainID string, vote *types.Vote) error {
|
||||
ve.mtx.Lock()
|
||||
defer ve.mtx.Unlock()
|
||||
return ve.signer.SignVote(chainID, vote)
|
||||
}
|
||||
|
||||
// SignProposal implements PrivValidator.
|
||||
func (ve *SignerValidatorEndpoint) SignProposal(chainID string, proposal *types.Proposal) error {
|
||||
ve.mtx.Lock()
|
||||
defer ve.mtx.Unlock()
|
||||
return ve.signer.SignProposal(chainID, proposal)
|
||||
}
|
||||
|
||||
//--------------------------------------------------------
|
||||
// More thread safe methods proxied to the signer
|
||||
|
||||
// Ping is used to check connection health.
|
||||
func (ve *SignerValidatorEndpoint) Ping() error {
|
||||
ve.mtx.Lock()
|
||||
defer ve.mtx.Unlock()
|
||||
return ve.signer.Ping()
|
||||
}
|
||||
|
||||
// Close closes the underlying net.Conn.
|
||||
func (ve *SignerValidatorEndpoint) Close() {
|
||||
ve.mtx.Lock()
|
||||
defer ve.mtx.Unlock()
|
||||
if ve.signer != nil {
|
||||
if err := ve.signer.Close(); err != nil {
|
||||
ve.Logger.Error("OnStop", "err", err)
|
||||
}
|
||||
}
|
||||
|
||||
if ve.listener != nil {
|
||||
if err := ve.listener.Close(); err != nil {
|
||||
ve.Logger.Error("OnStop", "err", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
//--------------------------------------------------------
|
||||
// Service start and stop
|
||||
|
||||
// OnStart implements cmn.Service.
|
||||
func (ve *SignerValidatorEndpoint) OnStart() error {
|
||||
if closed, err := ve.reset(); err != nil {
|
||||
closed, err := ve.connect()
|
||||
// TODO: Improve. Connection state should be kept in a variable
|
||||
|
||||
if err != nil {
|
||||
ve.Logger.Error("OnStart", "err", err)
|
||||
return err
|
||||
} else if closed {
|
||||
}
|
||||
|
||||
if closed {
|
||||
return fmt.Errorf("listener is closed")
|
||||
}
|
||||
|
||||
// Start a routine to keep the connection alive
|
||||
ve.cancelPingCh = make(chan struct{}, 1)
|
||||
ve.pingTicker = time.NewTicker(ve.heartbeatPeriod)
|
||||
|
||||
// TODO: Move subroutine to another place?
|
||||
go func() {
|
||||
for {
|
||||
select {
|
||||
case <-ve.pingTicker.C:
|
||||
err := ve.Ping()
|
||||
err := ve.ping()
|
||||
if err != nil {
|
||||
ve.Logger.Error("Ping", "err", err)
|
||||
if err == ErrUnexpectedResponse {
|
||||
return
|
||||
}
|
||||
|
||||
closed, err := ve.reset()
|
||||
closed, err := ve.connect()
|
||||
if err != nil {
|
||||
ve.Logger.Error("Reconnecting to remote signer failed", "err", err)
|
||||
continue
|
||||
@@ -173,42 +122,118 @@ func (ve *SignerValidatorEndpoint) OnStop() {
|
||||
if ve.cancelPingCh != nil {
|
||||
close(ve.cancelPingCh)
|
||||
}
|
||||
ve.Close()
|
||||
_ = ve.Close()
|
||||
}
|
||||
|
||||
//--------------------------------------------------------
|
||||
// Connection and signer management
|
||||
// Close closes the underlying net.Conn.
|
||||
func (ve *SignerValidatorEndpoint) Close() error {
|
||||
ve.mtx.Lock()
|
||||
defer ve.mtx.Unlock()
|
||||
|
||||
if ve.conn != nil {
|
||||
if err := ve.conn.Close(); err != nil {
|
||||
ve.Logger.Error("Closing connection", "err", err)
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
if ve.listener != nil {
|
||||
if err := ve.listener.Close(); err != nil {
|
||||
ve.Logger.Error("Closing Listener", "err", err)
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// SendRequest sends a request and waits for a response
|
||||
func (ve *SignerValidatorEndpoint) SendRequest(request RemoteSignerMsg) (RemoteSignerMsg, error) {
|
||||
ve.mtx.Lock()
|
||||
defer ve.mtx.Unlock()
|
||||
|
||||
err := ve.writeMessage(request)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
res, err := ve.readMessage()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return res, nil
|
||||
}
|
||||
|
||||
// Ping is used to check connection health.
|
||||
func (ve *SignerValidatorEndpoint) ping() error {
|
||||
response, err := ve.SendRequest(&PingRequest{})
|
||||
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
_, ok := response.(*PingResponse)
|
||||
if !ok {
|
||||
return ErrUnexpectedResponse
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (ve *SignerValidatorEndpoint) readMessage() (msg RemoteSignerMsg, err error) {
|
||||
// TODO: Check connection status
|
||||
|
||||
const maxRemoteSignerMsgSize = 1024 * 10
|
||||
_, err = cdc.UnmarshalBinaryLengthPrefixedReader(ve.conn, &msg, maxRemoteSignerMsgSize)
|
||||
if _, ok := err.(timeoutError); ok {
|
||||
err = cmn.ErrorWrap(ErrConnTimeout, err.Error())
|
||||
}
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
func (ve *SignerValidatorEndpoint) writeMessage(msg RemoteSignerMsg) (err error) {
|
||||
// TODO: Check connection status
|
||||
|
||||
_, err = cdc.MarshalBinaryLengthPrefixedWriter(ve.conn, msg)
|
||||
if _, ok := err.(timeoutError); ok {
|
||||
err = cmn.ErrorWrap(ErrConnTimeout, err.Error())
|
||||
}
|
||||
|
||||
return
|
||||
}
|
||||
|
||||
// waits to accept and sets a new connection.
|
||||
// connection is closed in OnStop.
|
||||
// returns true if the listener is closed
|
||||
// (ie. it returns a nil conn).
|
||||
func (ve *SignerValidatorEndpoint) reset() (closed bool, err error) {
|
||||
// returns true if the listener is closed (ie. it returns a nil conn).
|
||||
// TODO: Improve this
|
||||
func (ve *SignerValidatorEndpoint) connect() (closed bool, err error) {
|
||||
ve.mtx.Lock()
|
||||
defer ve.mtx.Unlock()
|
||||
|
||||
// first check if the conn already exists and close it.
|
||||
if ve.signer != nil {
|
||||
if tmpErr := ve.signer.Close(); tmpErr != nil {
|
||||
ve.Logger.Error("error closing socket val connection during reset", "err", tmpErr)
|
||||
if ve.conn != nil {
|
||||
if tmpErr := ve.conn.Close(); tmpErr != nil {
|
||||
ve.Logger.Error("error closing socket val connection during connect", "err", tmpErr)
|
||||
}
|
||||
}
|
||||
|
||||
// wait for a new conn
|
||||
conn, err := ve.acceptConnection()
|
||||
ve.conn, err = ve.acceptConnection()
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
// listener is closed
|
||||
if conn == nil {
|
||||
if ve.conn == nil {
|
||||
return true, nil
|
||||
}
|
||||
|
||||
ve.signer, err = NewSignerRemote(conn)
|
||||
if err != nil {
|
||||
// TODO: This does not belong here... but maybe we need to inform the owner that a connection has been received
|
||||
// failed to fetch the pubkey. close out the connection.
|
||||
if tmpErr := conn.Close(); tmpErr != nil {
|
||||
if tmpErr := ve.conn.Close(); tmpErr != nil {
|
||||
ve.Logger.Error("error closing connection", "err", tmpErr)
|
||||
}
|
||||
return false, err
|
||||
@@ -216,8 +241,9 @@ func (ve *SignerValidatorEndpoint) reset() (closed bool, err error) {
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// Attempt to accept a connection.
|
||||
// Times out after the listener's timeoutAccept
|
||||
// acceptConnection attempts to accept a connection
|
||||
// it will timeout after the listener's timeoutAccept
|
||||
// TODO: There is no reason for this separate accept
|
||||
func (ve *SignerValidatorEndpoint) acceptConnection() (net.Conn, error) {
|
||||
conn, err := ve.listener.Accept()
|
||||
if err != nil {
|
||||
|
||||
@@ -8,11 +8,9 @@ import (
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/tendermint/tendermint/crypto/ed25519"
|
||||
cmn "github.com/tendermint/tendermint/libs/common"
|
||||
"github.com/tendermint/tendermint/libs/common"
|
||||
"github.com/tendermint/tendermint/libs/log"
|
||||
|
||||
"github.com/tendermint/tendermint/types"
|
||||
)
|
||||
|
||||
@@ -31,11 +29,414 @@ type socketTestCase struct {
|
||||
dialer SocketDialer
|
||||
}
|
||||
|
||||
func socketTestCases(t *testing.T) []socketTestCase {
|
||||
// TestSignerRemoteRetryTCPOnly will test connection retry attempts over TCP. We
|
||||
// don't need this for Unix sockets because the OS instantly knows the state of
|
||||
// both ends of the socket connection. This basically causes the
|
||||
// SignerServiceEndpoint.dialer() call inside SignerServiceEndpoint.connect() to return
|
||||
// successfully immediately, putting an instant stop to any retry attempts.
|
||||
func TestSignerRemoteRetryTCPOnly(t *testing.T) {
|
||||
var (
|
||||
attemptCh = make(chan int)
|
||||
retries = 2
|
||||
)
|
||||
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
|
||||
go func(ln net.Listener, attemptCh chan<- int) {
|
||||
attempts := 0
|
||||
|
||||
for {
|
||||
conn, err := ln.Accept()
|
||||
require.NoError(t, err)
|
||||
|
||||
err = conn.Close()
|
||||
require.NoError(t, err)
|
||||
|
||||
attempts++
|
||||
|
||||
if attempts == retries {
|
||||
attemptCh <- attempts
|
||||
break
|
||||
}
|
||||
}
|
||||
}(ln, attemptCh)
|
||||
|
||||
serviceEndpoint := NewSignerServiceEndpoint(
|
||||
log.TestingLogger(),
|
||||
common.RandStr(12),
|
||||
types.NewMockPV(),
|
||||
DialTCPFn(ln.Addr().String(), testTimeoutReadWrite, ed25519.GenPrivKey()),
|
||||
)
|
||||
defer serviceEndpoint.Stop()
|
||||
|
||||
SignerServiceEndpointTimeoutReadWrite(time.Millisecond)(serviceEndpoint)
|
||||
SignerServiceEndpointConnRetries(retries)(serviceEndpoint)
|
||||
|
||||
assert.Equal(t, serviceEndpoint.Start(), ErrDialRetryMax)
|
||||
|
||||
select {
|
||||
case attempts := <-attemptCh:
|
||||
assert.Equal(t, retries, attempts)
|
||||
case <-time.After(100 * time.Millisecond):
|
||||
t.Error("expected remote to observe connection attempts")
|
||||
}
|
||||
}
|
||||
|
||||
//func TestSocketPVAddress(t *testing.T) {
|
||||
// for _, tc := range getTestCases(t) {
|
||||
// // Execute the test within a closure to ensure the deferred statements
|
||||
// // are called between each for loop iteration, for isolated test cases.
|
||||
// func() {
|
||||
// var (
|
||||
// chainID = common.RandStr(12)
|
||||
// validatorEndpoint, serviceEndpoint = getMockEndpoints(t, chainID, types.NewMockPV(), tc.addr, tc.dialer)
|
||||
// )
|
||||
// defer validatorEndpoint.Stop()
|
||||
// defer serviceEndpoint.Stop()
|
||||
//
|
||||
// serviceAddr := serviceEndpoint.privVal.GetPubKey().Address()
|
||||
// validatorAddr := validatorEndpoint.GetPubKey().Address()
|
||||
//
|
||||
// assert.Equal(t, serviceAddr, validatorAddr)
|
||||
// }()
|
||||
// }
|
||||
//}
|
||||
//
|
||||
//
|
||||
//func TestSocketPVProposal(t *testing.T) {
|
||||
// for _, tc := range getTestCases(t) {
|
||||
// func() {
|
||||
// var (
|
||||
// chainID = common.RandStr(12)
|
||||
// validatorEndpoint, serviceEndpoint = getMockEndpoints(
|
||||
// t,
|
||||
// chainID,
|
||||
// types.NewMockPV(),
|
||||
// tc.addr,
|
||||
// tc.dialer)
|
||||
//
|
||||
// ts = time.Now()
|
||||
// privProposal = &types.Proposal{Timestamp: ts}
|
||||
// clientProposal = &types.Proposal{Timestamp: ts}
|
||||
// )
|
||||
// defer validatorEndpoint.Stop()
|
||||
// defer serviceEndpoint.Stop()
|
||||
//
|
||||
// require.NoError(t, serviceEndpoint.privVal.SignProposal(chainID, privProposal))
|
||||
// require.NoError(t, validatorEndpoint.SignProposal(chainID, clientProposal))
|
||||
//
|
||||
// assert.Equal(t, privProposal.Signature, clientProposal.Signature)
|
||||
// }()
|
||||
// }
|
||||
//}
|
||||
//
|
||||
//func TestSocketPVVote(t *testing.T) {
|
||||
// for _, tc := range getTestCases(t) {
|
||||
// func() {
|
||||
// var (
|
||||
// chainID = common.RandStr(12)
|
||||
// validatorEndpoint, serviceEndpoint = getMockEndpoints(
|
||||
// t,
|
||||
// chainID,
|
||||
// types.NewMockPV(),
|
||||
// tc.addr,
|
||||
// tc.dialer)
|
||||
//
|
||||
// ts = time.Now()
|
||||
// vType = types.PrecommitType
|
||||
// want = &types.Vote{Timestamp: ts, Type: vType}
|
||||
// have = &types.Vote{Timestamp: ts, Type: vType}
|
||||
// )
|
||||
// defer validatorEndpoint.Stop()
|
||||
// defer serviceEndpoint.Stop()
|
||||
//
|
||||
// require.NoError(t, serviceEndpoint.privVal.SignVote(chainID, want))
|
||||
// require.NoError(t, validatorEndpoint.SignVote(chainID, have))
|
||||
// assert.Equal(t, want.Signature, have.Signature)
|
||||
// }()
|
||||
// }
|
||||
//}
|
||||
//
|
||||
//func TestSocketPVVoteResetDeadline(t *testing.T) {
|
||||
// for _, tc := range getTestCases(t) {
|
||||
// func() {
|
||||
// var (
|
||||
// chainID = common.RandStr(12)
|
||||
// validatorEndpoint, serviceEndpoint = getMockEndpoints(
|
||||
// t,
|
||||
// chainID,
|
||||
// types.NewMockPV(),
|
||||
// tc.addr,
|
||||
// tc.dialer)
|
||||
//
|
||||
// ts = time.Now()
|
||||
// vType = types.PrecommitType
|
||||
// want = &types.Vote{Timestamp: ts, Type: vType}
|
||||
// have = &types.Vote{Timestamp: ts, Type: vType}
|
||||
// )
|
||||
// defer validatorEndpoint.Stop()
|
||||
// defer serviceEndpoint.Stop()
|
||||
//
|
||||
// time.Sleep(testTimeoutReadWrite2o3)
|
||||
//
|
||||
// require.NoError(t, serviceEndpoint.privVal.SignVote(chainID, want))
|
||||
// require.NoError(t, validatorEndpoint.SignVote(chainID, have))
|
||||
// assert.Equal(t, want.Signature, have.Signature)
|
||||
//
|
||||
// // This would exceed the deadline if it was not extended by the previous message
|
||||
// time.Sleep(testTimeoutReadWrite2o3)
|
||||
//
|
||||
// require.NoError(t, serviceEndpoint.privVal.SignVote(chainID, want))
|
||||
// require.NoError(t, validatorEndpoint.SignVote(chainID, have))
|
||||
// assert.Equal(t, want.Signature, have.Signature)
|
||||
// }()
|
||||
// }
|
||||
//}
|
||||
//
|
||||
//func TestSocketPVVoteKeepalive(t *testing.T) {
|
||||
// for _, tc := range getTestCases(t) {
|
||||
// func() {
|
||||
// var (
|
||||
// chainID = common.RandStr(12)
|
||||
// validatorEndpoint, serviceEndpoint = getMockEndpoints(
|
||||
// t,
|
||||
// chainID,
|
||||
// types.NewMockPV(),
|
||||
// tc.addr,
|
||||
// tc.dialer)
|
||||
//
|
||||
// ts = time.Now()
|
||||
// vType = types.PrecommitType
|
||||
// want = &types.Vote{Timestamp: ts, Type: vType}
|
||||
// have = &types.Vote{Timestamp: ts, Type: vType}
|
||||
// )
|
||||
// defer validatorEndpoint.Stop()
|
||||
// defer serviceEndpoint.Stop()
|
||||
//
|
||||
// time.Sleep(testTimeoutReadWrite * 2)
|
||||
//
|
||||
// require.NoError(t, serviceEndpoint.privVal.SignVote(chainID, want))
|
||||
// require.NoError(t, validatorEndpoint.SignVote(chainID, have))
|
||||
// assert.Equal(t, want.Signature, have.Signature)
|
||||
// }()
|
||||
// }
|
||||
//}
|
||||
//
|
||||
//func TestSocketPVDeadline(t *testing.T) {
|
||||
// for _, tc := range getTestCases(t) {
|
||||
// func() {
|
||||
// var (
|
||||
// listenc = make(chan struct{})
|
||||
// thisConnTimeout = 100 * time.Millisecond
|
||||
// validatorEndpoint = newSignerValidatorEndpoint(log.TestingLogger(), tc.addr, thisConnTimeout)
|
||||
// )
|
||||
//
|
||||
// go func(sc *SignerValidatorEndpoint) {
|
||||
// defer close(listenc)
|
||||
//
|
||||
// // Note: the TCP connection times out at the accept() phase,
|
||||
// // whereas the Unix domain sockets connection times out while
|
||||
// // attempting to fetch the remote signer's public key.
|
||||
// assert.True(t, IsConnTimeout(sc.Start()))
|
||||
//
|
||||
// assert.False(t, sc.IsRunning())
|
||||
// }(validatorEndpoint)
|
||||
//
|
||||
// for {
|
||||
// _, err := common.Connect(tc.addr)
|
||||
// if err == nil {
|
||||
// break
|
||||
// }
|
||||
// }
|
||||
//
|
||||
// <-listenc
|
||||
// }()
|
||||
// }
|
||||
//}
|
||||
//
|
||||
//func TestRemoteSignVoteErrors(t *testing.T) {
|
||||
// for _, tc := range getTestCases(t) {
|
||||
// func() {
|
||||
// var (
|
||||
// chainID = common.RandStr(12)
|
||||
// validatorEndpoint, serviceEndpoint = getMockEndpoints(
|
||||
// t,
|
||||
// chainID,
|
||||
// types.NewErroringMockPV(),
|
||||
// tc.addr,
|
||||
// tc.dialer)
|
||||
//
|
||||
// ts = time.Now()
|
||||
// vType = types.PrecommitType
|
||||
// vote = &types.Vote{Timestamp: ts, Type: vType}
|
||||
// )
|
||||
// defer validatorEndpoint.Stop()
|
||||
// defer serviceEndpoint.Stop()
|
||||
//
|
||||
// err := validatorEndpoint.SignVote("", vote)
|
||||
// require.Equal(t, err.(*RemoteSignerError).Description, types.ErroringMockPVErr.Error())
|
||||
//
|
||||
// err = serviceEndpoint.privVal.SignVote(chainID, vote)
|
||||
// require.Error(t, err)
|
||||
// err = validatorEndpoint.SignVote(chainID, vote)
|
||||
// require.Error(t, err)
|
||||
// }()
|
||||
// }
|
||||
//}
|
||||
//
|
||||
//func TestRemoteSignProposalErrors(t *testing.T) {
|
||||
// for _, tc := range getTestCases(t) {
|
||||
// func() {
|
||||
// var (
|
||||
// chainID = common.RandStr(12)
|
||||
// validatorEndpoint, serviceEndpoint = getMockEndpoints(
|
||||
// t,
|
||||
// chainID,
|
||||
// types.NewErroringMockPV(),
|
||||
// tc.addr,
|
||||
// tc.dialer)
|
||||
//
|
||||
// ts = time.Now()
|
||||
// proposal = &types.Proposal{Timestamp: ts}
|
||||
// )
|
||||
// defer validatorEndpoint.Stop()
|
||||
// defer serviceEndpoint.Stop()
|
||||
//
|
||||
// err := validatorEndpoint.SignProposal("", proposal)
|
||||
// require.Equal(t, err.(*RemoteSignerError).Description, types.ErroringMockPVErr.Error())
|
||||
//
|
||||
// err = serviceEndpoint.privVal.SignProposal(chainID, proposal)
|
||||
// require.Error(t, err)
|
||||
//
|
||||
// err = validatorEndpoint.SignProposal(chainID, proposal)
|
||||
// require.Error(t, err)
|
||||
// }()
|
||||
// }
|
||||
//}
|
||||
//
|
||||
//func TestErrUnexpectedResponse(t *testing.T) {
|
||||
// for _, tc := range getTestCases(t) {
|
||||
// func() {
|
||||
// var (
|
||||
// logger = log.TestingLogger()
|
||||
// chainID = common.RandStr(12)
|
||||
// readyCh = make(chan struct{})
|
||||
// errCh = make(chan error, 1)
|
||||
//
|
||||
// serviceEndpoint = NewSignerServiceEndpoint(
|
||||
// logger,
|
||||
// chainID,
|
||||
// types.NewMockPV(),
|
||||
// tc.dialer,
|
||||
// )
|
||||
//
|
||||
// validatorEndpoint = newSignerValidatorEndpoint(
|
||||
// logger,
|
||||
// tc.addr,
|
||||
// testTimeoutReadWrite)
|
||||
// )
|
||||
//
|
||||
// getStartEndpoint(t, readyCh, validatorEndpoint)
|
||||
// defer validatorEndpoint.Stop()
|
||||
// SignerServiceEndpointTimeoutReadWrite(time.Millisecond)(serviceEndpoint)
|
||||
// SignerServiceEndpointConnRetries(100)(serviceEndpoint)
|
||||
// // we do not want to Start() the remote signer here and instead use the connection to
|
||||
// // reply with intentionally wrong replies below:
|
||||
// rsConn, err := serviceEndpoint.connect()
|
||||
// defer rsConn.Close()
|
||||
// require.NoError(t, err)
|
||||
// require.NotNil(t, rsConn)
|
||||
// // send over public key to get the remote signer running:
|
||||
// go testReadWriteResponse(t, &PubKeyResponse{}, rsConn)
|
||||
// <-readyCh
|
||||
//
|
||||
// // Proposal:
|
||||
// go func(errc chan error) {
|
||||
// errc <- validatorEndpoint.SignProposal(chainID, &types.Proposal{})
|
||||
// }(errCh)
|
||||
//
|
||||
// // read request and write wrong response:
|
||||
// go testReadWriteResponse(t, &SignedVoteResponse{}, rsConn)
|
||||
// err = <-errCh
|
||||
// require.Error(t, err)
|
||||
// require.Equal(t, err, ErrUnexpectedResponse)
|
||||
//
|
||||
// // Vote:
|
||||
// go func(errc chan error) {
|
||||
// errc <- validatorEndpoint.SignVote(chainID, &types.Vote{})
|
||||
// }(errCh)
|
||||
// // read request and write wrong response:
|
||||
// go testReadWriteResponse(t, &SignedProposalResponse{}, rsConn)
|
||||
// err = <-errCh
|
||||
// require.Error(t, err)
|
||||
// require.Equal(t, err, ErrUnexpectedResponse)
|
||||
// }()
|
||||
// }
|
||||
//}
|
||||
//
|
||||
//func TestRetryConnToRemoteSigner(t *testing.T) {
|
||||
// for _, tc := range getTestCases(t) {
|
||||
// func() {
|
||||
// var (
|
||||
// logger = log.TestingLogger()
|
||||
// chainID = common.RandStr(12)
|
||||
// readyCh = make(chan struct{})
|
||||
//
|
||||
// serviceEndpoint = NewSignerServiceEndpoint(
|
||||
// logger,
|
||||
// chainID,
|
||||
// types.NewMockPV(),
|
||||
// tc.dialer,
|
||||
// )
|
||||
// thisConnTimeout = testTimeoutReadWrite
|
||||
// validatorEndpoint = newSignerValidatorEndpoint(logger, tc.addr, thisConnTimeout)
|
||||
// )
|
||||
// // Ping every:
|
||||
// SignerValidatorEndpointSetHeartbeat(testTimeoutHeartbeat)(validatorEndpoint)
|
||||
//
|
||||
// SignerServiceEndpointTimeoutReadWrite(testTimeoutReadWrite)(serviceEndpoint)
|
||||
// SignerServiceEndpointConnRetries(10)(serviceEndpoint)
|
||||
//
|
||||
// getStartEndpoint(t, readyCh, validatorEndpoint)
|
||||
// defer validatorEndpoint.Stop()
|
||||
// require.NoError(t, serviceEndpoint.Start())
|
||||
// assert.True(t, serviceEndpoint.IsRunning())
|
||||
//
|
||||
// <-readyCh
|
||||
// time.Sleep(testTimeoutHeartbeat * 2)
|
||||
//
|
||||
// serviceEndpoint.Stop()
|
||||
// rs2 := NewSignerServiceEndpoint(
|
||||
// logger,
|
||||
// chainID,
|
||||
// types.NewMockPV(),
|
||||
// tc.dialer,
|
||||
// )
|
||||
// // let some pings pass
|
||||
// time.Sleep(testTimeoutHeartbeat3o2)
|
||||
// require.NoError(t, rs2.Start())
|
||||
// assert.True(t, rs2.IsRunning())
|
||||
// defer rs2.Stop()
|
||||
//
|
||||
// // give the client some time to re-establish the conn to the remote signer
|
||||
// // should see sth like this in the logs:
|
||||
// //
|
||||
// // E[10016-01-10|17:12:46.128] Ping err="remote signer timed out"
|
||||
// // I[10016-01-10|17:16:42.447] Re-created connection to remote signer impl=SocketVal
|
||||
// time.Sleep(testTimeoutReadWrite * 2)
|
||||
// }()
|
||||
// }
|
||||
//}
|
||||
|
||||
///////////////////////////////////
|
||||
|
||||
func getTestCases(t *testing.T) []socketTestCase {
|
||||
tcpAddr := fmt.Sprintf("tcp://%s", testFreeTCPAddr(t))
|
||||
unixFilePath, err := testUnixAddr()
|
||||
require.NoError(t, err)
|
||||
unixAddr := fmt.Sprintf("unix://%s", unixFilePath)
|
||||
|
||||
return []socketTestCase{
|
||||
{
|
||||
addr: tcpAddr,
|
||||
@@ -48,376 +449,8 @@ func socketTestCases(t *testing.T) []socketTestCase {
|
||||
}
|
||||
}
|
||||
|
||||
func TestSocketPVAddress(t *testing.T) {
|
||||
for _, tc := range socketTestCases(t) {
|
||||
// Execute the test within a closure to ensure the deferred statements
|
||||
// are called between each for loop iteration, for isolated test cases.
|
||||
func() {
|
||||
var (
|
||||
chainID = cmn.RandStr(12)
|
||||
validatorEndpoint, serviceEndpoint = testSetupSocketPair(t, chainID, types.NewMockPV(), tc.addr, tc.dialer)
|
||||
)
|
||||
defer validatorEndpoint.Stop()
|
||||
defer serviceEndpoint.Stop()
|
||||
|
||||
serviceAddr := serviceEndpoint.privVal.GetPubKey().Address()
|
||||
validatorAddr := validatorEndpoint.GetPubKey().Address()
|
||||
|
||||
assert.Equal(t, serviceAddr, validatorAddr)
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
func TestSocketPVPubKey(t *testing.T) {
|
||||
for _, tc := range socketTestCases(t) {
|
||||
func() {
|
||||
var (
|
||||
chainID = cmn.RandStr(12)
|
||||
validatorEndpoint, serviceEndpoint = testSetupSocketPair(
|
||||
t,
|
||||
chainID,
|
||||
types.NewMockPV(),
|
||||
tc.addr,
|
||||
tc.dialer)
|
||||
)
|
||||
defer validatorEndpoint.Stop()
|
||||
defer serviceEndpoint.Stop()
|
||||
|
||||
clientKey := validatorEndpoint.GetPubKey()
|
||||
privvalPubKey := serviceEndpoint.privVal.GetPubKey()
|
||||
|
||||
assert.Equal(t, privvalPubKey, clientKey)
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
func TestSocketPVProposal(t *testing.T) {
|
||||
for _, tc := range socketTestCases(t) {
|
||||
func() {
|
||||
var (
|
||||
chainID = cmn.RandStr(12)
|
||||
validatorEndpoint, serviceEndpoint = testSetupSocketPair(
|
||||
t,
|
||||
chainID,
|
||||
types.NewMockPV(),
|
||||
tc.addr,
|
||||
tc.dialer)
|
||||
|
||||
ts = time.Now()
|
||||
privProposal = &types.Proposal{Timestamp: ts}
|
||||
clientProposal = &types.Proposal{Timestamp: ts}
|
||||
)
|
||||
defer validatorEndpoint.Stop()
|
||||
defer serviceEndpoint.Stop()
|
||||
|
||||
require.NoError(t, serviceEndpoint.privVal.SignProposal(chainID, privProposal))
|
||||
require.NoError(t, validatorEndpoint.SignProposal(chainID, clientProposal))
|
||||
|
||||
assert.Equal(t, privProposal.Signature, clientProposal.Signature)
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
func TestSocketPVVote(t *testing.T) {
|
||||
for _, tc := range socketTestCases(t) {
|
||||
func() {
|
||||
var (
|
||||
chainID = cmn.RandStr(12)
|
||||
validatorEndpoint, serviceEndpoint = testSetupSocketPair(
|
||||
t,
|
||||
chainID,
|
||||
types.NewMockPV(),
|
||||
tc.addr,
|
||||
tc.dialer)
|
||||
|
||||
ts = time.Now()
|
||||
vType = types.PrecommitType
|
||||
want = &types.Vote{Timestamp: ts, Type: vType}
|
||||
have = &types.Vote{Timestamp: ts, Type: vType}
|
||||
)
|
||||
defer validatorEndpoint.Stop()
|
||||
defer serviceEndpoint.Stop()
|
||||
|
||||
require.NoError(t, serviceEndpoint.privVal.SignVote(chainID, want))
|
||||
require.NoError(t, validatorEndpoint.SignVote(chainID, have))
|
||||
assert.Equal(t, want.Signature, have.Signature)
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
func TestSocketPVVoteResetDeadline(t *testing.T) {
|
||||
for _, tc := range socketTestCases(t) {
|
||||
func() {
|
||||
var (
|
||||
chainID = cmn.RandStr(12)
|
||||
validatorEndpoint, serviceEndpoint = testSetupSocketPair(
|
||||
t,
|
||||
chainID,
|
||||
types.NewMockPV(),
|
||||
tc.addr,
|
||||
tc.dialer)
|
||||
|
||||
ts = time.Now()
|
||||
vType = types.PrecommitType
|
||||
want = &types.Vote{Timestamp: ts, Type: vType}
|
||||
have = &types.Vote{Timestamp: ts, Type: vType}
|
||||
)
|
||||
defer validatorEndpoint.Stop()
|
||||
defer serviceEndpoint.Stop()
|
||||
|
||||
time.Sleep(testTimeoutReadWrite2o3)
|
||||
|
||||
require.NoError(t, serviceEndpoint.privVal.SignVote(chainID, want))
|
||||
require.NoError(t, validatorEndpoint.SignVote(chainID, have))
|
||||
assert.Equal(t, want.Signature, have.Signature)
|
||||
|
||||
// This would exceed the deadline if it was not extended by the previous message
|
||||
time.Sleep(testTimeoutReadWrite2o3)
|
||||
|
||||
require.NoError(t, serviceEndpoint.privVal.SignVote(chainID, want))
|
||||
require.NoError(t, validatorEndpoint.SignVote(chainID, have))
|
||||
assert.Equal(t, want.Signature, have.Signature)
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
func TestSocketPVVoteKeepalive(t *testing.T) {
|
||||
for _, tc := range socketTestCases(t) {
|
||||
func() {
|
||||
var (
|
||||
chainID = cmn.RandStr(12)
|
||||
validatorEndpoint, serviceEndpoint = testSetupSocketPair(
|
||||
t,
|
||||
chainID,
|
||||
types.NewMockPV(),
|
||||
tc.addr,
|
||||
tc.dialer)
|
||||
|
||||
ts = time.Now()
|
||||
vType = types.PrecommitType
|
||||
want = &types.Vote{Timestamp: ts, Type: vType}
|
||||
have = &types.Vote{Timestamp: ts, Type: vType}
|
||||
)
|
||||
defer validatorEndpoint.Stop()
|
||||
defer serviceEndpoint.Stop()
|
||||
|
||||
time.Sleep(testTimeoutReadWrite * 2)
|
||||
|
||||
require.NoError(t, serviceEndpoint.privVal.SignVote(chainID, want))
|
||||
require.NoError(t, validatorEndpoint.SignVote(chainID, have))
|
||||
assert.Equal(t, want.Signature, have.Signature)
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
func TestSocketPVDeadline(t *testing.T) {
|
||||
for _, tc := range socketTestCases(t) {
|
||||
func() {
|
||||
var (
|
||||
listenc = make(chan struct{})
|
||||
thisConnTimeout = 100 * time.Millisecond
|
||||
validatorEndpoint = newSignerValidatorEndpoint(log.TestingLogger(), tc.addr, thisConnTimeout)
|
||||
)
|
||||
|
||||
go func(sc *SignerValidatorEndpoint) {
|
||||
defer close(listenc)
|
||||
|
||||
// Note: the TCP connection times out at the accept() phase,
|
||||
// whereas the Unix domain sockets connection times out while
|
||||
// attempting to fetch the remote signer's public key.
|
||||
assert.True(t, IsConnTimeout(sc.Start()))
|
||||
|
||||
assert.False(t, sc.IsRunning())
|
||||
}(validatorEndpoint)
|
||||
|
||||
for {
|
||||
_, err := cmn.Connect(tc.addr)
|
||||
if err == nil {
|
||||
break
|
||||
}
|
||||
}
|
||||
|
||||
<-listenc
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoteSignVoteErrors(t *testing.T) {
|
||||
for _, tc := range socketTestCases(t) {
|
||||
func() {
|
||||
var (
|
||||
chainID = cmn.RandStr(12)
|
||||
validatorEndpoint, serviceEndpoint = testSetupSocketPair(
|
||||
t,
|
||||
chainID,
|
||||
types.NewErroringMockPV(),
|
||||
tc.addr,
|
||||
tc.dialer)
|
||||
|
||||
ts = time.Now()
|
||||
vType = types.PrecommitType
|
||||
vote = &types.Vote{Timestamp: ts, Type: vType}
|
||||
)
|
||||
defer validatorEndpoint.Stop()
|
||||
defer serviceEndpoint.Stop()
|
||||
|
||||
err := validatorEndpoint.SignVote("", vote)
|
||||
require.Equal(t, err.(*RemoteSignerError).Description, types.ErroringMockPVErr.Error())
|
||||
|
||||
err = serviceEndpoint.privVal.SignVote(chainID, vote)
|
||||
require.Error(t, err)
|
||||
err = validatorEndpoint.SignVote(chainID, vote)
|
||||
require.Error(t, err)
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoteSignProposalErrors(t *testing.T) {
|
||||
for _, tc := range socketTestCases(t) {
|
||||
func() {
|
||||
var (
|
||||
chainID = cmn.RandStr(12)
|
||||
validatorEndpoint, serviceEndpoint = testSetupSocketPair(
|
||||
t,
|
||||
chainID,
|
||||
types.NewErroringMockPV(),
|
||||
tc.addr,
|
||||
tc.dialer)
|
||||
|
||||
ts = time.Now()
|
||||
proposal = &types.Proposal{Timestamp: ts}
|
||||
)
|
||||
defer validatorEndpoint.Stop()
|
||||
defer serviceEndpoint.Stop()
|
||||
|
||||
err := validatorEndpoint.SignProposal("", proposal)
|
||||
require.Equal(t, err.(*RemoteSignerError).Description, types.ErroringMockPVErr.Error())
|
||||
|
||||
err = serviceEndpoint.privVal.SignProposal(chainID, proposal)
|
||||
require.Error(t, err)
|
||||
|
||||
err = validatorEndpoint.SignProposal(chainID, proposal)
|
||||
require.Error(t, err)
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
func TestErrUnexpectedResponse(t *testing.T) {
|
||||
for _, tc := range socketTestCases(t) {
|
||||
func() {
|
||||
var (
|
||||
logger = log.TestingLogger()
|
||||
chainID = cmn.RandStr(12)
|
||||
readyCh = make(chan struct{})
|
||||
errCh = make(chan error, 1)
|
||||
|
||||
serviceEndpoint = NewSignerServiceEndpoint(
|
||||
logger,
|
||||
chainID,
|
||||
types.NewMockPV(),
|
||||
tc.dialer,
|
||||
)
|
||||
|
||||
validatorEndpoint = newSignerValidatorEndpoint(
|
||||
logger,
|
||||
tc.addr,
|
||||
testTimeoutReadWrite)
|
||||
)
|
||||
|
||||
testStartEndpoint(t, readyCh, validatorEndpoint)
|
||||
defer validatorEndpoint.Stop()
|
||||
SignerServiceEndpointTimeoutReadWrite(time.Millisecond)(serviceEndpoint)
|
||||
SignerServiceEndpointConnRetries(100)(serviceEndpoint)
|
||||
// we do not want to Start() the remote signer here and instead use the connection to
|
||||
// reply with intentionally wrong replies below:
|
||||
rsConn, err := serviceEndpoint.connect()
|
||||
defer rsConn.Close()
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, rsConn)
|
||||
// send over public key to get the remote signer running:
|
||||
go testReadWriteResponse(t, &PubKeyResponse{}, rsConn)
|
||||
<-readyCh
|
||||
|
||||
// Proposal:
|
||||
go func(errc chan error) {
|
||||
errc <- validatorEndpoint.SignProposal(chainID, &types.Proposal{})
|
||||
}(errCh)
|
||||
|
||||
// read request and write wrong response:
|
||||
go testReadWriteResponse(t, &SignedVoteResponse{}, rsConn)
|
||||
err = <-errCh
|
||||
require.Error(t, err)
|
||||
require.Equal(t, err, ErrUnexpectedResponse)
|
||||
|
||||
// Vote:
|
||||
go func(errc chan error) {
|
||||
errc <- validatorEndpoint.SignVote(chainID, &types.Vote{})
|
||||
}(errCh)
|
||||
// read request and write wrong response:
|
||||
go testReadWriteResponse(t, &SignedProposalResponse{}, rsConn)
|
||||
err = <-errCh
|
||||
require.Error(t, err)
|
||||
require.Equal(t, err, ErrUnexpectedResponse)
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
func TestRetryConnToRemoteSigner(t *testing.T) {
|
||||
for _, tc := range socketTestCases(t) {
|
||||
func() {
|
||||
var (
|
||||
logger = log.TestingLogger()
|
||||
chainID = cmn.RandStr(12)
|
||||
readyCh = make(chan struct{})
|
||||
|
||||
serviceEndpoint = NewSignerServiceEndpoint(
|
||||
logger,
|
||||
chainID,
|
||||
types.NewMockPV(),
|
||||
tc.dialer,
|
||||
)
|
||||
thisConnTimeout = testTimeoutReadWrite
|
||||
validatorEndpoint = newSignerValidatorEndpoint(logger, tc.addr, thisConnTimeout)
|
||||
)
|
||||
// Ping every:
|
||||
SignerValidatorEndpointSetHeartbeat(testTimeoutHeartbeat)(validatorEndpoint)
|
||||
|
||||
SignerServiceEndpointTimeoutReadWrite(testTimeoutReadWrite)(serviceEndpoint)
|
||||
SignerServiceEndpointConnRetries(10)(serviceEndpoint)
|
||||
|
||||
testStartEndpoint(t, readyCh, validatorEndpoint)
|
||||
defer validatorEndpoint.Stop()
|
||||
require.NoError(t, serviceEndpoint.Start())
|
||||
assert.True(t, serviceEndpoint.IsRunning())
|
||||
|
||||
<-readyCh
|
||||
time.Sleep(testTimeoutHeartbeat * 2)
|
||||
|
||||
serviceEndpoint.Stop()
|
||||
rs2 := NewSignerServiceEndpoint(
|
||||
logger,
|
||||
chainID,
|
||||
types.NewMockPV(),
|
||||
tc.dialer,
|
||||
)
|
||||
// let some pings pass
|
||||
time.Sleep(testTimeoutHeartbeat3o2)
|
||||
require.NoError(t, rs2.Start())
|
||||
assert.True(t, rs2.IsRunning())
|
||||
defer rs2.Stop()
|
||||
|
||||
// give the client some time to re-establish the conn to the remote signer
|
||||
// should see sth like this in the logs:
|
||||
//
|
||||
// E[10016-01-10|17:12:46.128] Ping err="remote signer timed out"
|
||||
// I[10016-01-10|17:16:42.447] Re-created connection to remote signer impl=SocketVal
|
||||
time.Sleep(testTimeoutReadWrite * 2)
|
||||
}()
|
||||
}
|
||||
}
|
||||
|
||||
func newSignerValidatorEndpoint(logger log.Logger, addr string, timeoutReadWrite time.Duration) *SignerValidatorEndpoint {
|
||||
proto, address := cmn.ProtocolAndAddress(addr)
|
||||
proto, address := common.ProtocolAndAddress(addr)
|
||||
|
||||
ln, err := net.Listen(proto, address)
|
||||
logger.Info("Listening at", "proto", proto, "address", address)
|
||||
@@ -442,17 +475,26 @@ func newSignerValidatorEndpoint(logger log.Logger, addr string, timeoutReadWrite
|
||||
return NewSignerValidatorEndpoint(logger, listener)
|
||||
}
|
||||
|
||||
func testSetupSocketPair(
|
||||
func getStartEndpoint(t *testing.T, readyCh chan struct{}, sv *SignerValidatorEndpoint) {
|
||||
go func(sv *SignerValidatorEndpoint) {
|
||||
require.NoError(t, sv.Start())
|
||||
assert.True(t, sv.IsRunning())
|
||||
readyCh <- struct{}{}
|
||||
}(sv)
|
||||
}
|
||||
|
||||
func getMockEndpoints(
|
||||
t *testing.T,
|
||||
chainID string,
|
||||
privValidator types.PrivValidator,
|
||||
addr string,
|
||||
socketDialer SocketDialer,
|
||||
) (*SignerValidatorEndpoint, *SignerServiceEndpoint) {
|
||||
|
||||
var (
|
||||
logger = log.TestingLogger()
|
||||
privVal = privValidator
|
||||
readyc = make(chan struct{})
|
||||
readyCh = make(chan struct{})
|
||||
serviceEndpoint = NewSignerServiceEndpoint(
|
||||
logger,
|
||||
chainID,
|
||||
@@ -460,46 +502,27 @@ func testSetupSocketPair(
|
||||
socketDialer,
|
||||
)
|
||||
|
||||
thisConnTimeout = testTimeoutReadWrite
|
||||
validatorEndpoint = newSignerValidatorEndpoint(logger, addr, thisConnTimeout)
|
||||
validatorEndpoint = newSignerValidatorEndpoint(logger, addr, testTimeoutReadWrite)
|
||||
)
|
||||
|
||||
SignerValidatorEndpointSetHeartbeat(testTimeoutHeartbeat)(validatorEndpoint)
|
||||
SignerServiceEndpointTimeoutReadWrite(testTimeoutReadWrite)(serviceEndpoint)
|
||||
SignerServiceEndpointConnRetries(1e6)(serviceEndpoint)
|
||||
|
||||
testStartEndpoint(t, readyc, validatorEndpoint)
|
||||
getStartEndpoint(t, readyCh, validatorEndpoint)
|
||||
|
||||
require.NoError(t, serviceEndpoint.Start())
|
||||
assert.True(t, serviceEndpoint.IsRunning())
|
||||
|
||||
<-readyc
|
||||
<-readyCh
|
||||
|
||||
return validatorEndpoint, serviceEndpoint
|
||||
}
|
||||
|
||||
func testReadWriteResponse(t *testing.T, resp RemoteSignerMsg, rsConn net.Conn) {
|
||||
_, err := readMsg(rsConn)
|
||||
require.NoError(t, err)
|
||||
|
||||
err = writeMsg(rsConn, resp)
|
||||
require.NoError(t, err)
|
||||
}
|
||||
|
||||
func testStartEndpoint(t *testing.T, readyCh chan struct{}, sc *SignerValidatorEndpoint) {
|
||||
go func(sc *SignerValidatorEndpoint) {
|
||||
require.NoError(t, sc.Start())
|
||||
assert.True(t, sc.IsRunning())
|
||||
|
||||
readyCh <- struct{}{}
|
||||
}(sc)
|
||||
}
|
||||
|
||||
// testFreeTCPAddr claims a free port so we don't block on listener being ready.
|
||||
func testFreeTCPAddr(t *testing.T) string {
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
defer ln.Close()
|
||||
|
||||
return fmt.Sprintf("127.0.0.1:%d", ln.Addr().(*net.TCPAddr).Port)
|
||||
}
|
||||
//func testReadWriteResponse(t *testing.T, resp RemoteSignerMsg, rsConn net.Conn) {
|
||||
// _, err := readMessage(rsConn)
|
||||
// require.NoError(t, err)
|
||||
//
|
||||
// err = writeMessage(rsConn, resp)
|
||||
// require.NoError(t, err)
|
||||
//}
|
||||
|
||||
@@ -2,6 +2,8 @@ package privval
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"github.com/stretchr/testify/require"
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
@@ -12,3 +14,12 @@ func TestIsConnTimeoutForNonTimeoutErrors(t *testing.T) {
|
||||
assert.False(t, IsConnTimeout(cmn.ErrorWrap(ErrDialRetryMax, "max retries exceeded")))
|
||||
assert.False(t, IsConnTimeout(fmt.Errorf("completely irrelevant error")))
|
||||
}
|
||||
|
||||
// testFreeTCPAddr claims a free port so we don't block on listener being ready.
|
||||
func testFreeTCPAddr(t *testing.T) string {
|
||||
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
||||
require.NoError(t, err)
|
||||
defer ln.Close()
|
||||
|
||||
return fmt.Sprintf("127.0.0.1:%d", ln.Addr().(*net.TCPAddr).Port)
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user