diff --git a/api/handlers/mock/signing.go b/api/handlers/mock/signing.go index b3796e39..8afbba5d 100644 --- a/api/handlers/mock/signing.go +++ b/api/handlers/mock/signing.go @@ -51,3 +51,39 @@ func (mr *MockSignatureCacherMockRecorder) Subscribe(ctx, id, sigChannel any) *g mr.mock.ctrl.T.Helper() return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Subscribe", reflect.TypeOf((*MockSignatureCacher)(nil).Subscribe), ctx, id, sigChannel) } + +// MockSignatureRemover is a mock of SignatureRemover interface. +type MockSignatureRemover struct { + ctrl *gomock.Controller + recorder *MockSignatureRemoverMockRecorder + isgomock struct{} +} + +// MockSignatureRemoverMockRecorder is the mock recorder for MockSignatureRemover. +type MockSignatureRemoverMockRecorder struct { + mock *MockSignatureRemover +} + +// NewMockSignatureRemover creates a new mock instance. +func NewMockSignatureRemover(ctrl *gomock.Controller) *MockSignatureRemover { + mock := &MockSignatureRemover{ctrl: ctrl} + mock.recorder = &MockSignatureRemoverMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockSignatureRemover) EXPECT() *MockSignatureRemoverMockRecorder { + return m.recorder +} + +// Remove mocks base method. +func (m *MockSignatureRemover) Remove(id string) { + m.ctrl.T.Helper() + m.ctrl.Call(m, "Remove", id) +} + +// Remove indicates an expected call of Remove. +func (mr *MockSignatureRemoverMockRecorder) Remove(id any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Remove", reflect.TypeOf((*MockSignatureRemover)(nil).Remove), id) +} diff --git a/api/handlers/signing.go b/api/handlers/signing.go index ba617f47..ce7da387 100644 --- a/api/handlers/signing.go +++ b/api/handlers/signing.go @@ -12,6 +12,7 @@ import ( "github.com/gorilla/mux" evmMessage "github.com/sprintertech/sprinter-signing/chains/evm/message" lighterMessage "github.com/sprintertech/sprinter-signing/chains/lighter/message" + "github.com/sprintertech/sprinter-signing/tss/ecdsa/signing" "github.com/sygmaprotocol/sygma-core/relayer/message" ) @@ -41,14 +42,16 @@ type SigningBody struct { } type SigningHandler struct { - msgChan chan []*message.Message - chains map[uint64]struct{} + msgChan chan []*message.Message + chains map[uint64]struct{} + sigCache SignatureRemover } -func NewSigningHandler(msgChan chan []*message.Message, chains map[uint64]struct{}) *SigningHandler { +func NewSigningHandler(msgChan chan []*message.Message, chains map[uint64]struct{}, sigCache SignatureRemover) *SigningHandler { return &SigningHandler{ - msgChan: msgChan, - chains: chains, + msgChan: msgChan, + chains: chains, + sigCache: sigCache, } } @@ -142,6 +145,7 @@ func (h *SigningHandler) HandleSigning(w http.ResponseWriter, r *http.Request) { JSONError(w, fmt.Errorf("invalid protocol %s", b.Protocol), http.StatusBadRequest) return } + h.sigCache.Remove(signing.SessionID(b.ChainId, b.DepositId)) h.msgChan <- []*message.Message{m} err = <-errChn @@ -196,6 +200,10 @@ type SignatureCacher interface { Subscribe(ctx context.Context, id string, sigChannel chan []byte) } +type SignatureRemover interface { + Remove(id string) +} + type StatusHandler struct { cache SignatureCacher chains map[uint64]struct{} @@ -234,7 +242,7 @@ func (h *StatusHandler) HandleRequest(w http.ResponseWriter, r *http.Request) { ctx := r.Context() sigChn := make(chan []byte, 1) - h.cache.Subscribe(ctx, fmt.Sprintf("%d-%s", chainId, depositId), sigChn) + h.cache.Subscribe(ctx, signing.SessionID(chainId.Uint64(), depositId), sigChn) for { select { case <-r.Context().Done(): diff --git a/api/handlers/signing_test.go b/api/handlers/signing_test.go index 820a46a9..0e4dceb6 100644 --- a/api/handlers/signing_test.go +++ b/api/handlers/signing_test.go @@ -16,6 +16,7 @@ import ( "github.com/sprintertech/sprinter-signing/api/handlers" mock_handlers "github.com/sprintertech/sprinter-signing/api/handlers/mock" across "github.com/sprintertech/sprinter-signing/chains/evm/message" + lighterChain "github.com/sprintertech/sprinter-signing/chains/lighter" lighter "github.com/sprintertech/sprinter-signing/chains/lighter/message" "github.com/stretchr/testify/suite" "github.com/sygmaprotocol/sygma-core/relayer/message" @@ -25,7 +26,8 @@ import ( type SigningHandlerTestSuite struct { suite.Suite - chains map[uint64]struct{} + chains map[uint64]struct{} + mockRemover *mock_handlers.MockSignatureRemover } func TestRunSigningHandlerTestSuite(t *testing.T) { @@ -33,14 +35,17 @@ func TestRunSigningHandlerTestSuite(t *testing.T) { } func (s *SigningHandlerTestSuite) SetupTest() { + ctrl := gomock.NewController(s.T()) chains := make(map[uint64]struct{}) chains[1] = struct{}{} + chains[lighterChain.LIGHTER_DOMAIN_ID] = struct{}{} s.chains = chains + s.mockRemover = mock_handlers.NewMockSignatureRemover(ctrl) } func (s *SigningHandlerTestSuite) Test_HandleSigning_MissingDepositID() { msgChn := make(chan []*message.Message) - handler := handlers.NewSigningHandler(msgChn, s.chains) + handler := handlers.NewSigningHandler(msgChn, s.chains, s.mockRemover) input := handlers.SigningBody{ Protocol: "across", @@ -71,7 +76,7 @@ func (s *SigningHandlerTestSuite) Test_HandleSigning_MissingDepositID() { func (s *SigningHandlerTestSuite) Test_HandleSigning_MissingCaller() { msgChn := make(chan []*message.Message) - handler := handlers.NewSigningHandler(msgChn, s.chains) + handler := handlers.NewSigningHandler(msgChn, s.chains, s.mockRemover) input := handlers.SigningBody{ Protocol: "across", @@ -102,7 +107,7 @@ func (s *SigningHandlerTestSuite) Test_HandleSigning_MissingCaller() { func (s *SigningHandlerTestSuite) Test_HandleSigning_MissingLiquidityPool() { msgChn := make(chan []*message.Message) - handler := handlers.NewSigningHandler(msgChn, s.chains) + handler := handlers.NewSigningHandler(msgChn, s.chains, s.mockRemover) input := handlers.SigningBody{ Protocol: "across", @@ -133,7 +138,7 @@ func (s *SigningHandlerTestSuite) Test_HandleSigning_MissingLiquidityPool() { func (s *SigningHandlerTestSuite) Test_HandleSigning_InvalidChainID() { msgChn := make(chan []*message.Message) - handler := handlers.NewSigningHandler(msgChn, s.chains) + handler := handlers.NewSigningHandler(msgChn, s.chains, s.mockRemover) input := handlers.SigningBody{ DepositId: "1000", @@ -165,7 +170,7 @@ func (s *SigningHandlerTestSuite) Test_HandleSigning_InvalidChainID() { func (s *SigningHandlerTestSuite) Test_HandleSigning_ChainNotSupported() { msgChn := make(chan []*message.Message) - handler := handlers.NewSigningHandler(msgChn, s.chains) + handler := handlers.NewSigningHandler(msgChn, s.chains, s.mockRemover) input := handlers.SigningBody{ DepositId: "1000", @@ -197,7 +202,7 @@ func (s *SigningHandlerTestSuite) Test_HandleSigning_ChainNotSupported() { func (s *SigningHandlerTestSuite) Test_HandleSigning_InvalidProtocol() { msgChn := make(chan []*message.Message) - handler := handlers.NewSigningHandler(msgChn, s.chains) + handler := handlers.NewSigningHandler(msgChn, s.chains, s.mockRemover) input := handlers.SigningBody{ DepositId: "1000", @@ -229,7 +234,7 @@ func (s *SigningHandlerTestSuite) Test_HandleSigning_InvalidProtocol() { func (s *SigningHandlerTestSuite) Test_HandleSigning_ErrorHandlingMessage() { msgChn := make(chan []*message.Message) - handler := handlers.NewSigningHandler(msgChn, s.chains) + handler := handlers.NewSigningHandler(msgChn, s.chains, s.mockRemover) input := handlers.SigningBody{ DepositId: "1000", @@ -257,6 +262,8 @@ func (s *SigningHandlerTestSuite) Test_HandleSigning_ErrorHandlingMessage() { ad.ErrChn <- fmt.Errorf("error handling message") }() + s.mockRemover.EXPECT().Remove("1-1000") + handler.HandleSigning(recorder, req) s.Equal(http.StatusInternalServerError, recorder.Code) @@ -264,7 +271,7 @@ func (s *SigningHandlerTestSuite) Test_HandleSigning_ErrorHandlingMessage() { func (s *SigningHandlerTestSuite) Test_HandleSigning_AcrossSuccess() { msgChn := make(chan []*message.Message) - handler := handlers.NewSigningHandler(msgChn, s.chains) + handler := handlers.NewSigningHandler(msgChn, s.chains, s.mockRemover) input := handlers.SigningBody{ DepositId: "1000", @@ -294,6 +301,8 @@ func (s *SigningHandlerTestSuite) Test_HandleSigning_AcrossSuccess() { ad.ErrChn <- nil }() + s.mockRemover.EXPECT().Remove("1-1000") + handler.HandleSigning(recorder, req) s.Equal(http.StatusAccepted, recorder.Code) @@ -301,7 +310,7 @@ func (s *SigningHandlerTestSuite) Test_HandleSigning_AcrossSuccess() { func (s *SigningHandlerTestSuite) Test_HandleSigning_LifiSuccess() { msgChn := make(chan []*message.Message) - handler := handlers.NewSigningHandler(msgChn, s.chains) + handler := handlers.NewSigningHandler(msgChn, s.chains, s.mockRemover) input := handlers.SigningBody{ DepositId: "depositID", @@ -330,6 +339,8 @@ func (s *SigningHandlerTestSuite) Test_HandleSigning_LifiSuccess() { ad.ErrChn <- nil }() + s.mockRemover.EXPECT().Remove("1-depositID") + handler.HandleSigning(recorder, req) s.Equal(http.StatusAccepted, recorder.Code) @@ -337,7 +348,7 @@ func (s *SigningHandlerTestSuite) Test_HandleSigning_LifiSuccess() { func (s *SigningHandlerTestSuite) Test_HandleSigning_LighterSuccess() { msgChn := make(chan []*message.Message) - handler := handlers.NewSigningHandler(msgChn, s.chains) + handler := handlers.NewSigningHandler(msgChn, s.chains, s.mockRemover) input := handlers.SigningBody{ DepositId: "depositID", @@ -354,7 +365,7 @@ func (s *SigningHandlerTestSuite) Test_HandleSigning_LighterSuccess() { req := httptest.NewRequest(http.MethodPost, "/v1/chains/1/signatures", bytes.NewReader(body)) req = mux.SetURLVars(req, map[string]string{ - "chainId": "1", + "chainId": fmt.Sprintf("%d", lighterChain.LIGHTER_DOMAIN_ID), }) req.Header.Set("Content-Type", "application/json") @@ -366,6 +377,8 @@ func (s *SigningHandlerTestSuite) Test_HandleSigning_LighterSuccess() { ad.ErrChn <- nil }() + s.mockRemover.EXPECT().Remove(fmt.Sprintf("%d-depositID", lighterChain.LIGHTER_DOMAIN_ID)) + handler.HandleSigning(recorder, req) s.Equal(http.StatusAccepted, recorder.Code) @@ -373,7 +386,7 @@ func (s *SigningHandlerTestSuite) Test_HandleSigning_LighterSuccess() { func (s *SigningHandlerTestSuite) Test_HandleSigning_SprinterSuccess() { msgChn := make(chan []*message.Message) - handler := handlers.NewSigningHandler(msgChn, s.chains) + handler := handlers.NewSigningHandler(msgChn, s.chains, s.mockRemover) input := handlers.SigningBody{ DepositId: "depositID", @@ -402,6 +415,8 @@ func (s *SigningHandlerTestSuite) Test_HandleSigning_SprinterSuccess() { ad.ErrChn <- nil }() + s.mockRemover.EXPECT().Remove("1-depositID") + handler.HandleSigning(recorder, req) s.Equal(http.StatusAccepted, recorder.Code) diff --git a/app/app.go b/app/app.go index f4e77dd6..4d6c4c38 100644 --- a/app/app.go +++ b/app/app.go @@ -435,7 +435,7 @@ func Run() error { log.Info().Msg("Relayer not part of MPC. Waiting for refresh event...") } - signingHandler := handlers.NewSigningHandler(msgChan, supportedChains) + signingHandler := handlers.NewSigningHandler(msgChan, supportedChains, signatureCache) statusHandler := handlers.NewStatusHandler(signatureCache, supportedChains) confirmationsHandler := handlers.NewConfirmationsHandler(confirmationsPerChain) unlockHandler := handlers.NewUnlockHandler(msgChan, supportedChains) diff --git a/cache/signature.go b/cache/signature.go index 04b07495..8f272083 100644 --- a/cache/signature.go +++ b/cache/signature.go @@ -68,6 +68,10 @@ func (s *SignatureCache) Subscribe(ctx context.Context, id string, sigChannel ch } } +func (s *SignatureCache) Remove(id string) { + s.sigCache.Delete(id) +} + func (s *SignatureCache) Signature(id string) ([]byte, error) { sig := s.sigCache.Get(id) if sig == nil { diff --git a/cache/signature_test.go b/cache/signature_test.go index 8f43c163..09cd768b 100644 --- a/cache/signature_test.go +++ b/cache/signature_test.go @@ -123,6 +123,24 @@ func (s *SignatureCacheTestSuite) Test_Subscribe_ValidMessage_EarlyExit() { s.Equal(sig, expectedSig.Signature) } +func (s *SignatureCacheTestSuite) Test_Remove_DeletesCachedSignature() { + expectedSig := signing.EcdsaSignature{ + Signature: []byte("signature"), + ID: "signatureID", + } + s.mockMetrics.EXPECT().EndProcess(expectedSig.ID) + s.sigChn <- expectedSig + time.Sleep(time.Millisecond * 100) + + _, err := s.sc.Signature(expectedSig.ID) + s.Nil(err) + + s.sc.Remove(expectedSig.ID) + + _, err = s.sc.Signature(expectedSig.ID) + s.NotNil(err) +} + func (s *SignatureCacheTestSuite) Test_Subscribe_ValidMessage() { expectedSig := signing.EcdsaSignature{ Signature: []byte("signature"), diff --git a/chains/evm/message/across.go b/chains/evm/message/across.go index 806818eb..3632b364 100644 --- a/chains/evm/message/across.go +++ b/chains/evm/message/across.go @@ -197,7 +197,7 @@ func (h *AcrossMessageHandler) HandleMessage(m *message.Message) (*proposal.Prop return nil, err } - sessionID := fmt.Sprintf("%d-%s", sourceChainID, data.DepositId) + sessionID := signing.SessionID(sourceChainID, data.DepositId.String()) signing, err := signing.NewSigning( new(big.Int).SetBytes(unlockHash), sessionID, diff --git a/chains/evm/message/lifiEscrow.go b/chains/evm/message/lifiEscrow.go index 3a6e6956..2100429f 100644 --- a/chains/evm/message/lifiEscrow.go +++ b/chains/evm/message/lifiEscrow.go @@ -180,7 +180,7 @@ func (h *LifiEscrowMessageHandler) HandleMessage(m *message.Message) (*proposal. return nil, err } - sessionID := fmt.Sprintf("%d-%s", h.chainID, data.OrderID) + sessionID := signing.SessionID(h.chainID, data.OrderID) signing, err := signing.NewSigning( new(big.Int).SetBytes(unlockHash), sessionID, diff --git a/chains/evm/message/sprinter.go b/chains/evm/message/sprinter.go index d71bdc76..1bae1b7f 100644 --- a/chains/evm/message/sprinter.go +++ b/chains/evm/message/sprinter.go @@ -94,7 +94,7 @@ func (h *SprinterCreditMessageHandler) HandleMessage(m *message.Message) (*propo } data.ErrChn <- nil - sessionID := fmt.Sprintf("%d-%s", h.chainID, data.DepositID) + sessionID := signing.SessionID(h.chainID, data.DepositID) signing, err := signing.NewSigning( new(big.Int).SetBytes(unlockHash), sessionID, diff --git a/chains/lighter/message/lighter.go b/chains/lighter/message/lighter.go index 12db97a2..c75d7d8f 100644 --- a/chains/lighter/message/lighter.go +++ b/chains/lighter/message/lighter.go @@ -124,7 +124,7 @@ func (h *LighterMessageHandler) HandleMessage(m *message.Message) (*proposal.Pro return nil, err } - sessionID := fmt.Sprintf("%d-%s", lighterChain.LIGHTER_DOMAIN_ID, data.OrderHash) + sessionID := signing.SessionID(lighterChain.LIGHTER_DOMAIN_ID, data.OrderHash) signing, err := signing.NewSigning( new(big.Int).SetBytes(unlockHash), sessionID, diff --git a/tss/ecdsa/signing/signing.go b/tss/ecdsa/signing/signing.go index ea7fe6e3..b9c86739 100644 --- a/tss/ecdsa/signing/signing.go +++ b/tss/ecdsa/signing/signing.go @@ -50,6 +50,10 @@ type Signing struct { subscriptionID comm.SubscriptionID } +func SessionID(chainID uint64, depositID string) string { + return fmt.Sprintf("%d-%s", chainID, depositID) +} + func NewSigning( msg *big.Int, messageID string,