From 36d4b97ee58030db682b7050b5813df0320b786d Mon Sep 17 00:00:00 2001 From: AntiD2ta Date: Wed, 7 Jan 2026 11:22:32 +0100 Subject: [PATCH 1/6] When listing accounts, query all available endpoints and deduplicate Co-authored-by: Jacob Shufro --- account.go | 13 +- distributedaccount.go | 1 + grpc.go | 227 +++++++++++++++------- grpc_internal_test.go | 44 ++++- mock/listerserver.go | 411 ++++++++++++++++++++++++---------------- wallet_internal_test.go | 4 +- 6 files changed, 456 insertions(+), 244 deletions(-) diff --git a/account.go b/account.go index 20f9ff9..3d7d9b7 100644 --- a/account.go +++ b/account.go @@ -24,12 +24,13 @@ import ( ) type account struct { - wallet *wallet - id uuid.UUID - name string - pubKey e2types.PublicKey - version uint - mutex *sync.RWMutex + wallet *wallet + id uuid.UUID + name string + pubKey e2types.PublicKey + version uint + mutex *sync.RWMutex + endpoint *Endpoint } func newAccount(wallet *wallet, diff --git a/distributedaccount.go b/distributedaccount.go index 0a4e7b6..8ca64b7 100644 --- a/distributedaccount.go +++ b/distributedaccount.go @@ -36,6 +36,7 @@ type distributedAccount struct { participantConns map[uint64]*grpc.ClientConn version uint mutex *sync.RWMutex + endpoint *Endpoint } func newDistributedAccount(wallet *wallet, diff --git a/grpc.go b/grpc.go index e4fede2..8abcb7c 100644 --- a/grpc.go +++ b/grpc.go @@ -18,6 +18,7 @@ import ( "crypto/tls" "crypto/x509" "encoding/binary" + "encoding/hex" "fmt" "os" "runtime" @@ -102,74 +103,162 @@ func (w *wallet) List(ctx context.Context, accountPath string) ([]e2wtypes.Accou return nil, errors.New("wallet has no endpoints") } - var resp *pb.ListAccountsResponse - var err error - ctx, cancelFunc := context.WithTimeout(ctx, w.timeout) - defer cancelFunc() + var wg sync.WaitGroup + var errsWaitGroup sync.WaitGroup + errs := make([]error, 0) + errChan := make(chan error) + + // add a tick to the wait group to ensure errs is populated after + // wg.Wait() and errsWaitGroup.Wait() are called. + errsWaitGroup.Add(1) + go func(errChan chan error) { + defer errsWaitGroup.Done() + for err := range errChan { + errs = append(errs, err) + } + }(errChan) + + respsMap := new(sync.Map) for i := range len(w.endpoints) { - var conn *grpc.ClientConn - var release func() - conn, release, err = w.connectionProvider.Connection(ctx, w.endpoints[i]) - if err != nil { - w.log.Debug().Stringer("endpoint", w.endpoints[i]).Str("path", path).Err(err).Msg("Failed to obtain connection") + wg.Add(1) + go func() { + var resp *pb.ListAccountsResponse + var err error + + i := i + defer wg.Done() + + ctx, cancelFunc := context.WithTimeout(ctx, w.timeout) + defer cancelFunc() + + var conn *grpc.ClientConn + var release func() + conn, release, err = w.connectionProvider.Connection(ctx, w.endpoints[i]) + if err != nil { + w.log.Debug().Stringer("endpoint", w.endpoints[i]).Str("path", path).Err(err).Msg("Failed to obtain connection") + errChan <- err + return + } + + listerClient := pb.NewListerClient(conn) + req := &pb.ListAccountsRequest{ + Paths: []string{ + path, + }, + } + resp, err = listerClient.ListAccounts(ctx, req) + release() + if err != nil { + w.log.Debug().Stringer("endpoint", w.endpoints[i]).Str("path", path).Err(err).Msg("Failed to list accounts") + errChan <- err + return + } + + if resp.GetState() != pb.ResponseState_SUCCEEDED { + errChan <- errors.New(fmt.Sprintf("request to list wallet accounts returned state %v", resp.GetState())) + return + } + + respsMap.Store(i, resp) + }() + } + wg.Wait() + close(errChan) + errsWaitGroup.Wait() + + resps := make([]*pb.ListAccountsResponse, 0, len(w.endpoints)) + respEndpoints := make([]*Endpoint, 0, len(w.endpoints)) + for i := range w.endpoints { + respAny, ok := respsMap.Load(i) + if !ok { continue } - listerClient := pb.NewListerClient(conn) - req := &pb.ListAccountsRequest{ - Paths: []string{ - path, - }, + resp := respAny.(*pb.ListAccountsResponse) + resps = append(resps, resp) + respEndpoints = append(respEndpoints, w.endpoints[i]) + } + + if len(resps) == 0 { + if len(errs) == 0 { + return nil, errors.New("internal error: no responses and no errors from endpoints") } - resp, err = listerClient.ListAccounts(ctx, req) - release() - if err == nil { - // Success. - break + // Wrap the errors in a single error. + wrappedErr := errs[0] + for _, err := range errs[1:] { + wrappedErr = errors.Wrap(wrappedErr, err.Error()) } - w.log.Debug().Stringer("endpoint", w.endpoints[i]).Str("path", path).Err(err).Msg("Failed to list accounts") - } - if err != nil { - return nil, errors.Wrap(err, "failed to access dirk") - } - if resp.GetState() != pb.ResponseState_SUCCEEDED { - return nil, fmt.Errorf("request to list wallet accounts returned state %v", resp.GetState()) + return nil, errors.Wrap(wrappedErr, "failed to access dirk") } + span.AddEvent("Obtained accounts") // sem := semaphore.NewWeighted(int64(runtime.GOMAXPROCS(0))) - var wg sync.WaitGroup + distributedAccountsMap := make(map[[48]byte]bool) + regularAccountsMap := make(map[[48]byte]bool) accounts := make([]e2wtypes.Account, 0) var accountsMu sync.Mutex - for _, respAccount := range resp.GetAccounts() { - wg.Add(1) - go func(respAccount *pb.Account, wg *sync.WaitGroup, mu *sync.Mutex) { - defer wg.Done() - account, err := w.obtainAccount(respAccount) - if err != nil { - w.log.Error().Err(err).Msg("Failed to obtain account") - } + for respIdx, resp := range resps { + endpoint := respEndpoints[respIdx] + for _, respAccount := range resp.GetAccounts() { + wg.Add(1) + go func(respAccount *pb.Account, endpoint *Endpoint, wg *sync.WaitGroup, mu *sync.Mutex) { + defer wg.Done() - mu.Lock() - accounts = append(accounts, account) - mu.Unlock() - }(respAccount, &wg, &accountsMu) - } - for _, respAccount := range resp.GetDistributedAccounts() { - wg.Add(1) - go func(respAccount *pb.DistributedAccount, wg *sync.WaitGroup, mu *sync.Mutex) { - defer wg.Done() + account, err := w.obtainAccount(respAccount, endpoint) + if err != nil { + w.log.Error().Err(err).Msg("Failed to obtain account") + } - account, err := w.obtainDistributedAccount(respAccount) - if err != nil { - w.log.Error().Err(err).Msg("Failed to obtain distributed account") - } + var pubkey [48]byte + copy(pubkey[:], respAccount.GetPublicKey()) + + mu.Lock() + defer mu.Unlock() + _, ok := regularAccountsMap[pubkey] + if ok { + w.log.Warn().Str("account", respAccount.GetName()).Str("pubkey", hex.EncodeToString(pubkey[:])).Msg("Duplicate pubkey found, ignoring") + return + } + regularAccountsMap[pubkey] = true + accounts = append(accounts, account) + }(respAccount, endpoint, &wg, &accountsMu) + } + for _, respAccount := range resp.GetDistributedAccounts() { + wg.Add(1) + go func(respAccount *pb.DistributedAccount, endpoint *Endpoint, wg *sync.WaitGroup, mu *sync.Mutex) { + defer wg.Done() + + account, err := w.obtainDistributedAccount(respAccount, endpoint) + if err != nil { + w.log.Error().Err(err).Msg("Failed to obtain distributed account") + } - mu.Lock() - accounts = append(accounts, account) - mu.Unlock() - }(respAccount, &wg, &accountsMu) + var pubkey [48]byte + copy(pubkey[:], respAccount.GetPublicKey()) + var compositePubKey [48]byte + copy(compositePubKey[:], respAccount.GetCompositePublicKey()) + + mu.Lock() + defer mu.Unlock() + _, ok := distributedAccountsMap[pubkey] + if ok { + // It's not normal to find duplicate distributed public keys. + w.log.Warn().Str("account", respAccount.GetName()).Str("pubkey", hex.EncodeToString(pubkey[:])).Msg("Duplicate distributed pubkey found, ignoring") + return + } + distributedAccountsMap[pubkey] = true + _, ok = distributedAccountsMap[compositePubKey] + if ok { + // It's normal to find duplicate composite public keys. + // It just means we've already tracked the account. + return + } + distributedAccountsMap[compositePubKey] = true + accounts = append(accounts, account) + }(respAccount, endpoint, &wg, &accountsMu) + } } wg.Wait() span.AddEvent("Processed accounts") @@ -1625,17 +1714,21 @@ func blsID(id uint64) *bls.ID { return &res } -func (w *wallet) obtainAccount(respAccount *pb.Account) ( +func (w *wallet) obtainAccount(respAccount *pb.Account, endpoint *Endpoint) ( e2wtypes.Account, error, ) { var key [48]byte copy(key[:], respAccount.GetPublicKey()) w.accountMapMu.RLock() - account, exists := w.accountMap[key] + cachedAccount, exists := w.accountMap[key] w.accountMapMu.RUnlock() if exists { - return account, nil + // Ensure endpoint is set even for cached accounts + if acc, ok := cachedAccount.(*account); ok { + acc.endpoint = endpoint + } + return cachedAccount, nil } pubKey, err := e2types.BLSPublicKeyFromBytes(respAccount.GetPublicKey()) @@ -1657,26 +1750,31 @@ func (w *wallet) obtainAccount(respAccount *pb.Account) ( name = respAccount.GetName() } - account = newAccount(w, uuid, name, pubKey, 1) + acc := newAccount(w, uuid, name, pubKey, 1) + acc.endpoint = endpoint w.accountMapMu.Lock() - w.accountMap[key] = account + w.accountMap[key] = acc w.accountMapMu.Unlock() - return account, nil + return acc, nil } -func (w *wallet) obtainDistributedAccount(respAccount *pb.DistributedAccount) ( +func (w *wallet) obtainDistributedAccount(respAccount *pb.DistributedAccount, endpoint *Endpoint) ( e2wtypes.Account, error, ) { var key [48]byte copy(key[:], respAccount.GetPublicKey()) w.accountMapMu.RLock() - account, exists := w.accountMap[key] + cachedAccount, exists := w.accountMap[key] w.accountMapMu.RUnlock() if exists { - return account, nil + // Ensure endpoint is set even for cached accounts + if acc, ok := cachedAccount.(*distributedAccount); ok { + acc.endpoint = endpoint + } + return cachedAccount, nil } pubKey, err := e2types.BLSPublicKeyFromBytes(respAccount.GetPublicKey()) @@ -1709,11 +1807,12 @@ func (w *wallet) obtainDistributedAccount(respAccount *pb.DistributedAccount) ( } } - account = newDistributedAccount(w, uuid, name, pubKey, compositePubKey, respAccount.GetSigningThreshold(), participants, 1) + acc := newDistributedAccount(w, uuid, name, pubKey, compositePubKey, respAccount.GetSigningThreshold(), participants, 1) + acc.endpoint = endpoint w.accountMapMu.Lock() - w.accountMap[key] = account + w.accountMap[key] = acc w.accountMapMu.Unlock() - return account, nil + return acc, nil } diff --git a/grpc_internal_test.go b/grpc_internal_test.go index b57037b..9dd9b12 100644 --- a/grpc_internal_test.go +++ b/grpc_internal_test.go @@ -84,12 +84,13 @@ func (c *BufConnectionProvider) Connection(ctx context.Context, endpoint *Endpoi pb.RegisterListerServer(server, c.listerServers[int(endpoint.port)%len(c.listerServers)]) } c.servers[serverAddress] = server - c.listeners[serverAddress] = bufconn.Listen(bufSize) - go func() { - if err := server.Serve(c.listeners[serverAddress]); err != nil { + listener := bufconn.Listen(bufSize) + c.listeners[serverAddress] = listener + go func(listener *bufconn.Listener) { + if err := server.Serve(listener); err != nil { log.Fatalf("Buffer server error: %v", err) } - }() + }(listener) } c.mutex.Unlock() @@ -119,6 +120,39 @@ func TestListGRPC(t *testing.T) { require.Equal(t, 8, accounts) } +func TestListGRPCDeduplication(t *testing.T) { + require.NoError(t, e2types.InitBLS()) + ctx := context.Background() + connectionProvider, err := NewBufConnectionProvider(ctx, []pb.ListerServer{&mock.MockListerServer{}}) + require.NoError(t, err) + w, err := OpenWallet(ctx, "Test wallet", credentials.NewTLS(nil), []*Endpoint{{host: "localhost", port: 12345}, {host: "localhost", port: 12346}}) + w.(*wallet).SetConnectionProvider(connectionProvider) + require.NoError(t, err) + accounts := 0 + for range w.Accounts(ctx) { + accounts++ + } + require.Equal(t, 8, accounts) +} + +func TestListGRPCAccountsFromSecondEndpoint(t *testing.T) { + mockListerServer := &mock.MockListerServerOverlappingAccounts{} + require.NoError(t, e2types.InitBLS()) + ctx := context.Background() + connectionProvider, err := NewBufConnectionProvider(ctx, []pb.ListerServer{mockListerServer}) + require.NoError(t, err) + w, err := OpenWallet(ctx, "Test wallet", credentials.NewTLS(nil), []*Endpoint{{host: "localhost", port: 12345}, {host: "localhost", port: 12346}}) + w.(*wallet).SetConnectionProvider(connectionProvider) + require.NoError(t, err) + accounts := 0 + for range w.Accounts(ctx) { + accounts++ + } + // The usual 8, plus an extra interop account, and an extra distributed account. + require.Equal(t, 10, accounts) + require.Equal(t, 2, mockListerServer.RequestsReceived) +} + func TestListGRPCErroring(t *testing.T) { require.NoError(t, e2types.InitBLS()) ctx := context.Background() @@ -264,4 +298,4 @@ func TestListGRPCErroring(t *testing.T) { // }, // ) // require.EqualError(t, err, "failed to obtain signature: not enough signatures: 1 signed, 0 denied, 0 failed, 2 errored") -// } +// } \ No newline at end of file diff --git a/mock/listerserver.go b/mock/listerserver.go index b27b4ea..46fe24d 100644 --- a/mock/listerserver.go +++ b/mock/listerserver.go @@ -19,10 +19,193 @@ import ( "errors" "fmt" "strings" + "sync" pb "github.com/wealdtech/eth2-signer-api/pb/v1" ) +type InteropAccountRegistry map[string]*pb.Account +type DistributedAccountRegistry map[string]*pb.DistributedAccount + +func (r InteropAccountRegistry) GetAccounts(in *pb.ListAccountsRequest) []*pb.Account { + accounts := make([]*pb.Account, 0) + for _, account := range r { + if len(in.GetPaths()) == 0 { + accounts = append(accounts, account) + } else { + for _, path := range in.GetPaths() { + if !strings.Contains(path, "/") { + accounts = append(accounts, account) + break + } + if strings.HasSuffix(path, fmt.Sprintf("/%s", account.GetName())) { + accounts = append(accounts, account) + break + } + } + } + } + return accounts +} + +func (r DistributedAccountRegistry) GetDistributedAccounts(in *pb.ListAccountsRequest) []*pb.DistributedAccount { + distributedAccounts := make([]*pb.DistributedAccount, 0) + for _, account := range r { + if len(in.GetPaths()) == 0 { + distributedAccounts = append(distributedAccounts, account) + } else { + for _, path := range in.GetPaths() { + if !strings.Contains(path, "/") { + distributedAccounts = append(distributedAccounts, account) + break + } + if strings.HasSuffix(path, fmt.Sprintf("/%s", account.GetName())) { + distributedAccounts = append(distributedAccounts, account) + break + } + } + } + } + return distributedAccounts +} + +var interop0 = &pb.Account{ + Name: "Interop 0", + PublicKey: _byte("0xa99a76ed7796f7be22d5b7e85deeb7c5677e88e511e0b337618f8c4eb61349b4bf2d153f649f7b53359fe8b94a38e44c"), + Uuid: _byte("0x00000000000000000000000000000000"), +} + +var interopAccounts = InteropAccountRegistry{ + "Interop 0": interop0, + "Interop 1": { + Name: "Interop 1", + PublicKey: _byte("0xb89bebc699769726a318c8e9971bd3171297c61aea4a6578a7a4f94b547dcba5bac16a89108b6b6a1fe3695d1a874a0b"), + Uuid: _byte("0x00000000000000000000000000000001"), + }, + "Interop 2": { + Name: "Interop 2", + PublicKey: _byte("0xa3a32b0f8b4ddb83f1a0a853d81dd725dfe577d4f4c3db8ece52ce2b026eca84815c1a7e8e92a4de3d755733bf7e4a9b"), + Uuid: _byte("0x00000000000000000000000000000002"), + }, + "Interop 3": { + Name: "Interop 3", + PublicKey: _byte("0x88c141df77cd9d8d7a71a75c826c41a9c9f03c6ee1b180f3e7852f6a280099ded351b58d66e653af8e42816a4d8f532e"), + Uuid: _byte("0x00000000000000000000000000000003"), + }, + "Interop 4": { + Name: "Interop 4", + PublicKey: _byte("0x81283b7a20e1ca460ebd9bbd77005d557370cabb1f9a44f530c4c4c66230f675f8df8b4c2818851aa7d77a80ca5a4a5e"), + Uuid: _byte("0x00000000000000000000000000000004"), + }, +} + +var distributed0 = &pb.DistributedAccount{ + Name: "Distributed 0", + PublicKey: _byte("0xaaf4abea98732aa9da46a4ddd8c56c03ec173a4daae90424e986be61d2b07999db746e103d6f505dc98716e91d4f946a"), + CompositePublicKey: _byte("0xa155a5fb0a6d732fa0f4d3714a8550ee5b90690475e010fbf89277e98e060203d69eba05fa71b2d0fa6aa6d091172f1e"), + SigningThreshold: 3, + Participants: []*pb.Endpoint{ + { + Id: 1, + Name: "signer-test01", + Port: 12001, + }, + { + Id: 2, + Name: "signer-test02", + Port: 12002, + }, + { + Id: 3, + Name: "signer-test03", + Port: 12003, + }, + { + Id: 4, + Name: "signer-test04", + Port: 12004, + }, + { + Id: 5, + Name: "signer-test05", + Port: 12005, + }, + }, + Uuid: _byte("0x01000000000000000000000000000000"), +} + +var allDistributedAccounts = DistributedAccountRegistry{ + "Interop 0": distributed0, + "Interop 1": { + Name: "Distributed 1", + PublicKey: _byte("0x98bc7c7596d70a27a243e6b6acc4a96bf1666428783671cb8545ced08a10c641fae1afbc83b525fce9357f6be667129e"), + CompositePublicKey: _byte("0x93c98077de26a2d382910c64664bb34ca3e29a5a6e3222c590b28efe9bc554b607677947cb6a44b168b2da5c74237fba"), + SigningThreshold: 3, + Participants: []*pb.Endpoint{ + { + Id: 1, + Name: "signer-test01", + Port: 12001, + }, + { + Id: 2, + Name: "signer-test02", + Port: 12002, + }, + { + Id: 3, + Name: "signer-test03", + Port: 12003, + }, + { + Id: 4, + Name: "signer-test04", + Port: 12004, + }, + { + Id: 5, + Name: "signer-test05", + Port: 12005, + }, + }, + Uuid: _byte("0x01000000000000000000000000000001"), + }, + "Interop 2": { + Name: "Distributed 2", + PublicKey: _byte("0x98552b2bdb1860c0c6363111477e3d738220988c9a2ee25fdeaa9971077d2ecde772c87dce07ce5f73754dcb585c43bf"), + CompositePublicKey: _byte("0xb75f33e5bd36841eb79f5018ad9f48494ddcc5b71bb671a59effdd2a139f8be18287df69bc028ce9b21e57d37bea5ffa"), + SigningThreshold: 3, + Participants: []*pb.Endpoint{ + { + Id: 1, + Name: "signer-test01", + Port: 12001, + }, + { + Id: 2, + Name: "signer-test02", + Port: 12002, + }, + { + Id: 3, + Name: "signer-test03", + Port: 12003, + }, + { + Id: 4, + Name: "signer-test04", + Port: 12004, + }, + { + Id: 5, + Name: "signer-test05", + Port: 12005, + }, + }, + Uuid: _byte("0x01000000000000000000000000000002"), + }, +} + func _byte(input string) []byte { res, _ := hex.DecodeString(strings.TrimPrefix(input, "0x")) return res @@ -59,179 +242,73 @@ type MockListerServer struct { // ListAccounts returns static accounts. func (s *MockListerServer) ListAccounts(_ context.Context, in *pb.ListAccountsRequest) (*pb.ListAccountsResponse, error) { - interopAccounts := map[string]*pb.Account{ - "Interop 0": { - Name: "Interop 0", - PublicKey: _byte("0xa99a76ed7796f7be22d5b7e85deeb7c5677e88e511e0b337618f8c4eb61349b4bf2d153f649f7b53359fe8b94a38e44c"), - Uuid: _byte("0x00000000000000000000000000000000"), - }, - "Interop 1": { - Name: "Interop 1", - PublicKey: _byte("0xb89bebc699769726a318c8e9971bd3171297c61aea4a6578a7a4f94b547dcba5bac16a89108b6b6a1fe3695d1a874a0b"), - Uuid: _byte("0x00000000000000000000000000000001"), - }, - "Interop 2": { - Name: "Interop 2", - PublicKey: _byte("0xa3a32b0f8b4ddb83f1a0a853d81dd725dfe577d4f4c3db8ece52ce2b026eca84815c1a7e8e92a4de3d755733bf7e4a9b"), - Uuid: _byte("0x00000000000000000000000000000002"), - }, - "Interop 3": { - Name: "Interop 3", - PublicKey: _byte("0x88c141df77cd9d8d7a71a75c826c41a9c9f03c6ee1b180f3e7852f6a280099ded351b58d66e653af8e42816a4d8f532e"), - Uuid: _byte("0x00000000000000000000000000000003"), - }, - "Interop 4": { - Name: "Interop 4", - PublicKey: _byte("0x81283b7a20e1ca460ebd9bbd77005d557370cabb1f9a44f530c4c4c66230f675f8df8b4c2818851aa7d77a80ca5a4a5e"), - Uuid: _byte("0x00000000000000000000000000000004"), - }, - } - allDistributedAccounts := map[string]*pb.DistributedAccount{ - "Interop 0": { - Name: "Distributed 0", - PublicKey: _byte("0xaaf4abea98732aa9da46a4ddd8c56c03ec173a4daae90424e986be61d2b07999db746e103d6f505dc98716e91d4f946a"), - CompositePublicKey: _byte("0xa155a5fb0a6d732fa0f4d3714a8550ee5b90690475e010fbf89277e98e060203d69eba05fa71b2d0fa6aa6d091172f1e"), - SigningThreshold: 3, - Participants: []*pb.Endpoint{ - { - Id: 1, - Name: "signer-test01", - Port: 12001, - }, - { - Id: 2, - Name: "signer-test02", - Port: 12002, - }, - { - Id: 3, - Name: "signer-test03", - Port: 12003, - }, - { - Id: 4, - Name: "signer-test04", - Port: 12004, - }, - { - Id: 5, - Name: "signer-test05", - Port: 12005, - }, - }, - Uuid: _byte("0x01000000000000000000000000000000"), - }, - "Interop 1": { - Name: "Distributed 1", - PublicKey: _byte("0x98bc7c7596d70a27a243e6b6acc4a96bf1666428783671cb8545ced08a10c641fae1afbc83b525fce9357f6be667129e"), - CompositePublicKey: _byte("0x93c98077de26a2d382910c64664bb34ca3e29a5a6e3222c590b28efe9bc554b607677947cb6a44b168b2da5c74237fba"), - SigningThreshold: 3, - Participants: []*pb.Endpoint{ - { - Id: 1, - Name: "signer-test01", - Port: 12001, - }, - { - Id: 2, - Name: "signer-test02", - Port: 12002, - }, - { - Id: 3, - Name: "signer-test03", - Port: 12003, - }, - { - Id: 4, - Name: "signer-test04", - Port: 12004, - }, - { - Id: 5, - Name: "signer-test05", - Port: 12005, - }, - }, - Uuid: _byte("0x01000000000000000000000000000001"), - }, - "Interop 2": { - Name: "Distributed 2", - PublicKey: _byte("0x98552b2bdb1860c0c6363111477e3d738220988c9a2ee25fdeaa9971077d2ecde772c87dce07ce5f73754dcb585c43bf"), - CompositePublicKey: _byte("0xb75f33e5bd36841eb79f5018ad9f48494ddcc5b71bb671a59effdd2a139f8be18287df69bc028ce9b21e57d37bea5ffa"), - SigningThreshold: 3, - Participants: []*pb.Endpoint{ - { - Id: 1, - Name: "signer-test01", - Port: 12001, - }, - { - Id: 2, - Name: "signer-test02", - Port: 12002, - }, - { - Id: 3, - Name: "signer-test03", - Port: 12003, - }, - { - Id: 4, - Name: "signer-test04", - Port: 12004, - }, - { - Id: 5, - Name: "signer-test05", - Port: 12005, - }, - }, - Uuid: _byte("0x01000000000000000000000000000002"), - }, - } + accounts := interopAccounts.GetAccounts(in) + distributedAccounts := allDistributedAccounts.GetDistributedAccounts(in) - accounts := make([]*pb.Account, 0) - for _, account := range interopAccounts { - if len(in.GetPaths()) == 0 { - accounts = append(accounts, account) - } else { - for _, path := range in.GetPaths() { - if !strings.Contains(path, "/") { - // Wallet only. - accounts = append(accounts, account) - break - } - if strings.HasSuffix(path, fmt.Sprintf("/%s", account.GetName())) { - accounts = append(accounts, account) - break - } - } - } - } + return &pb.ListAccountsResponse{ + State: pb.ResponseState_SUCCEEDED, + Accounts: accounts, + DistributedAccounts: distributedAccounts, + }, nil +} - distributedAccounts := make([]*pb.DistributedAccount, 0) - for _, account := range allDistributedAccounts { - if len(in.GetPaths()) == 0 { - distributedAccounts = append(distributedAccounts, account) - } else { - for _, path := range in.GetPaths() { - if !strings.Contains(path, "/") { - // Wallet only. - distributedAccounts = append(distributedAccounts, account) - break - } - if strings.HasSuffix(path, fmt.Sprintf("/%s", account.GetName())) { - distributedAccounts = append(distributedAccounts, account) - break - } - } - } +type MockListerServerOverlappingAccounts struct { + mutex sync.Mutex + RequestsReceived int + + pb.UnimplementedListerServer +} + +var interopAccountOverlapping = InteropAccountRegistry{ + "Interop 0": interop0, + "Extra Interop": { + Name: "Interop 0", + PublicKey: _byte("0x8c713b8d713b5166d85fdcae8d701ed1e6e5f2495e29bfcfadc2577a38230de27824a22ec18c798cbe6829144c5c9aef"), + Uuid: _byte("0x00000000000000000000000000000000"), + }, +} + +var distributedAccountOverlapping = DistributedAccountRegistry{ + "Interop 0": distributed0, + "Extra Distributed Deduped by Composite Public Key": { + Name: "Extra Distributed Deduped by Composite Public Key", + PublicKey: _byte("0x8695d953c289c41a3a1067fc2baab8f36b9390c1b1148fec4814d693e24e56bf8c25fbe947baf26d81c965c2804d1e61"), + // Same as allDistributedAccounts["Interop 2"].CompositePublicKey + CompositePublicKey: _byte("0xb75f33e5bd36841eb79f5018ad9f48494ddcc5b71bb671a59effdd2a139f8be18287df69bc028ce9b21e57d37bea5ffa"), + SigningThreshold: 3, + Participants: []*pb.Endpoint{}, + }, + "Unseen Distributed": { + Name: "Unseen Distributed", + PublicKey: _byte("0x80d206fbd30fd175f482b0b6e36da30ce66f6d31a5748da4d64925273873c16de1066c104512dcecac9e458a9540de6d"), + // Unique + CompositePublicKey: _byte("0xb3dc4efd216af7836942b0f84692e56bba286241295b7a3f37db261d42d9f2d2179c0847b15cb8009c83f2c9abc54172"), + SigningThreshold: 3, + Participants: []*pb.Endpoint{}, + }, +} + +// ListAccounts returns static accounts. +func (s *MockListerServerOverlappingAccounts) ListAccounts(_ context.Context, in *pb.ListAccountsRequest) (*pb.ListAccountsResponse, error) { + s.mutex.Lock() + defer func() { + s.RequestsReceived++ + s.mutex.Unlock() + }() + + registry := interopAccounts + distributedRegistry := allDistributedAccounts + if s.RequestsReceived%2 == 0 { + registry = interopAccountOverlapping + distributedRegistry = distributedAccountOverlapping } + accounts := registry.GetAccounts(in) + distributedAccounts := distributedRegistry.GetDistributedAccounts(in) + return &pb.ListAccountsResponse{ State: pb.ResponseState_SUCCEEDED, Accounts: accounts, DistributedAccounts: distributedAccounts, }, nil -} +} \ No newline at end of file diff --git a/wallet_internal_test.go b/wallet_internal_test.go index b2a69ec..36c36e0 100644 --- a/wallet_internal_test.go +++ b/wallet_internal_test.go @@ -93,7 +93,7 @@ func TestListDenyingServer(t *testing.T) { require.NoError(t, err) w.(*wallet).SetConnectionProvider(connectionProvider) _, err = w.(*wallet).List(ctx, "") - require.EqualError(t, err, "request to list wallet accounts returned state DENIED") + require.EqualError(t, err, "failed to access dirk: request to list wallet accounts returned state DENIED") } func TestList(t *testing.T) { @@ -185,4 +185,4 @@ func TestAccountByID(t *testing.T) { _, err = w.(e2wtypes.WalletAccountByIDProvider).AccountByID(ctx, uuid.MustParse("00000000-0000-0000-0000-000000000002")) require.EqualError(t, err, "not supported") -} +} \ No newline at end of file From 1099237499fd949a79fc03b413a07ed669974951 Mon Sep 17 00:00:00 2001 From: AntiD2ta Date: Thu, 8 Jan 2026 10:41:57 +0100 Subject: [PATCH 2/6] Migrate golangci-lint to version 2 --- .golangci.yml | 198 +++++++++++++++--------------------------- account.go | 1 + connectionprovider.go | 7 ++ distributedaccount.go | 1 + metrics.go | 2 + mock/listerserver.go | 20 +++-- parameters.go | 6 ++ wallet.go | 6 ++ 8 files changed, 106 insertions(+), 135 deletions(-) diff --git a/.golangci.yml b/.golangci.yml index 1b7c941..919bedd 100644 --- a/.golangci.yml +++ b/.golangci.yml @@ -1,149 +1,87 @@ -# This file contains all available configuration options -# with their default values (in comments). -# -# This file is not a configuration example, -# it contains the exhaustive configuration with explanations of the options. - -issues: - # Which files to skip: they will be analyzed, but issues from them won't be reported. - # Default value is empty list, - # but there is no need to include all autogenerated files, - # we confidently recognize autogenerated files. - # If it's not please let us know. - # "/" will be replaced by current OS file path separator to properly work on Windows. - exclude-files: - - ".*_ssz\\.go$" - -# Options for analysis running. +version: "2" run: - # The default concurrency value is the number of available CPU. - # concurrency: 4 - - # Timeout for analysis, e.g. 30s, 5m. - # Default: 1m - timeout: 10m - - # Exit code when at least one issue was found. - # Default: 1 - # issues-exit-code: 2 - - # Include test files or not. - # Default: true - tests: false - - # List of build tags, all linters use it. - # Default: []. - # build-tags: - # - mytag - - # Which dirs to skip: issues from them won't be reported. - # Can use regexp here: `generated.*`, regexp is applied on full path. - # Default value is empty list, - # but default dirs are skipped independently of this option's value (see skip-dirs-use-default). - # "/" will be replaced by current OS file path separator to properly work on Windows. - # skip-dirs: - # - autogenerated_by_my_lib - - # Enables skipping of directories: - # - vendor$, third_party$, testdata$, examples$, Godeps$, builtin$ - # Default: true - # skip-dirs-use-default: false - - # If set we pass it to "go list -mod={option}". From "go help modules": - # If invoked with -mod=readonly, the go command is disallowed from the implicit - # automatic updating of go.mod described above. Instead, it fails when any changes - # to go.mod are needed. This setting is most useful to check that go.mod does - # not need updates, such as in a continuous integration and testing system. - # If invoked with -mod=vendor, the go command assumes that the vendor - # directory holds the correct copies of dependencies and ignores - # the dependency descriptions in go.mod. - # - # Allowed values: readonly|vendor|mod - # By default, it isn't set. modules-download-mode: readonly - - # Allow multiple parallel golangci-lint instances running. - # If false (default) - golangci-lint acquires file lock on start. + tests: false allow-parallel-runners: true - - # Define the Go version limit. - # Mainly related to generics support since go1.18. - # Default: use Go version from the go.mod file, fallback on the env var `GOVERSION`, fallback on 1.18 - # go: '1.21' - - -# output configuration options -output: - # Format: colored-line-number|line-number|json|tab|checkstyle|code-climate|junit-xml|github-actions - # - # Multiple can be specified by separating them by comma, output can be provided - # for each of them by separating format name and path by colon symbol. - # Output path can be either `stdout`, `stderr` or path to the file to write to. - # Example: "checkstyle:report.json,colored-line-number" - # - # Default: colored-line-number - # format: json - - # Print lines of code with issue. - # Default: true - # print-issued-lines: false - - # Print linter name in the end of issue text. - # Default: true - # print-linter-name: false - - # Make issues output unique by line. - # Default: true - # uniq-by-line: false - - # Add a prefix to the output file references. - # Default is no prefix. - # path-prefix: "" - - # Sort results by: filepath, line and column. - # sort-results: true - - -# All available settings of specific linters. -linters-settings: - lll: - line-length: 132 - - stylecheck: - checks: [ "all", "-ST1000" ] - - tagliatelle: - case: - # use-field-name: true - rules: - json: snake - yaml: snake - - nlreturn: - # Allow two-line blocks without requiring a newline - block-size: 3 - linters: - # Enable all available linters. - # Default: false - enable-all: true - # Disable specific linter - # https://golangci-lint.run/usage/linters/#disabled-by-default + default: all disable: - cyclop - depguard - dupl - err113 - - execinquery - exhaustruct - - exportloopref - funlen - gochecknoglobals - gocognit - - gomnd - ireturn - lll - mnd - perfsprint - varnamelen - wsl + - maintidx + - noinlineerr + settings: + lll: + line-length: 132 + nlreturn: + block-size: 3 + staticcheck: + checks: + - all + - -ST1000 + tagliatelle: + case: + rules: + json: snake + yaml: snake + wsl_v5: + disable: + - assign + - branch + - decl + - defer + - expr + - for + - go + - if + - inc-dec + - label + - range + - return + - select + - send + - switch + - type-switch + - append + - assign-exclusive + - assign-expr + - err + - leading-whitespace + - trailing-whitespace + exclusions: + generated: lax + presets: + - comments + - common-false-positives + - legacy + - std-error-handling + paths: + - .*_ssz\.go$ + - third_party$ + - builtin$ + - examples$ +formatters: + enable: + - gci + - gofmt + - gofumpt + - goimports + exclusions: + generated: lax + paths: + - .*_ssz\.go$ + - third_party$ + - builtin$ + - examples$ diff --git a/account.go b/account.go index 3d7d9b7..fda56cb 100644 --- a/account.go +++ b/account.go @@ -85,6 +85,7 @@ func (a *account) Unlock(ctx context.Context, passphrase []byte) error { if err != nil { return errors.Wrap(err, "failed attempt to unlock account") } + if !unlocked { return errors.New("unlock attempt failed") } diff --git a/connectionprovider.go b/connectionprovider.go index e4aeb27..80a1480 100644 --- a/connectionprovider.go +++ b/connectionprovider.go @@ -58,8 +58,11 @@ func (c *PuddleConnectionProvider) Connection(ctx context.Context, endpoint *End func (c *PuddleConnectionProvider) obtainOrCreatePool(address string) *puddle.Pool[*grpc.ClientConn] { connectionPoolsMu.RLock() + pool, exists := connectionPools[address] + connectionPoolsMu.RUnlock() + if !exists { constructor := func(_ context.Context) (*grpc.ClientConn, error) { conn, err := grpc.NewClient(address, []grpc.DialOption{ @@ -77,6 +80,7 @@ func (c *PuddleConnectionProvider) obtainOrCreatePool(address string) *puddle.Po if err != nil { return nil, errors.Wrap(err, "failed to construct connection") } + incConnections(address) return conn, nil @@ -92,8 +96,11 @@ func (c *PuddleConnectionProvider) obtainOrCreatePool(address string) *puddle.Po Destructor: destructor, MaxSize: c.poolConnections, }) + connectionPoolsMu.Lock() + connectionPools[address] = pool + connectionPoolsMu.Unlock() } diff --git a/distributedaccount.go b/distributedaccount.go index 8ca64b7..f6d94ca 100644 --- a/distributedaccount.go +++ b/distributedaccount.go @@ -118,6 +118,7 @@ func (a *distributedAccount) Unlock(ctx context.Context, passphrase []byte) erro if err != nil { return errors.Wrap(err, "failed attempt to unlock account") } + if !unlocked { return errors.New("unlock attempt failed") } diff --git a/metrics.go b/metrics.go index e309387..c34919b 100644 --- a/metrics.go +++ b/metrics.go @@ -34,10 +34,12 @@ func registerMetrics(ctx context.Context, monitor Metrics) error { // Already registered. return nil } + if monitor == nil { // No monitor. return nil } + if monitor.Presenter() == "prometheus" { return registerPrometheusMetrics(ctx) } diff --git a/mock/listerserver.go b/mock/listerserver.go index 46fe24d..1a9b79f 100644 --- a/mock/listerserver.go +++ b/mock/listerserver.go @@ -24,11 +24,14 @@ import ( pb "github.com/wealdtech/eth2-signer-api/pb/v1" ) -type InteropAccountRegistry map[string]*pb.Account -type DistributedAccountRegistry map[string]*pb.DistributedAccount +type ( + InteropAccountRegistry map[string]*pb.Account + DistributedAccountRegistry map[string]*pb.DistributedAccount +) func (r InteropAccountRegistry) GetAccounts(in *pb.ListAccountsRequest) []*pb.Account { accounts := make([]*pb.Account, 0) + for _, account := range r { if len(in.GetPaths()) == 0 { accounts = append(accounts, account) @@ -38,6 +41,7 @@ func (r InteropAccountRegistry) GetAccounts(in *pb.ListAccountsRequest) []*pb.Ac accounts = append(accounts, account) break } + if strings.HasSuffix(path, fmt.Sprintf("/%s", account.GetName())) { accounts = append(accounts, account) break @@ -45,11 +49,13 @@ func (r InteropAccountRegistry) GetAccounts(in *pb.ListAccountsRequest) []*pb.Ac } } } + return accounts } func (r DistributedAccountRegistry) GetDistributedAccounts(in *pb.ListAccountsRequest) []*pb.DistributedAccount { distributedAccounts := make([]*pb.DistributedAccount, 0) + for _, account := range r { if len(in.GetPaths()) == 0 { distributedAccounts = append(distributedAccounts, account) @@ -59,6 +65,7 @@ func (r DistributedAccountRegistry) GetDistributedAccounts(in *pb.ListAccountsRe distributedAccounts = append(distributedAccounts, account) break } + if strings.HasSuffix(path, fmt.Sprintf("/%s", account.GetName())) { distributedAccounts = append(distributedAccounts, account) break @@ -66,6 +73,7 @@ func (r DistributedAccountRegistry) GetDistributedAccounts(in *pb.ListAccountsRe } } } + return distributedAccounts } @@ -253,10 +261,10 @@ func (s *MockListerServer) ListAccounts(_ context.Context, in *pb.ListAccountsRe } type MockListerServerOverlappingAccounts struct { + pb.UnimplementedListerServer + mutex sync.Mutex RequestsReceived int - - pb.UnimplementedListerServer } var interopAccountOverlapping = InteropAccountRegistry{ @@ -291,6 +299,7 @@ var distributedAccountOverlapping = DistributedAccountRegistry{ // ListAccounts returns static accounts. func (s *MockListerServerOverlappingAccounts) ListAccounts(_ context.Context, in *pb.ListAccountsRequest) (*pb.ListAccountsResponse, error) { s.mutex.Lock() + defer func() { s.RequestsReceived++ s.mutex.Unlock() @@ -298,6 +307,7 @@ func (s *MockListerServerOverlappingAccounts) ListAccounts(_ context.Context, in registry := interopAccounts distributedRegistry := allDistributedAccounts + if s.RequestsReceived%2 == 0 { registry = interopAccountOverlapping distributedRegistry = distributedAccountOverlapping @@ -311,4 +321,4 @@ func (s *MockListerServerOverlappingAccounts) ListAccounts(_ context.Context, in Accounts: accounts, DistributedAccounts: distributedAccounts, }, nil -} \ No newline at end of file +} diff --git a/parameters.go b/parameters.go index fe9681a..b22ae66 100644 --- a/parameters.go +++ b/parameters.go @@ -99,6 +99,7 @@ func parseAndCheckParameters(params ...Parameter) (*parameters, error) { poolConnections: 128, monitor: &nullMetrics{}, } + for _, p := range params { if params != nil { p.apply(¶meters) @@ -108,18 +109,23 @@ func parseAndCheckParameters(params ...Parameter) (*parameters, error) { if parameters.monitor == nil { return nil, errors.New("no monitor specified") } + if parameters.timeout == 0 { return nil, errors.New("no timeout specified") } + if parameters.name == "" { return nil, errors.New("no name specified") } + if parameters.credentials == nil { return nil, errors.New("no credentials specified") } + if len(parameters.endpoints) == 0 { return nil, errors.New("no endpoints specified") } + if parameters.poolConnections < 1 { return nil, errors.New("no pool connections specified") } diff --git a/wallet.go b/wallet.go index 8fc37d7..296c65f 100644 --- a/wallet.go +++ b/wallet.go @@ -81,6 +81,7 @@ func Open(ctx context.Context, wallet.name = parameters.name wallet.timeout = parameters.timeout wallet.endpoints = make([]*Endpoint, len(parameters.endpoints)) + wallet.connectionProvider = &PuddleConnectionProvider{ name: parameters.name, poolConnections: parameters.poolConnections, @@ -92,6 +93,7 @@ func Open(ctx context.Context, port: parameters.endpoints[i].port, } } + wallet.log.Trace().Str("name", wallet.name).Msg("Opened wallet") return wallet, nil @@ -103,6 +105,7 @@ func OpenWallet(_ context.Context, name string, credentials credentials.Transpor wallet := newWallet() wallet.name = name wallet.endpoints = make([]*Endpoint, len(endpoints)) + wallet.connectionProvider = &PuddleConnectionProvider{ poolConnections: 32, credentials: credentials.Clone(), @@ -159,6 +162,7 @@ func (w *wallet) IsUnlocked(_ context.Context) (bool, error) { // Accounts provides all accounts in the wallet. func (w *wallet) Accounts(ctx context.Context) <-chan e2wtypes.Account { ch := make(chan e2wtypes.Account, 1024) + go func() { accounts, err := w.List(ctx, "") if err != nil { @@ -168,6 +172,7 @@ func (w *wallet) Accounts(ctx context.Context) <-chan e2wtypes.Account { ch <- account } } + close(ch) }() @@ -181,6 +186,7 @@ func (w *wallet) AccountByName(ctx context.Context, name string) (e2wtypes.Accou if err != nil { return nil, errors.Wrap(err, "failed to obtain account") } + if len(accounts) == 0 { return nil, errors.New("not found") } From c11dddac6b986f595b2ea8b16894144646379d22 Mon Sep 17 00:00:00 2001 From: AntiD2ta Date: Thu, 8 Jan 2026 10:43:29 +0100 Subject: [PATCH 3/6] Persist proper endpoints to accounts for correct signing --- grpc.go | 244 +++++++++++++++++++++++++++++++++++++++--- grpc_internal_test.go | 146 ++++++++++++++++++++++++- mock/listerserver.go | 36 +++++++ mock/signerserver.go | 60 +++++++++++ 4 files changed, 468 insertions(+), 18 deletions(-) create mode 100644 mock/signerserver.go diff --git a/grpc.go b/grpc.go index 8abcb7c..67e65a0 100644 --- a/grpc.go +++ b/grpc.go @@ -47,10 +47,12 @@ func ComposeCredentials(ctx context.Context, certPath string, keyPath string, ca if err != nil { return nil, errors.Wrap(err, "failed to obtain client certificate") } + clientKey, err := os.ReadFile(keyPath) if err != nil { return nil, errors.Wrap(err, "failed to obtain client key") } + var caCert []byte if caCertPath != "" { caCert, err = os.ReadFile(caCertPath) @@ -79,6 +81,7 @@ func Credentials(_ context.Context, clientCert []byte, clientKey []byte, caCert if !cp.AppendCertsFromPEM(caCert) { return nil, errors.New("failed to add CA certificate") } + tlsCfg.RootCAs = cp } @@ -103,36 +106,47 @@ func (w *wallet) List(ctx context.Context, accountPath string) ([]e2wtypes.Accou return nil, errors.New("wallet has no endpoints") } - var wg sync.WaitGroup - var errsWaitGroup sync.WaitGroup + var ( + wg sync.WaitGroup + errsWaitGroup sync.WaitGroup + ) + errs := make([]error, 0) errChan := make(chan error) // add a tick to the wait group to ensure errs is populated after // wg.Wait() and errsWaitGroup.Wait() are called. errsWaitGroup.Add(1) + go func(errChan chan error) { defer errsWaitGroup.Done() + for err := range errChan { errs = append(errs, err) } }(errChan) respsMap := new(sync.Map) + for i := range len(w.endpoints) { wg.Add(1) - go func() { - var resp *pb.ListAccountsResponse - var err error + go func() { + var ( + resp *pb.ListAccountsResponse + err error + ) i := i defer wg.Done() ctx, cancelFunc := context.WithTimeout(ctx, w.timeout) defer cancelFunc() - var conn *grpc.ClientConn - var release func() + var ( + conn *grpc.ClientConn + release func() + ) + conn, release, err = w.connectionProvider.Connection(ctx, w.endpoints[i]) if err != nil { w.log.Debug().Stringer("endpoint", w.endpoints[i]).Str("path", path).Err(err).Msg("Failed to obtain connection") @@ -148,6 +162,7 @@ func (w *wallet) List(ctx context.Context, accountPath string) ([]e2wtypes.Accou } resp, err = listerClient.ListAccounts(ctx, req) release() + if err != nil { w.log.Debug().Stringer("endpoint", w.endpoints[i]).Str("path", path).Err(err).Msg("Failed to list accounts") errChan <- err @@ -155,18 +170,20 @@ func (w *wallet) List(ctx context.Context, accountPath string) ([]e2wtypes.Accou } if resp.GetState() != pb.ResponseState_SUCCEEDED { - errChan <- errors.New(fmt.Sprintf("request to list wallet accounts returned state %v", resp.GetState())) + errChan <- fmt.Errorf("request to list wallet accounts returned state %v", resp.GetState()) return } respsMap.Store(i, resp) }() } + wg.Wait() close(errChan) errsWaitGroup.Wait() resps := make([]*pb.ListAccountsResponse, 0, len(w.endpoints)) + respEndpoints := make([]*Endpoint, 0, len(w.endpoints)) for i := range w.endpoints { respAny, ok := respsMap.Load(i) @@ -174,7 +191,11 @@ func (w *wallet) List(ctx context.Context, accountPath string) ([]e2wtypes.Accou continue } - resp := respAny.(*pb.ListAccountsResponse) + resp, ok := respAny.(*pb.ListAccountsResponse) + if !ok { + continue + } + resps = append(resps, resp) respEndpoints = append(respEndpoints, w.endpoints[i]) } @@ -188,6 +209,7 @@ func (w *wallet) List(ctx context.Context, accountPath string) ([]e2wtypes.Accou for _, err := range errs[1:] { wrappedErr = errors.Wrap(wrappedErr, err.Error()) } + return nil, errors.Wrap(wrappedErr, "failed to access dirk") } @@ -197,12 +219,15 @@ func (w *wallet) List(ctx context.Context, accountPath string) ([]e2wtypes.Accou distributedAccountsMap := make(map[[48]byte]bool) regularAccountsMap := make(map[[48]byte]bool) accounts := make([]e2wtypes.Account, 0) + var accountsMu sync.Mutex for respIdx, resp := range resps { endpoint := respEndpoints[respIdx] + for _, respAccount := range resp.GetAccounts() { wg.Add(1) + go func(respAccount *pb.Account, endpoint *Endpoint, wg *sync.WaitGroup, mu *sync.Mutex) { defer wg.Done() @@ -216,17 +241,22 @@ func (w *wallet) List(ctx context.Context, accountPath string) ([]e2wtypes.Accou mu.Lock() defer mu.Unlock() + _, ok := regularAccountsMap[pubkey] if ok { w.log.Warn().Str("account", respAccount.GetName()).Str("pubkey", hex.EncodeToString(pubkey[:])).Msg("Duplicate pubkey found, ignoring") return } + regularAccountsMap[pubkey] = true + accounts = append(accounts, account) }(respAccount, endpoint, &wg, &accountsMu) } + for _, respAccount := range resp.GetDistributedAccounts() { wg.Add(1) + go func(respAccount *pb.DistributedAccount, endpoint *Endpoint, wg *sync.WaitGroup, mu *sync.Mutex) { defer wg.Done() @@ -237,29 +267,36 @@ func (w *wallet) List(ctx context.Context, accountPath string) ([]e2wtypes.Accou var pubkey [48]byte copy(pubkey[:], respAccount.GetPublicKey()) + var compositePubKey [48]byte copy(compositePubKey[:], respAccount.GetCompositePublicKey()) mu.Lock() defer mu.Unlock() + _, ok := distributedAccountsMap[pubkey] if ok { // It's not normal to find duplicate distributed public keys. w.log.Warn().Str("account", respAccount.GetName()).Str("pubkey", hex.EncodeToString(pubkey[:])).Msg("Duplicate distributed pubkey found, ignoring") return } + distributedAccountsMap[pubkey] = true + _, ok = distributedAccountsMap[compositePubKey] if ok { // It's normal to find duplicate composite public keys. // It just means we've already tracked the account. return } + distributedAccountsMap[compositePubKey] = true + accounts = append(accounts, account) }(respAccount, endpoint, &wg, &accountsMu) } } + wg.Wait() span.AddEvent("Processed accounts") @@ -285,12 +322,15 @@ func (w *wallet) UnlockAccount(ctx context.Context, accountName string, passphra Account: fmt.Sprintf("%s/%s", w.Name(), accountName), Passphrase: passphrase, } + ctx, cancelFunc := context.WithTimeout(ctx, w.timeout) defer cancelFunc() + resp, err := accountManagerClient.Unlock(ctx, req) if err != nil { return false, errors.Wrap(err, "failed to access dirk") } + if resp.GetState() == pb.ResponseState_FAILED { return false, errors.New("request to unlock account failed") } @@ -316,12 +356,15 @@ func (w *wallet) LockAccount(ctx context.Context, accountName string) error { req := &pb.LockAccountRequest{ Account: fmt.Sprintf("%s/%s", w.Name(), accountName), } + ctx, cancelFunc := context.WithTimeout(ctx, w.timeout) defer cancelFunc() + resp, err := accountManagerClient.Lock(ctx, req) if err != nil { return errors.Wrap(err, "failed to access dirk") } + if resp.GetState() == pb.ResponseState_FAILED { return errors.New("request to lock account failed") } @@ -353,7 +396,14 @@ func (a *account) SignGRPC(ctx context.Context, Domain: domain, } - conn, release, err := a.wallet.connectionProvider.Connection(ctx, a.wallet.endpoints[0]) + // Use the endpoint set in the account if available, + // otherwise use the first endpoint from the wallet (for backwards compatibility). + endpoint := a.endpoint + if endpoint == nil { + endpoint = a.wallet.endpoints[0] + } + + conn, release, err := a.wallet.connectionProvider.Connection(ctx, endpoint) if err != nil { return nil, errors.Wrap(err, "failed to connect to endpoint") } @@ -363,20 +413,25 @@ func (a *account) SignGRPC(ctx context.Context, if client == nil { return nil, errors.New("failed to set up signing client") } + ctx, cancelFunc := context.WithTimeout(ctx, a.wallet.timeout) defer cancelFunc() + resp, err := client.Sign(ctx, req) if err != nil { return nil, errors.Wrap(err, "failed to obtain signature") } + if resp.GetState() == pb.ResponseState_FAILED { span.SetStatus(codes.Error, "Failed to request signature bytes") return nil, errors.New("request to obtain signature failed") } + if resp.GetState() == pb.ResponseState_DENIED { span.SetStatus(codes.Error, "Request signature bytes denied") return nil, errors.New("request to obtain signature denied") } + span.AddEvent("Obtained signature bytes") sig, err := e2types.BLSSignatureFromBytes(resp.GetSignature()) @@ -384,10 +439,12 @@ func (a *account) SignGRPC(ctx context.Context, span.SetStatus(codes.Error, "Invalid signature bytes received") return nil, errors.Wrap(err, "invalid signature received") } + if sig == nil { span.SetStatus(codes.Error, "No signature received") return nil, fmt.Errorf("no signature received") } + span.AddEvent("Generated signature from bytes") return sig, nil @@ -419,6 +476,7 @@ func (a *distributedAccount) SignGRPC(ctx context.Context, ctx, cancelFunc := context.WithTimeout(ctx, a.wallet.timeout) defer cancelFunc() + sig, err := a.thresholdSign(ctx, req) if err != nil { return nil, errors.Wrap(err, "failed to obtain signature") @@ -457,13 +515,21 @@ func (a *account) SignMultiGRPC(ctx context.Context, if !isAccount { return nil, errors.New("account not of required type") } + req.Requests[i] = &pb.SignRequest{ Id: &pb.SignRequest_Account{Account: fmt.Sprintf("%s/%s", assertedAccount.wallet.Name(), accounts[i].Name())}, Data: data[i], Domain: domain, } } - endpoint := a.wallet.endpoints[0] + + // Use the endpoint set in the account if available, + // otherwise use the first endpoint from the wallet (for backwards compatibility). + endpoint := a.endpoint + if endpoint == nil { + endpoint = a.wallet.endpoints[0] + } + conn, release, err := a.wallet.connectionProvider.Connection(ctx, endpoint) if err != nil { return nil, errors.Wrap(err, "failed to connect to endpoint") @@ -474,26 +540,32 @@ func (a *account) SignMultiGRPC(ctx context.Context, if client == nil { return nil, errors.New("failed to set up signing client") } + ctx, cancelFunc := context.WithTimeout(ctx, a.wallet.timeout) defer cancelFunc() + resp, err := client.Multisign(ctx, req) if err != nil { return nil, errors.Wrap(err, "failed to obtain signatures") } sigs := make([]e2types.Signature, len(accounts)) + for i, response := range resp.GetResponses() { if response.GetState() == pb.ResponseState_FAILED { return nil, errors.New("request to obtain signatures failed") } + if response.GetState() == pb.ResponseState_DENIED { return nil, errors.New("request to obtain signatures denied") } + sigs[i], err = e2types.BLSSignatureFromBytes(response.GetSignature()) if err != nil { return nil, errors.Wrap(err, fmt.Sprintf("invalid signature received from %v", endpoint)) } } + span.AddEvent("Generated signatures from bytes") return sigs, nil @@ -530,6 +602,7 @@ func (a *distributedAccount) SignMultiGRPC(ctx context.Context, if !isAccount { return nil, errors.New("account not of required type") } + thresholds[i] = assertedAccount.signingThreshold req.Requests[i] = &pb.SignRequest{ Id: &pb.SignRequest_Account{Account: fmt.Sprintf("%s/%s", assertedAccount.wallet.Name(), accounts[i].Name())}, @@ -540,6 +613,7 @@ func (a *distributedAccount) SignMultiGRPC(ctx context.Context, ctx, cancelFunc := context.WithTimeout(ctx, a.wallet.timeout) defer cancelFunc() + sigs, err := a.thresholdMultiSign(ctx, req, thresholds) if err != nil { return nil, errors.Wrap(err, "failed to obtain signature") @@ -579,7 +653,13 @@ func (a *account) SignBeaconProposalGRPC(ctx context.Context, Domain: domain, } - endpoint := a.wallet.endpoints[0] + // Use the endpoint set in the account if available, + // otherwise use the first endpoint from the wallet (for backwards compatibility). + endpoint := a.endpoint + if endpoint == nil { + endpoint = a.wallet.endpoints[0] + } + conn, release, err := a.wallet.connectionProvider.Connection(ctx, endpoint) if err != nil { return nil, errors.Wrap(err, "failed to connect to endpoint") @@ -590,15 +670,19 @@ func (a *account) SignBeaconProposalGRPC(ctx context.Context, if client == nil { return nil, errors.New("failed to set up signing client") } + ctx, cancelFunc := context.WithTimeout(ctx, a.wallet.timeout) defer cancelFunc() + resp, err := client.SignBeaconProposal(ctx, req) if err != nil { return nil, errors.Wrap(err, "failed to obtain signature") } + if resp.GetState() == pb.ResponseState_FAILED { return nil, errors.New("request to obtain signature failed") } + if resp.GetState() == pb.ResponseState_DENIED { return nil, errors.New("request to obtain signature denied") } @@ -607,6 +691,7 @@ func (a *account) SignBeaconProposalGRPC(ctx context.Context, if err != nil { return nil, errors.Wrap(err, fmt.Sprintf("invalid signature received from %v", endpoint)) } + if sig == nil { return nil, fmt.Errorf("no signature received from %v", endpoint) } @@ -647,6 +732,7 @@ func (a *distributedAccount) SignBeaconProposalGRPC(ctx context.Context, ctx, cancelFunc := context.WithTimeout(ctx, a.wallet.timeout) defer cancelFunc() + sig, err := a.thresholdSignBeaconProposal(ctx, req) if err != nil { return nil, errors.Wrap(err, "failed to obtain signature") @@ -694,7 +780,13 @@ func (a *account) SignBeaconAttestationGRPC(ctx context.Context, Domain: domain, } - endpoint := a.wallet.endpoints[0] + // Use the endpoint set in the account if available, + // otherwise use the first endpoint from the wallet (for backwards compatibility). + endpoint := a.endpoint + if endpoint == nil { + endpoint = a.wallet.endpoints[0] + } + conn, release, err := a.wallet.connectionProvider.Connection(ctx, endpoint) if err != nil { return nil, errors.Wrap(err, "failed to connect to endpoint") @@ -705,15 +797,19 @@ func (a *account) SignBeaconAttestationGRPC(ctx context.Context, if client == nil { return nil, errors.New("failed to set up signing client") } + ctx, cancelFunc := context.WithTimeout(ctx, a.wallet.timeout) defer cancelFunc() + resp, err := client.SignBeaconAttestation(ctx, req) if err != nil { return nil, errors.Wrap(err, "failed to obtain signature") } + if resp.GetState() == pb.ResponseState_FAILED { return nil, errors.New("request to obtain signature failed") } + if resp.GetState() == pb.ResponseState_DENIED { return nil, errors.New("request to obtain signature denied") } @@ -722,6 +818,7 @@ func (a *account) SignBeaconAttestationGRPC(ctx context.Context, if err != nil { return nil, errors.Wrap(err, fmt.Sprintf("invalid signature received from %v", endpoint)) } + if sig == nil { return nil, fmt.Errorf("no signature received from %v", endpoint) } @@ -812,6 +909,7 @@ func (a *account) SignBeaconAttestationsGRPC(ctx context.Context, if !isAccount { return nil, errors.New("account not of required type") } + req.Requests[i] = &pb.SignBeaconAttestationRequest{ Id: &pb.SignBeaconAttestationRequest_Account{Account: fmt.Sprintf("%s/%s", account.wallet.Name(), accounts[i].Name())}, Data: &pb.AttestationData{ @@ -831,7 +929,13 @@ func (a *account) SignBeaconAttestationsGRPC(ctx context.Context, } } - endpoint := a.wallet.endpoints[0] + // Use the endpoint set in the account if available, + // otherwise use the first endpoint from the wallet (for backwards compatibility). + endpoint := a.endpoint + if endpoint == nil { + endpoint = a.wallet.endpoints[0] + } + conn, release, err := a.wallet.connectionProvider.Connection(ctx, endpoint) if err != nil { return nil, errors.Wrap(err, "failed to connect to endpoint") @@ -845,19 +949,23 @@ func (a *account) SignBeaconAttestationsGRPC(ctx context.Context, ctx, cancelFunc := context.WithTimeout(ctx, a.wallet.timeout) defer cancelFunc() + resp, err := client.SignBeaconAttestations(ctx, req) if err != nil { return nil, errors.Wrap(err, "failed to obtain signatures") } sigs := make([]e2types.Signature, len(accounts)) + for i, response := range resp.GetResponses() { if response.GetState() == pb.ResponseState_FAILED { return nil, errors.New("request to obtain signatures failed") } + if response.GetState() == pb.ResponseState_DENIED { return nil, errors.New("request to obtain signatures denied") } + sigs[i], err = e2types.BLSSignatureFromBytes(response.GetSignature()) if err != nil { return nil, errors.Wrap(err, fmt.Sprintf("invalid signature received from %v", endpoint)) @@ -896,6 +1004,7 @@ func (a *distributedAccount) SignBeaconAttestationsGRPC(ctx context.Context, } thresholds := make([]uint32, len(accounts)) + req := &pb.SignBeaconAttestationsRequest{ Requests: make([]*pb.SignBeaconAttestationRequest, len(accounts)), } @@ -904,6 +1013,7 @@ func (a *distributedAccount) SignBeaconAttestationsGRPC(ctx context.Context, if !isAccount { return nil, errors.New("account not of required type") } + thresholds[i] = account.signingThreshold req.Requests[i] = &pb.SignBeaconAttestationRequest{ Id: &pb.SignBeaconAttestationRequest_Account{Account: fmt.Sprintf("%s/%s", account.wallet.Name(), accounts[i].Name())}, @@ -948,7 +1058,10 @@ func (w *wallet) GenerateDistributedAccount(ctx context.Context, )) defer span.End() - conn, release, err := w.connectionProvider.Connection(ctx, w.endpoints[0]) + // Use the first endpoint from the wallet (for backwards compatibility). + endpoint := w.endpoints[0] + + conn, release, err := w.connectionProvider.Connection(ctx, endpoint) if err != nil { return nil, errors.Wrap(err, "failed to connect to endpoint") } @@ -961,8 +1074,10 @@ func (w *wallet) GenerateDistributedAccount(ctx context.Context, SigningThreshold: signingThreshold, Passphrase: passphrase, } + ctx, cancelFunc := context.WithTimeout(ctx, w.timeout) defer cancelFunc() + resp, err := accountClient.Generate(ctx, req) if err != nil { return nil, errors.Wrap(err, "failed to access dirk") @@ -986,6 +1101,7 @@ func (w *wallet) GenerateDistributedAccount(ctx context.Context, if err != nil { return nil, errors.New("failed to confirm created account") } + if len(accountList) == 0 { return nil, errors.New("failed to obtain created account") } @@ -1006,6 +1122,7 @@ func (a *distributedAccount) thresholdSign(ctx context.Context, req *pb.SignRequ return nil, errors.Wrap(err, fmt.Sprintf("failed to connect to endpoint %v", endpoint)) } defer release() + span.AddEvent("Obtained connection") clients[id] = pb.NewSignerClient(conn) @@ -1013,27 +1130,34 @@ func (a *distributedAccount) thresholdSign(ctx context.Context, req *pb.SignRequ return nil, fmt.Errorf("failed to set up signing client for %v", endpoint) } } + span.AddEvent("Obtained connections") type thresholdSignResponse struct { id uint64 resp *pb.SignResponse } + respChannel := make(chan *thresholdSignResponse, len(clients)) type thresholdSignError struct { id uint64 err error } + errChannel := make(chan *thresholdSignError, len(clients)) span.AddEvent("Ready to contact servers") + ctx, cancelFunc := context.WithTimeout(ctx, a.wallet.timeout) defer cancelFunc() + for id, client := range clients { go func(client pb.SignerClient, id uint64, req *pb.SignRequest) { resp, err := client.Sign(ctx, req) + span.AddEvent("Received response") + if err != nil { errChannel <- &thresholdSignError{ id: id, @@ -1045,9 +1169,11 @@ func (a *distributedAccount) thresholdSign(ctx context.Context, req *pb.SignRequ resp: resp, } } + span.AddEvent("Processed response") }(client, id, req) } + span.AddEvent("Contacted all servers") // Wait for enough responses (or timeout) @@ -1056,6 +1182,7 @@ func (a *distributedAccount) thresholdSign(ctx context.Context, req *pb.SignRequ failed := 0 errored := 0 ids := make([]bls.ID, a.signingThreshold) + signatures := make([]bls.Sign, a.signingThreshold) for signed < int(a.signingThreshold) && signed+denied+failed+errored != len(clients) { select { @@ -1063,6 +1190,7 @@ func (a *distributedAccount) thresholdSign(ctx context.Context, req *pb.SignRequ return nil, errors.New("context done") case resp := <-errChannel: a.wallet.log.Warn().Uint64("server_id", resp.id).Err(resp.err).Msg("Received error") + errored++ case resp := <-respChannel: switch resp.resp.GetState() { @@ -1078,16 +1206,19 @@ func (a *distributedAccount) thresholdSign(ctx context.Context, req *pb.SignRequ if err := signatures[signed].Deserialize(resp.resp.GetSignature()); err != nil { return nil, errors.Wrap(err, fmt.Sprintf("invalid signature received from %d", resp.id)) } + signed++ } } } + span.AddEvent("Received responses", trace.WithAttributes( attribute.Int("signed", signed), attribute.Int("denied", denied), attribute.Int("failed", failed), attribute.Int("errored", errored), )) + if signed < int(a.signingThreshold) { return nil, fmt.Errorf("not enough signatures: %d signed, %d denied, %d failed, %d errored", signed, denied, failed, errored) } @@ -1096,6 +1227,7 @@ func (a *distributedAccount) thresholdSign(ctx context.Context, req *pb.SignRequ if err := signature.Recover(signatures, ids); err != nil { return nil, errors.Wrap(err, "failed to recover composite signature") } + span.AddEvent("Recovered signature") //nolint:wrapcheck @@ -1125,16 +1257,19 @@ func (a *distributedAccount) thresholdMultiSign(ctx context.Context, req *pb.Mul id uint64 resp *pb.MultisignResponse } + respChannel := make(chan *thresholdSignResponse, len(clients)) type thresholdSignError struct { id uint64 err error } + errChannel := make(chan *thresholdSignError, len(clients)) ctx, cancelFunc := context.WithTimeout(ctx, a.wallet.timeout) defer cancelFunc() + for id, client := range clients { go func(client pb.SignerClient, id uint64, req *pb.MultisignRequest) { resp, err := client.Multisign(ctx, req) @@ -1151,6 +1286,7 @@ func (a *distributedAccount) thresholdMultiSign(ctx context.Context, req *pb.Mul } }(client, id, req) } + span.AddEvent("Contacted all servers") // Wait for enough responses (or context done). @@ -1160,23 +1296,28 @@ func (a *distributedAccount) thresholdMultiSign(ctx context.Context, req *pb.Mul failed := make([]int, len(thresholds)) errored := make([]int, len(thresholds)) ids := make([][]bls.ID, len(thresholds)) + signatureBytes := make([][][]byte, len(thresholds)) for i := range ids { ids[i] = make([]bls.ID, 0, len(clients)) signatureBytes[i] = make([][]byte, 0, len(clients)) } + for { select { case <-ctx.Done(): return nil, errors.New("context done") case resp := <-errChannel: a.wallet.log.Warn().Uint64("server_id", resp.id).Err(resp.err).Msg("Received error") + for i := range errored { errored[i]++ } + responses++ case resp := <-respChannel: responses++ + for i, response := range resp.resp.GetResponses() { switch response.GetState() { case pb.ResponseState_DENIED: @@ -1202,26 +1343,32 @@ func (a *distributedAccount) thresholdMultiSign(ctx context.Context, req *pb.Mul // We could be done early if we have enough signatures. done := true + for i := range ids { if len(ids[i]) < int(thresholds[i]) { done = false break } } + if done { break } } + span.AddEvent("Received responses") // Take the signature bytes, turn them in to real signatures, then // recover the final signature from the components. sem := semaphore.NewWeighted(int64(runtime.GOMAXPROCS(0))) + var wg sync.WaitGroup + res := make([]e2types.Signature, len(thresholds)) - var err error + for i := range ids { wg.Add(1) + go func(_ context.Context, _ *semaphore.Weighted, wg *sync.WaitGroup, i int) { defer wg.Done() @@ -1231,6 +1378,8 @@ func (a *distributedAccount) thresholdMultiSign(ctx context.Context, req *pb.Mul Int("index", i). Logger() + var err error + if signed[i] < int(thresholds[i]) { log.Error(). Int("signed", signed[i]). @@ -1251,6 +1400,7 @@ func (a *distributedAccount) thresholdMultiSign(ctx context.Context, req *pb.Mul return } } + var signature bls.Sign if err := signature.Recover(components, ids[i][0:thresholds[i]]); err != nil { // Invalid components. @@ -1258,20 +1408,24 @@ func (a *distributedAccount) thresholdMultiSign(ctx context.Context, req *pb.Mul for i := range components { sigs[i] = fmt.Sprintf("%#x", components[i].Serialize()) } + log.Error().Err(err).Strs("sigs", sigs).Msg("Failed to recover signature") return } + res[i], err = e2types.BLSSignatureFromSig(signature) if err != nil { // Invalid composite signature. log.Error(). Err(err). Msg("Failed to recreate signature") + res[i] = nil } }(ctx, sem, &wg, i) } + wg.Wait() span.AddEvent("Recovered signatures") @@ -1306,16 +1460,19 @@ func (a *distributedAccount) thresholdSignBeaconAttestation(ctx context.Context, id uint64 resp *pb.SignResponse } + respChannel := make(chan *thresholdSignResponse, len(clients)) type thresholdSignError struct { id uint64 err error } + errChannel := make(chan *thresholdSignError, len(clients)) ctx, cancelFunc := context.WithTimeout(ctx, a.wallet.timeout) defer cancelFunc() + for id, client := range clients { go func(client pb.SignerClient, id uint64, req *pb.SignBeaconAttestationRequest) { resp, err := client.SignBeaconAttestation(ctx, req) @@ -1332,6 +1489,7 @@ func (a *distributedAccount) thresholdSignBeaconAttestation(ctx context.Context, } }(client, id, req) } + span.AddEvent("Contacted all servers") // Wait for enough responses (or context done). @@ -1340,6 +1498,7 @@ func (a *distributedAccount) thresholdSignBeaconAttestation(ctx context.Context, failed := 0 errored := 0 ids := make([]bls.ID, a.signingThreshold) + signatures := make([]bls.Sign, a.signingThreshold) for signed < int(a.signingThreshold) && signed+denied+failed+errored != len(clients) { select { @@ -1347,6 +1506,7 @@ func (a *distributedAccount) thresholdSignBeaconAttestation(ctx context.Context, return nil, errors.New("context done") case resp := <-errChannel: a.wallet.log.Warn().Uint64("server_id", resp.id).Err(resp.err).Msg("Received error") + errored++ case resp := <-respChannel: switch resp.resp.GetState() { @@ -1362,16 +1522,19 @@ func (a *distributedAccount) thresholdSignBeaconAttestation(ctx context.Context, if err := signatures[signed].Deserialize(resp.resp.GetSignature()); err != nil { return nil, errors.Wrap(err, fmt.Sprintf("invalid signature received from %d", resp.id)) } + signed++ } } } + span.AddEvent("Received responses", trace.WithAttributes( attribute.Int("signed", signed), attribute.Int("denied", denied), attribute.Int("failed", failed), attribute.Int("errored", errored), )) + if signed < int(a.signingThreshold) { a.wallet.log.Error(). Int("signed", signed). @@ -1389,10 +1552,12 @@ func (a *distributedAccount) thresholdSignBeaconAttestation(ctx context.Context, for i := range signatures { sigs[i] = fmt.Sprintf("%#x", signatures[i].Serialize()) } + a.wallet.log.Error().Err(err).Strs("sigs", sigs).Msg("Failed to recover signature") return nil, errors.Wrap(err, "failed to recover composite signature") } + span.AddEvent("Recovered signature") //nolint:wrapcheck @@ -1428,16 +1593,19 @@ func (a *distributedAccount) thresholdSignBeaconAttestations(ctx context.Context id uint64 resp *pb.MultisignResponse } + respChannel := make(chan *thresholdSignResponse, len(clients)) type thresholdSignError struct { id uint64 err error } + errChannel := make(chan *thresholdSignError, len(clients)) ctx, cancelFunc := context.WithTimeout(ctx, a.wallet.timeout) defer cancelFunc() + for id, client := range clients { go func(client pb.SignerClient, id uint64, req *pb.SignBeaconAttestationsRequest) { resp, err := client.SignBeaconAttestations(ctx, req) @@ -1454,6 +1622,7 @@ func (a *distributedAccount) thresholdSignBeaconAttestations(ctx context.Context } }(client, id, req) } + span.AddEvent("Contacted all servers") // Wait for enough responses (or context done). @@ -1463,23 +1632,28 @@ func (a *distributedAccount) thresholdSignBeaconAttestations(ctx context.Context failed := make([]int, len(thresholds)) errored := make([]int, len(thresholds)) ids := make([][]bls.ID, len(thresholds)) + signatureBytes := make([][][]byte, len(thresholds)) for i := range ids { ids[i] = make([]bls.ID, 0, len(clients)) signatureBytes[i] = make([][]byte, 0, len(clients)) } + for { select { case <-ctx.Done(): return nil, errors.New("context done") case resp := <-errChannel: a.wallet.log.Warn().Uint64("server_id", resp.id).Err(resp.err).Msg("Received error") + for i := range errored { errored[i]++ } + responses++ case resp := <-respChannel: responses++ + for i, response := range resp.resp.GetResponses() { switch response.GetState() { case pb.ResponseState_DENIED: @@ -1505,26 +1679,32 @@ func (a *distributedAccount) thresholdSignBeaconAttestations(ctx context.Context // We could be done early if we have enough signatures. done := true + for i := range ids { if len(ids[i]) < int(thresholds[i]) { done = false break } } + if done { break } } + span.AddEvent("Received responses") // Take the signature bytes, turn them in to real signatures, then // recover the final signature from the components. sem := semaphore.NewWeighted(int64(runtime.GOMAXPROCS(0))) + var wg sync.WaitGroup + res := make([]e2types.Signature, len(thresholds)) - var err error + for i := range ids { wg.Add(1) + go func(_ context.Context, _ *semaphore.Weighted, wg *sync.WaitGroup, i int) { defer wg.Done() @@ -1534,6 +1714,8 @@ func (a *distributedAccount) thresholdSignBeaconAttestations(ctx context.Context Int("index", i). Logger() + var err error + if signed[i] < int(thresholds[i]) { log.Error(). Int("signed", signed[i]). @@ -1554,6 +1736,7 @@ func (a *distributedAccount) thresholdSignBeaconAttestations(ctx context.Context return } } + var signature bls.Sign if err := signature.Recover(components, ids[i][0:thresholds[i]]); err != nil { // Invalid components. @@ -1561,20 +1744,24 @@ func (a *distributedAccount) thresholdSignBeaconAttestations(ctx context.Context for i := range components { sigs[i] = fmt.Sprintf("%#x", components[i].Serialize()) } + log.Error().Err(err).Strs("sigs", sigs).Msg("Failed to recover signature") return } + res[i], err = e2types.BLSSignatureFromSig(signature) if err != nil { // Invalid composite signature. log.Error(). Err(err). Msg("Failed to recreate signature") + res[i] = nil } }(ctx, sem, &wg, i) } + wg.Wait() span.AddEvent("Recovered signatures") @@ -1609,16 +1796,19 @@ func (a *distributedAccount) thresholdSignBeaconProposal(ctx context.Context, id uint64 resp *pb.SignResponse } + respChannel := make(chan *thresholdSignResponse, len(clients)) type thresholdSignError struct { id uint64 err error } + errChannel := make(chan *thresholdSignError, len(clients)) ctx, cancelFunc := context.WithTimeout(ctx, a.wallet.timeout) defer cancelFunc() + for id, client := range clients { go func(client pb.SignerClient, id uint64, req *pb.SignBeaconProposalRequest) { resp, err := client.SignBeaconProposal(ctx, req) @@ -1635,6 +1825,7 @@ func (a *distributedAccount) thresholdSignBeaconProposal(ctx context.Context, } }(client, id, req) } + span.AddEvent("Contacted all servers") // Wait for enough responses (or context done). @@ -1643,6 +1834,7 @@ func (a *distributedAccount) thresholdSignBeaconProposal(ctx context.Context, failed := 0 errored := 0 ids := make([]bls.ID, a.signingThreshold) + signatures := make([]bls.Sign, a.signingThreshold) for signed < int(a.signingThreshold) && signed+denied+failed+errored != len(clients) { select { @@ -1650,6 +1842,7 @@ func (a *distributedAccount) thresholdSignBeaconProposal(ctx context.Context, return nil, errors.New("context done") case resp := <-errChannel: a.wallet.log.Warn().Uint64("server_id", resp.id).Err(resp.err).Msg("Received error") + errored++ case resp := <-respChannel: switch resp.resp.GetState() { @@ -1665,16 +1858,19 @@ func (a *distributedAccount) thresholdSignBeaconProposal(ctx context.Context, if err := signatures[signed].Deserialize(resp.resp.GetSignature()); err != nil { return nil, errors.Wrap(err, fmt.Sprintf("invalid signature received from %d", resp.id)) } + signed++ } } } + span.AddEvent("Received responses", trace.WithAttributes( attribute.Int("signed", signed), attribute.Int("denied", denied), attribute.Int("failed", failed), attribute.Int("errored", errored), )) + if signed < int(a.signingThreshold) { a.wallet.log.Error(). Int("signed", signed). @@ -1692,10 +1888,12 @@ func (a *distributedAccount) thresholdSignBeaconProposal(ctx context.Context, for i := range signatures { sigs[i] = fmt.Sprintf("%#x", signatures[i].Serialize()) } + a.wallet.log.Error().Err(err).Strs("sigs", sigs).Msg("Failed to recover signature") return nil, errors.Wrap(err, "failed to recover composite signature") } + span.AddEvent("Recovered signature") //nolint:wrapcheck @@ -1705,8 +1903,10 @@ func (a *distributedAccount) thresholdSignBeaconProposal(ctx context.Context, // blsID turns a uint64 in to a BLS identifier. func blsID(id uint64) *bls.ID { var res bls.ID + buf := [8]byte{} binary.LittleEndian.PutUint64(buf[:], id) + if err := res.SetLittleEndian(buf[:]); err != nil { panic(err) } @@ -1723,11 +1923,13 @@ func (w *wallet) obtainAccount(respAccount *pb.Account, endpoint *Endpoint) ( w.accountMapMu.RLock() cachedAccount, exists := w.accountMap[key] w.accountMapMu.RUnlock() + if exists { // Ensure endpoint is set even for cached accounts if acc, ok := cachedAccount.(*account); ok { acc.endpoint = endpoint } + return cachedAccount, nil } @@ -1737,12 +1939,14 @@ func (w *wallet) obtainAccount(respAccount *pb.Account, endpoint *Endpoint) ( } var uuid uuid.UUID + err = uuid.UnmarshalBinary(respAccount.GetUuid()) if err != nil { return nil, errors.Wrap(err, fmt.Sprintf("uuid %x invalid", respAccount.GetUuid())) } walletPrefixLen := len(w.Name()) + 1 + var name string if strings.Contains(respAccount.GetName(), "/") { name = respAccount.GetName()[walletPrefixLen:] @@ -1769,11 +1973,13 @@ func (w *wallet) obtainDistributedAccount(respAccount *pb.DistributedAccount, en w.accountMapMu.RLock() cachedAccount, exists := w.accountMap[key] w.accountMapMu.RUnlock() + if exists { // Ensure endpoint is set even for cached accounts if acc, ok := cachedAccount.(*distributedAccount); ok { acc.endpoint = endpoint } + return cachedAccount, nil } @@ -1788,17 +1994,21 @@ func (w *wallet) obtainDistributedAccount(respAccount *pb.DistributedAccount, en } var uuid uuid.UUID + err = uuid.UnmarshalBinary(respAccount.GetUuid()) if err != nil { return nil, errors.Wrap(err, fmt.Sprintf("uuid %x invalid", respAccount.GetUuid())) } + var name string + walletPrefixLen := len(w.Name()) + 1 if strings.Contains(respAccount.GetName(), "/") { name = respAccount.GetName()[walletPrefixLen:] } else { name = respAccount.GetName() } + participants := make(map[uint64]*Endpoint, len(respAccount.GetParticipants())) for _, participant := range respAccount.GetParticipants() { participants[participant.GetId()] = &Endpoint{ diff --git a/grpc_internal_test.go b/grpc_internal_test.go index 9dd9b12..fee969a 100644 --- a/grpc_internal_test.go +++ b/grpc_internal_test.go @@ -15,9 +15,11 @@ package dirk import ( "context" + "encoding/hex" "fmt" "log" "net" + "strings" "sync" "testing" @@ -26,12 +28,18 @@ import ( pb "github.com/wealdtech/eth2-signer-api/pb/v1" e2types "github.com/wealdtech/go-eth2-types/v2" mock "github.com/wealdtech/go-eth2-wallet-dirk/mock" + e2wtypes "github.com/wealdtech/go-eth2-wallet-types/v2" "google.golang.org/grpc" "google.golang.org/grpc/credentials" "google.golang.org/grpc/credentials/insecure" "google.golang.org/grpc/test/bufconn" ) +func _byte(input string) []byte { + res, _ := hex.DecodeString(strings.TrimPrefix(input, "0x")) + return res +} + // ErroringConnectionProvider throws errors. type ErroringConnectionProvider struct { pb.UnimplementedListerServer @@ -53,6 +61,7 @@ type BufConnectionProvider struct { servers map[string]*grpc.Server listeners map[string]*bufconn.Listener listerServers []pb.ListerServer + signerServer pb.SignerServer } const bufSize = 1024 * 1024 @@ -68,6 +77,19 @@ func NewBufConnectionProvider(_ context.Context, }, nil } +// NewBufConnectionProviderWithSigner creates a new buffer connection provider with signer server support. +func NewBufConnectionProviderWithSigner(_ context.Context, + listerServers []pb.ListerServer, + signerServer pb.SignerServer, +) (*BufConnectionProvider, error) { + return &BufConnectionProvider{ + listerServers: listerServers, + signerServer: signerServer, + servers: make(map[string]*grpc.Server), + listeners: make(map[string]*bufconn.Listener), + }, nil +} + func (c *BufConnectionProvider) bufDialer(_ context.Context, in string) (net.Conn, error) { return c.listeners[in].Dial() } @@ -83,6 +105,9 @@ func (c *BufConnectionProvider) Connection(ctx context.Context, endpoint *Endpoi // Pick a server from the available list. pb.RegisterListerServer(server, c.listerServers[int(endpoint.port)%len(c.listerServers)]) } + if c.signerServer != nil { + pb.RegisterSignerServer(server, c.signerServer) + } c.servers[serverAddress] = server listener := bufconn.Listen(bufSize) c.listeners[serverAddress] = listener @@ -168,6 +193,125 @@ func TestListGRPCErroring(t *testing.T) { require.Equal(t, 8, accounts) } +// TestAccountUsesCorrectEndpointForSigning verifies that accounts use the endpoint +// that returned their data during the List operation for signing operations +func TestAccountUsesCorrectEndpointForSigning(t *testing.T) { + require.NoError(t, e2types.InitBLS()) + ctx := context.Background() + + // Create mock signer server to track signing calls + mockSigner := &mock.MockSignerServer{} + + // Create custom lister servers that return different accounts for different endpoints + // endpoint1Server returns one unique account + endpoint1Server := &mock.CustomListerServer{ + Accounts: []*pb.Account{ + { + Name: "Account Endpoint1", + PublicKey: _byte("0xa99a76ed7796f7be22d5b7e85deeb7c5677e88e511e0b337618f8c4eb61349b4bf2d153f649f7b53359fe8b94a38e44c"), + Uuid: _byte("0x00000000000000000000000000000001"), + }, + }, + } + + // endpoint2Server returns one unique account and one shared account (same as endpoint1) + endpoint2Server := &mock.CustomListerServer{ + Accounts: []*pb.Account{ + { + Name: "Account Endpoint1", // Same account as endpoint1 + PublicKey: _byte("0xa99a76ed7796f7be22d5b7e85deeb7c5677e88e511e0b337618f8c4eb61349b4bf2d153f649f7b53359fe8b94a38e44c"), + Uuid: _byte("0x00000000000000000000000000000001"), + }, + { + Name: "Account Endpoint2", + PublicKey: _byte("0xb89bebc699769726a318c8e9971bd3171297c61aea4a6578a7a4f94b547dcba5bac16a89108b6b6a1fe3695d1a874a0b"), + Uuid: _byte("0x00000000000000000000000000000002"), + }, + }, + } + + connectionProvider, err := NewBufConnectionProviderWithSigner(ctx, []pb.ListerServer{endpoint2Server, endpoint1Server}, mockSigner) + require.NoError(t, err) + + // Set up wallet with multiple endpoints + endpoints := []*Endpoint{ + + {host: "localhost", port: 12345}, // Will use endpoint1Server (port % 2 = 1) + {host: "localhost", port: 12346}, // Will use endpoint2Server (port % 2 = 0) + } + + w, err := OpenWallet(ctx, "Test wallet", credentials.NewTLS(nil), endpoints) + w.(*wallet).SetConnectionProvider(connectionProvider) + require.NoError(t, err) + + // List accounts - this should assign endpoints to accounts based on which endpoint returned them + accounts, err := w.(*wallet).List(ctx, "") + require.NoError(t, err) + + require.Equal(t, 2, len(accounts), "Should have 2 unique accounts (shared account not duplicated)") + + // Create a map to track accounts by endpoint + accountsByEndpoint := make(map[string][]e2wtypes.Account) + endpointByAccount := make(map[string]*Endpoint) + + for _, acct := range accounts { + var accountEndpoint *Endpoint + if acc, ok := acct.(*account); ok { + accountEndpoint = acc.endpoint + } + + require.NotNil(t, accountEndpoint, "Account should have an endpoint assigned") + + endpointKey := fmt.Sprintf("%s:%d", accountEndpoint.host, accountEndpoint.port) + accountsByEndpoint[endpointKey] = append(accountsByEndpoint[endpointKey], acct) + endpointByAccount[acct.Name()] = accountEndpoint + } + + // Verify that the shared account "Account Endpoint1" is assigned to one of the endpoints that returned it + sharedAccountEndpoint := fmt.Sprintf("%s:%d", endpointByAccount["Account Endpoint1"].host, endpointByAccount["Account Endpoint1"].port) + require.True(t, sharedAccountEndpoint == "localhost:12345" || sharedAccountEndpoint == "localhost:12346", + "Shared account should be assigned to one of the endpoints that returned it, got: %s", sharedAccountEndpoint) + + // Verify that Account Endpoint2 is assigned to endpoint 12346 (only endpoint that returns it) + require.Equal(t, "localhost:12346", fmt.Sprintf("%s:%d", endpointByAccount["Account Endpoint2"].host, endpointByAccount["Account Endpoint2"].port)) + + // Verify that we have exactly 2 accounts total (no duplication of the shared account) + require.Equal(t, 2, len(accounts), "Should have exactly 2 unique accounts") + + // The endpoint distribution depends on which goroutine finished last + // But we should have accounts assigned to their endpoints + totalAccountsAssigned := 0 + for _, accountsOnEndpoint := range accountsByEndpoint { + totalAccountsAssigned += len(accountsOnEndpoint) + } + require.Equal(t, 2, totalAccountsAssigned, "All accounts should be assigned to endpoints") + + // Test signing with both accounts to verify they use their assigned endpoints + for _, acct := range accounts { + // Sign with the account - this should use the endpoint assigned during List operation + _, err := acct.(e2wtypes.AccountProtectingSigner).SignGeneric(ctx, + []byte{0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, + 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00}, + []byte{0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, + 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00}, + ) + + // The signing should have succeeded (no error) + require.NoError(t, err) + } + + // Verify that the signer server was called with both accounts + endpointsUsed := mockSigner.GetEndpointsUsed() + require.Len(t, endpointsUsed, 2, "Signer server should have been called for both accounts") + require.Contains(t, endpointsUsed, "Test wallet/Account Endpoint1", "Should have signed with shared account") + require.Contains(t, endpointsUsed, "Test wallet/Account Endpoint2", "Should have signed with account from endpoint 2") + + // Verify that signing works with whatever endpoint the shared account got assigned to + // This demonstrates that endpoint reassignment works correctly and accounts use their assigned endpoint for signing + t.Logf("Shared account 'Account Endpoint1' was assigned to endpoint: %s:%d", + endpointByAccount["Account Endpoint1"].host, endpointByAccount["Account Endpoint1"].port) +} + // Disabled because it results in a link back to Dirk repository for // "github.com/attestantio/dirk/testing/daemon" // "github.com/attestantio/dirk/testing/resources" @@ -298,4 +442,4 @@ func TestListGRPCErroring(t *testing.T) { // }, // ) // require.EqualError(t, err, "failed to obtain signature: not enough signatures: 1 signed, 0 denied, 0 failed, 2 errored") -// } \ No newline at end of file +// } diff --git a/mock/listerserver.go b/mock/listerserver.go index 1a9b79f..50533d9 100644 --- a/mock/listerserver.go +++ b/mock/listerserver.go @@ -322,3 +322,39 @@ func (s *MockListerServerOverlappingAccounts) ListAccounts(_ context.Context, in DistributedAccounts: distributedAccounts, }, nil } + +// CustomListerServer returns accounts specific to an endpoint. +type CustomListerServer struct { + pb.UnimplementedListerServer + + Accounts []*pb.Account +} + +func (s *CustomListerServer) ListAccounts(_ context.Context, in *pb.ListAccountsRequest) (*pb.ListAccountsResponse, error) { + // Filter accounts based on the request + var filteredAccounts []*pb.Account + + for _, account := range s.Accounts { + if len(in.GetPaths()) == 0 { + filteredAccounts = append(filteredAccounts, account) + } else { + for _, path := range in.GetPaths() { + if !strings.Contains(path, "/") { + filteredAccounts = append(filteredAccounts, account) + break + } + + if strings.HasSuffix(path, fmt.Sprintf("/%s", account.GetName())) { + filteredAccounts = append(filteredAccounts, account) + break + } + } + } + } + + return &pb.ListAccountsResponse{ + State: pb.ResponseState_SUCCEEDED, + Accounts: filteredAccounts, + DistributedAccounts: nil, + }, nil +} diff --git a/mock/signerserver.go b/mock/signerserver.go new file mode 100644 index 0000000..71162f2 --- /dev/null +++ b/mock/signerserver.go @@ -0,0 +1,60 @@ +// Copyright © 2020, 2021 Weald Technology Trading. +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +// Package dirk provides mock implementations for testing. +package dirk + +import ( + "context" + "sync" + + pb "github.com/wealdtech/eth2-signer-api/pb/v1" +) + +// MockSignerServer tracks which endpoints are used for signing operations. +type MockSignerServer struct { + pb.UnimplementedSignerServer + + endpointsUsed []string + mutex sync.Mutex +} + +func (s *MockSignerServer) Sign(_ context.Context, req *pb.SignRequest) (*pb.SignResponse, error) { + s.mutex.Lock() + defer s.mutex.Unlock() + + // Extract account info from the request + s.endpointsUsed = append(s.endpointsUsed, req.GetAccount()) + + return &pb.SignResponse{ + State: pb.ResponseState_SUCCEEDED, + Signature: []byte{ + 0x8e, 0x16, 0x21, 0x7c, 0xfb, 0x18, 0xe2, 0xf2, 0xb2, 0xc8, 0x88, 0x5b, 0x02, 0xd7, 0x34, 0x36, + 0x00, 0x69, 0x03, 0xba, 0x77, 0x32, 0x0b, 0x43, 0xa8, 0xcd, 0x7b, 0x60, 0x30, 0xbe, 0x67, 0x94, + 0x95, 0x46, 0x38, 0x1e, 0xfb, 0xd0, 0x9e, 0x8d, 0x21, 0x47, 0x85, 0x5b, 0x05, 0xad, 0x8c, 0xc9, + 0x11, 0x93, 0x33, 0xf4, 0x28, 0x99, 0xaa, 0xf7, 0x45, 0xa7, 0x61, 0x1e, 0x4f, 0xad, 0x52, 0xaa, + 0x08, 0xe6, 0xa2, 0x80, 0xe1, 0xef, 0x4e, 0xf9, 0xc5, 0x3c, 0x42, 0x60, 0x28, 0xca, 0xbf, 0x5b, + 0x45, 0xd7, 0x3c, 0xb5, 0xbc, 0x8c, 0x34, 0x3c, 0xd9, 0x44, 0xa9, 0x99, 0xda, 0x1e, 0x6f, 0x4e, + }, + }, nil +} + +func (s *MockSignerServer) GetEndpointsUsed() []string { + s.mutex.Lock() + defer s.mutex.Unlock() + + result := make([]string, len(s.endpointsUsed)) + copy(result, s.endpointsUsed) + + return result +} From 62c14385d6684f88f7972bce4f3766c211494b7f Mon Sep 17 00:00:00 2001 From: AntiD2ta Date: Thu, 8 Jan 2026 11:04:56 +0100 Subject: [PATCH 4/6] Make signing test stricter with endpoint-specific account validation --- grpc_internal_test.go | 64 +++++++++++++++++++++++++++++++++-------- mock/signerserver.go | 44 ++++++++++++++++++++++++++-- wallet_internal_test.go | 2 +- 3 files changed, 94 insertions(+), 16 deletions(-) diff --git a/grpc_internal_test.go b/grpc_internal_test.go index fee969a..9add2f0 100644 --- a/grpc_internal_test.go +++ b/grpc_internal_test.go @@ -61,7 +61,7 @@ type BufConnectionProvider struct { servers map[string]*grpc.Server listeners map[string]*bufconn.Listener listerServers []pb.ListerServer - signerServer pb.SignerServer + signerServers map[int]pb.SignerServer } const bufSize = 1024 * 1024 @@ -81,10 +81,28 @@ func NewBufConnectionProvider(_ context.Context, func NewBufConnectionProviderWithSigner(_ context.Context, listerServers []pb.ListerServer, signerServer pb.SignerServer, +) (*BufConnectionProvider, error) { + signerServers := make(map[int]pb.SignerServer) + // If only one signer server provided, use it for all ports + for i := range listerServers { + signerServers[i] = signerServer + } + return &BufConnectionProvider{ + listerServers: listerServers, + signerServers: signerServers, + servers: make(map[string]*grpc.Server), + listeners: make(map[string]*bufconn.Listener), + }, nil +} + +// NewBufConnectionProviderWithSigners creates a new buffer connection provider with different signer servers for different endpoints. +func NewBufConnectionProviderWithSigners(_ context.Context, + listerServers []pb.ListerServer, + signerServers map[int]pb.SignerServer, ) (*BufConnectionProvider, error) { return &BufConnectionProvider{ listerServers: listerServers, - signerServer: signerServer, + signerServers: signerServers, servers: make(map[string]*grpc.Server), listeners: make(map[string]*bufconn.Listener), }, nil @@ -105,8 +123,10 @@ func (c *BufConnectionProvider) Connection(ctx context.Context, endpoint *Endpoi // Pick a server from the available list. pb.RegisterListerServer(server, c.listerServers[int(endpoint.port)%len(c.listerServers)]) } - if c.signerServer != nil { - pb.RegisterSignerServer(server, c.signerServer) + if c.signerServers != nil { + if signerServer, exists := c.signerServers[int(endpoint.port)%len(c.listerServers)]; exists && signerServer != nil { + pb.RegisterSignerServer(server, signerServer) + } } c.servers[serverAddress] = server listener := bufconn.Listen(bufSize) @@ -199,8 +219,11 @@ func TestAccountUsesCorrectEndpointForSigning(t *testing.T) { require.NoError(t, e2types.InitBLS()) ctx := context.Background() - // Create mock signer server to track signing calls - mockSigner := &mock.MockSignerServer{} + // Create separate mock signer servers for each endpoint with account restrictions + // endpoint1Signer (index 1) only accepts "Account Endpoint1" + endpoint1Signer := mock.NewMockSignerServerWithAccounts([]string{"Account Endpoint1"}) + // endpoint2Signer (index 0) accepts both "Account Endpoint1" and "Account Endpoint2" + endpoint2Signer := mock.NewMockSignerServerWithAccounts([]string{"Account Endpoint1", "Account Endpoint2"}) // Create custom lister servers that return different accounts for different endpoints // endpoint1Server returns one unique account @@ -230,7 +253,10 @@ func TestAccountUsesCorrectEndpointForSigning(t *testing.T) { }, } - connectionProvider, err := NewBufConnectionProviderWithSigner(ctx, []pb.ListerServer{endpoint2Server, endpoint1Server}, mockSigner) + connectionProvider, err := NewBufConnectionProviderWithSigners(ctx, []pb.ListerServer{endpoint2Server, endpoint1Server}, map[int]pb.SignerServer{ + 0: endpoint2Signer, // port % 2 == 0 uses endpoint2Signer + 1: endpoint1Signer, // port % 2 == 1 uses endpoint1Signer + }) require.NoError(t, err) // Set up wallet with multiple endpoints @@ -300,11 +326,25 @@ func TestAccountUsesCorrectEndpointForSigning(t *testing.T) { require.NoError(t, err) } - // Verify that the signer server was called with both accounts - endpointsUsed := mockSigner.GetEndpointsUsed() - require.Len(t, endpointsUsed, 2, "Signer server should have been called for both accounts") - require.Contains(t, endpointsUsed, "Test wallet/Account Endpoint1", "Should have signed with shared account") - require.Contains(t, endpointsUsed, "Test wallet/Account Endpoint2", "Should have signed with account from endpoint 2") + // Verify that the correct signer servers were called with the correct accounts + endpoint1Used := endpoint1Signer.GetEndpointsUsed() + endpoint2Used := endpoint2Signer.GetEndpointsUsed() + + // "Account Endpoint1" should be signed by whichever endpoint it was assigned to + // "Account Endpoint2" should only be signed by endpoint2Signer (port 12346) + if endpointByAccount["Account Endpoint1"].port == 12345 { + // Account Endpoint1 assigned to endpoint1 (port 12345) + require.Len(t, endpoint1Used, 1, "endpoint1Signer should have been called once for Account Endpoint1") + require.Contains(t, endpoint1Used, "Test wallet/Account Endpoint1", "endpoint1Signer should have signed Account Endpoint1") + require.Len(t, endpoint2Used, 1, "endpoint2Signer should have been called once for Account Endpoint2") + require.Contains(t, endpoint2Used, "Test wallet/Account Endpoint2", "endpoint2Signer should have signed Account Endpoint2") + } else { + // Account Endpoint1 assigned to endpoint2 (port 12346) + require.Len(t, endpoint1Used, 0, "endpoint1Signer should not have been called") + require.Len(t, endpoint2Used, 2, "endpoint2Signer should have been called twice") + require.Contains(t, endpoint2Used, "Test wallet/Account Endpoint1", "endpoint2Signer should have signed Account Endpoint1") + require.Contains(t, endpoint2Used, "Test wallet/Account Endpoint2", "endpoint2Signer should have signed Account Endpoint2") + } // Verify that signing works with whatever endpoint the shared account got assigned to // This demonstrates that endpoint reassignment works correctly and accounts use their assigned endpoint for signing diff --git a/mock/signerserver.go b/mock/signerserver.go index 71162f2..872c9c7 100644 --- a/mock/signerserver.go +++ b/mock/signerserver.go @@ -16,6 +16,8 @@ package dirk import ( "context" + "fmt" + "strings" "sync" pb "github.com/wealdtech/eth2-signer-api/pb/v1" @@ -25,15 +27,51 @@ import ( type MockSignerServer struct { pb.UnimplementedSignerServer - endpointsUsed []string - mutex sync.Mutex + endpointsUsed []string + mutex sync.Mutex + allowedAccounts map[string]bool // set of allowed account names (without wallet prefix) +} + +// NewMockSignerServer creates a new mock signer server. +func NewMockSignerServer() *MockSignerServer { + return &MockSignerServer{ + endpointsUsed: make([]string, 0), + allowedAccounts: make(map[string]bool), + } +} + +// NewMockSignerServerWithAccounts creates a new mock signer server that only accepts specific accounts. +func NewMockSignerServerWithAccounts(allowedAccounts []string) *MockSignerServer { + server := NewMockSignerServer() + for _, account := range allowedAccounts { + server.allowedAccounts[account] = true + } + + return server } func (s *MockSignerServer) Sign(_ context.Context, req *pb.SignRequest) (*pb.SignResponse, error) { s.mutex.Lock() defer s.mutex.Unlock() - // Extract account info from the request + // Extract account name from path (remove wallet prefix) + accountPath := req.GetAccount() + // Account paths are like "Test wallet/Account Endpoint1", we want just "Account Endpoint1" + var accountName string + if lastSlash := strings.LastIndex(accountPath, "/"); lastSlash >= 0 { + accountName = accountPath[lastSlash+1:] + } else { + accountName = accountPath + } + + // Check if this account is allowed on this signer server + if len(s.allowedAccounts) > 0 && !s.allowedAccounts[accountName] { + return &pb.SignResponse{ + State: pb.ResponseState_FAILED, + }, fmt.Errorf("account %s not available on this endpoint", accountName) + } + + // Record the signing request s.endpointsUsed = append(s.endpointsUsed, req.GetAccount()) return &pb.SignResponse{ diff --git a/wallet_internal_test.go b/wallet_internal_test.go index 36c36e0..3ca624e 100644 --- a/wallet_internal_test.go +++ b/wallet_internal_test.go @@ -185,4 +185,4 @@ func TestAccountByID(t *testing.T) { _, err = w.(e2wtypes.WalletAccountByIDProvider).AccountByID(ctx, uuid.MustParse("00000000-0000-0000-0000-000000000002")) require.EqualError(t, err, "not supported") -} \ No newline at end of file +} From 7a895c1e37baea963d9d2b658f9d4d20b4447b98 Mon Sep 17 00:00:00 2001 From: AntiD2ta Date: Thu, 8 Jan 2026 11:07:54 +0100 Subject: [PATCH 5/6] Update golangci-lint CI workflow to use version 2 --- .github/workflows/golangci-lint.yml | 12 +++++++----- 1 file changed, 7 insertions(+), 5 deletions(-) diff --git a/.github/workflows/golangci-lint.yml b/.github/workflows/golangci-lint.yml index 2af114b..b41c857 100644 --- a/.github/workflows/golangci-lint.yml +++ b/.github/workflows/golangci-lint.yml @@ -2,11 +2,12 @@ name: golangci-lint on: push: branches: - - master + - master pull_request: permissions: contents: read + pull-requests: read jobs: golangci: @@ -15,10 +16,11 @@ jobs: steps: - uses: actions/setup-go@v5 with: - go-version: "^1.22" + cache: false + go-version: '1.22.4' - uses: actions/checkout@v4 - name: golangci-lint - uses: golangci/golangci-lint-action@v6 + uses: golangci/golangci-lint-action@v8 with: - version: latest - args: --timeout=60m + only-new-issues: true + args: --timeout=10m From a1767900c719c3becba6d1569c92a9dba12ae876 Mon Sep 17 00:00:00 2001 From: AntiD2ta Date: Thu, 8 Jan 2026 14:45:03 +0100 Subject: [PATCH 6/6] Update license header --- mock/signerserver.go | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/mock/signerserver.go b/mock/signerserver.go index 872c9c7..076aaa1 100644 --- a/mock/signerserver.go +++ b/mock/signerserver.go @@ -1,4 +1,4 @@ -// Copyright © 2020, 2021 Weald Technology Trading. +// Copyright © 2026 Weald Technology Trading. // Licensed under the Apache License, Version 2.0 (the "License"); // you may not use this file except in compliance with the License. // You may obtain a copy of the License at