From 1e4524fe01e9169f13c13bebbd1205bd60e752c1 Mon Sep 17 00:00:00 2001 From: Gareth George Date: Thu, 24 Jul 2025 22:04:44 -0700 Subject: [PATCH] initial test coverage --- internal/api/syncapi/authmiddleware.go | 245 +++++++++++++++ internal/api/syncapi/authmiddleware_test.go | 198 +++++++++++++ internal/api/syncapi/client.go | 18 ++ internal/api/syncapi/cmdstreamutil.go | 111 ------- internal/api/syncapi/errors.go | 20 +- internal/api/syncapi/handler.go | 280 ++++++++++++++++++ .../api/syncapi/{ => old}/peerstate_test.go | 0 internal/api/syncapi/{ => old}/syncclient.go | 7 +- internal/api/syncapi/{ => old}/synccommon.go | 0 internal/api/syncapi/{ => old}/synchandler.go | 9 +- internal/api/syncapi/peerstate.go | 39 ++- .../{syncapi_test.go => syncapi_test.nogo} | 19 +- .../{syncmanager.go => syncmanager.nogo} | 5 +- ...cstatehandler.go => syncstatehandler.nogo} | 11 +- 14 files changed, 798 insertions(+), 164 deletions(-) create mode 100644 internal/api/syncapi/authmiddleware.go create mode 100644 internal/api/syncapi/authmiddleware_test.go create mode 100644 internal/api/syncapi/client.go delete mode 100644 internal/api/syncapi/cmdstreamutil.go create mode 100644 internal/api/syncapi/handler.go rename internal/api/syncapi/{ => old}/peerstate_test.go (100%) rename internal/api/syncapi/{ => old}/syncclient.go (98%) rename internal/api/syncapi/{ => old}/synccommon.go (100%) rename internal/api/syncapi/{ => old}/synchandler.go (98%) rename internal/api/syncapi/{syncapi_test.go => syncapi_test.nogo} (97%) rename internal/api/syncapi/{syncmanager.go => syncmanager.nogo} (97%) rename internal/api/syncapi/{syncstatehandler.go => syncstatehandler.nogo} (81%) diff --git a/internal/api/syncapi/authmiddleware.go b/internal/api/syncapi/authmiddleware.go new file mode 100644 index 00000000..2070297d --- /dev/null +++ b/internal/api/syncapi/authmiddleware.go @@ -0,0 +1,245 @@ +package syncapi + +import ( + "context" + "encoding/base64" + "encoding/binary" + "errors" + "fmt" + "net/http" + "slices" + "time" + + "connectrpc.com/connect" + v1 "github.com/garethgeorge/backrest/gen/go/v1" + "github.com/garethgeorge/backrest/gen/go/v1sync" + "github.com/garethgeorge/backrest/internal/config" + "github.com/garethgeorge/backrest/internal/cryptoutil" + "google.golang.org/protobuf/proto" +) + +var authTokenHeader = "Authorization" +var maxSignatureAge = 5 * time.Minute // Maximum age of a signature before it is considered invalid + +type peerContextKey string + +const PeerContextKey peerContextKey = "peer" + +func ContextWithPeer(ctx context.Context, peer *v1.Multihost_Peer) context.Context { + return context.WithValue(ctx, PeerContextKey, peer) +} + +func PeerFromContext(ctx context.Context) *v1.Multihost_Peer { + peer, ok := ctx.Value(PeerContextKey).(*v1.Multihost_Peer) + if !ok { + return nil + } + return peer +} + +func newAuthHandler(config *config.ConfigManager, next http.Handler) http.Handler { + return http.HandlerFunc(func(rw http.ResponseWriter, r *http.Request) { + config, err := config.Get() + if err != nil { + http.Error(rw, "internal error", http.StatusInternalServerError) + return + } + + authHeaderValue, err := createAuthHeader(config) + if err != nil { + http.Error(rw, fmt.Sprintf("internal error: %v", err), http.StatusInternalServerError) + return + } + rw.Header().Set(authTokenHeader, authHeaderValue) + + peer, err := decodeAndVerifyAuthHeader(r, config.Instance, config.GetMultihost().GetAuthorizedClients()) + if err != nil { + http.Error(rw, fmt.Sprintf("unauthorized: %v", err), http.StatusUnauthorized) + return + } + + next.ServeHTTP(rw, r.WithContext(context.WithValue(r.Context(), PeerContextKey, peer))) + }) +} + +func createAuthHeader(config *v1.Config) (string, error) { + if config == nil || config.GetMultihost().GetIdentity() == nil { + return "", errors.New("config missing multihost.identity") + } + + privKey, err := cryptoutil.NewPrivateKey(config.GetMultihost().GetIdentity()) + if err != nil { + return "", fmt.Errorf("load private key: %w", err) + } + + signedMessage, err := createSignedMessage([]byte(config.Instance), privKey) + if err != nil { + return "", fmt.Errorf("create signed message: %w", err) + } + + authToken := &v1sync.AuthorizationToken{ + InstanceId: signedMessage, + PublicKey: privKey.PublicKeyProto(), + } + + tokenBytes, err := proto.Marshal(authToken) + if err != nil { + return "", fmt.Errorf("marshal auth token: %w", err) + } + + return base64.StdEncoding.EncodeToString(tokenBytes), nil +} + +type authHeaderClient struct { + configManager *config.ConfigManager + delegate connect.HTTPClient + wantPeer *v1.Multihost_Peer +} + +func (c *authHeaderClient) Do(req *http.Request) (*http.Response, error) { + // create the header + cfg, err := c.configManager.Get() + if err != nil { + return nil, fmt.Errorf("get config: %w", err) + } + authHeaderValue, err := createAuthHeader(cfg) + if err != nil { + return nil, fmt.Errorf("create auth header: %w", err) + } + req.Header.Set(authTokenHeader, authHeaderValue) + + resp, err := c.delegate.Do(req) + // verify the response header + if err != nil { + return nil, fmt.Errorf("HTTP request failed: %w", err) + } + if resp.StatusCode != http.StatusOK { + return resp, fmt.Errorf("HTTP request failed with status %d: %s", resp.StatusCode, resp.Status) + } + peer, err := decodeAndVerifyAuthHeader(req, cfg.Instance, cfg.GetMultihost().GetAuthorizedClients()) + if err != nil { + return resp, fmt.Errorf("verify auth header: %w", err) + } + + // Check the peer matches the expected one. + if c.wantPeer == nil || c.wantPeer.GetInstanceId() != peer.GetInstanceId() { + return resp, fmt.Errorf("peer instance ID mismatch: expected %s, got %s", c.wantPeer.GetInstanceId(), peer.GetInstanceId()) + } + if c.wantPeer.GetKeyid() != peer.GetKeyid() { + return resp, fmt.Errorf("peer key ID mismatch: expected %s, got %s", c.wantPeer.GetKeyid(), peer.GetKeyid()) + } + return resp, nil +} + +func newHTTPClientWithConfig(cfg *config.ConfigManager, delegate connect.HTTPClient) (connect.HTTPClient, error) { + return &authHeaderClient{ + configManager: cfg, + delegate: delegate, + }, nil +} + +func decodeAndVerifyAuthHeader(r *http.Request, localInstanceID string, peers []*v1.Multihost_Peer) (*v1.Multihost_Peer, error) { + authHeader := r.Header.Get(authTokenHeader) + if len(authHeader) == 0 { + return nil, errors.New("missing authorization header") + } + + // Decode the auth token from the header + tokenBytes, err := base64.StdEncoding.DecodeString(authHeader) + if err != nil { + return nil, errors.New("invalid authorization header format") + } + + var token v1sync.AuthorizationToken + if err := proto.Unmarshal(tokenBytes, &token); err != nil { + return nil, fmt.Errorf("unmarshal authorization token: %w", err) + } + + // Load the public key from the token + publicKey, err := cryptoutil.NewPublicKey(token.GetPublicKey()) + if err != nil { + return nil, fmt.Errorf("load public key: %w", err) + } + if publicKey.KeyID() != token.InstanceId.GetKeyid() { + return nil, fmt.Errorf("instance ID must be signed with public key in token: expected %s, got %s", token.InstanceId.GetKeyid(), publicKey.KeyID()) + } + + // Verify the signed message + if err := verifySignedMessage(token.GetInstanceId(), publicKey); err != nil { + return nil, fmt.Errorf("verify signed message: %w", err) + } + + // Now that we've validated that the peer was able to sign the message, we can look it up in the config + peerIdx := slices.IndexFunc(peers, func(peer *v1.Multihost_Peer) bool { + return peer.Keyid == publicKey.KeyID() + }) + if peerIdx == -1 { + return nil, fmt.Errorf("peer with key ID %s not found in authorized clients", publicKey.KeyID()) + } + + // Finally check that the instance ID in the token matches the one in the config + peer := peers[peerIdx] + tokenInstanceID := string(token.GetInstanceId().GetPayload()) + if peer.InstanceId != tokenInstanceID { + return nil, fmt.Errorf("instance ID mismatch: expected %s, got %s", peer.InstanceId, tokenInstanceID) + } + + return peer, nil +} + +func createSignedMessage(payload []byte, identity *cryptoutil.PrivateKey) (*v1.SignedMessage, error) { + if len(payload) == 0 { + return nil, errors.New("payload must not be empty") + } + + timestampMillis := time.Now().UnixMilli() + + payloadWithTimestamp := make([]byte, 0, len(payload)+8) + binary.BigEndian.AppendUint64(payloadWithTimestamp, uint64(timestampMillis)) + payloadWithTimestamp = append(payloadWithTimestamp, payload...) + + signature, err := identity.Sign(payloadWithTimestamp) + if err != nil { + return nil, fmt.Errorf("signing payload: %w", err) + } + + return &v1.SignedMessage{ + Payload: payload, + Signature: signature, + Keyid: identity.KeyID(), + TimestampMillis: timestampMillis, + }, nil +} + +func verifySignedMessage(msg *v1.SignedMessage, publicKey *cryptoutil.PublicKey) error { + if msg == nil { + return errors.New("signed message must not be nil") + } + if len(msg.GetPayload()) == 0 { + return errors.New("signed message payload must not be empty") + } + if len(msg.GetSignature()) == 0 { + return errors.New("signed message signature must not be empty") + } + if len(msg.GetKeyid()) == 0 { + return errors.New("signed message key ID must not be empty") + } + + if publicKey.KeyID() != msg.GetKeyid() { + return fmt.Errorf("public key ID mismatch: expected %s, got %s", publicKey.KeyID(), msg.GetKeyid()) + } + + payloadWithTimestamp := make([]byte, 0, len(msg.GetPayload())+8) + binary.BigEndian.AppendUint64(payloadWithTimestamp, uint64(msg.GetTimestampMillis())) + payloadWithTimestamp = append(payloadWithTimestamp, msg.GetPayload()...) + + if err := publicKey.Verify(payloadWithTimestamp, msg.GetSignature()); err != nil { + return fmt.Errorf("verifying signed message: %w", err) + } + + if time.Since(time.UnixMilli(msg.GetTimestampMillis())) > maxSignatureAge { + return fmt.Errorf("signature is too old, max age is %s. Is the clock out of sync?", maxSignatureAge) + } + + return nil +} diff --git a/internal/api/syncapi/authmiddleware_test.go b/internal/api/syncapi/authmiddleware_test.go new file mode 100644 index 00000000..6b775155 --- /dev/null +++ b/internal/api/syncapi/authmiddleware_test.go @@ -0,0 +1,198 @@ +package syncapi + +import ( + "encoding/base64" + "fmt" + "net/http" + "net/http/httptest" + "testing" + "time" + + v1 "github.com/garethgeorge/backrest/gen/go/v1" + "github.com/garethgeorge/backrest/gen/go/v1sync" + "github.com/garethgeorge/backrest/internal/config" + "github.com/garethgeorge/backrest/internal/cryptoutil" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "google.golang.org/protobuf/proto" +) + +func TestAuthMiddleware(t *testing.T) { + serverPrivKey, err := cryptoutil.GeneratePrivateKey() + require.NoError(t, err) + + clientPrivKey, err := cryptoutil.GeneratePrivateKey() + require.NoError(t, err) + + // Create a mock config manager + cfgManager := &config.ConfigManager{ + Store: &config.MemoryStore{ + Config: &v1.Config{ + Instance: "test-instance", + Multihost: &v1.Multihost{ + Identity: serverPrivKey, + AuthorizedClients: []*v1.Multihost_Peer{ + { + InstanceId: "client-instance", + Keyid: clientPrivKey.Keyid, + }, + }, + }, + }, + }, + } + + // Create a mock handler + mockHandler := http.HandlerFunc(func(rw http.ResponseWriter, r *http.Request) { + peer := PeerFromContext(r.Context()) + require.NotNil(t, peer) + assert.Equal(t, "client-instance", peer.InstanceId) + rw.WriteHeader(http.StatusOK) + }) + + // Create the auth handler + authHandler := newAuthHandler(cfgManager, mockHandler) + + // Create a test server + server := httptest.NewServer(authHandler) + defer server.Close() + + t.Run("valid auth header", func(t *testing.T) { + // Create a request with a valid auth header + req, err := http.NewRequest("GET", server.URL, nil) + require.NoError(t, err) + + // Create a valid auth header + clientCfg := &v1.Config{ + Instance: "client-instance", + Multihost: &v1.Multihost{ + Identity: clientPrivKey, + }, + } + authHeader, err := createAuthHeader(clientCfg) + require.NoError(t, err) + req.Header.Set(authTokenHeader, authHeader) + + // Make the request + resp, err := http.DefaultClient.Do(req) + require.NoError(t, err) + defer resp.Body.Close() + + // Check the response + assert.Equal(t, http.StatusOK, resp.StatusCode) + }) + + t.Run("missing auth header", func(t *testing.T) { + req, err := http.NewRequest("GET", server.URL, nil) + require.NoError(t, err) + + resp, err := http.DefaultClient.Do(req) + require.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusUnauthorized, resp.StatusCode) + }) + + t.Run("invalid auth header", func(t *testing.T) { + req, err := http.NewRequest("GET", server.URL, nil) + require.NoError(t, err) + req.Header.Set(authTokenHeader, "invalid") + + resp, err := http.DefaultClient.Do(req) + require.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusUnauthorized, resp.StatusCode) + }) + + t.Run("unauthorized peer", func(t *testing.T) { + unauthorizedPrivKey, err := cryptoutil.GeneratePrivateKey() + require.NoError(t, err) + + req, err := http.NewRequest("GET", server.URL, nil) + require.NoError(t, err) + + clientCfg := &v1.Config{ + Instance: "unauthorized-instance", + Multihost: &v1.Multihost{ + Identity: unauthorizedPrivKey, + }, + } + authHeader, err := createAuthHeader(clientCfg) + require.NoError(t, err) + req.Header.Set(authTokenHeader, authHeader) + + resp, err := http.DefaultClient.Do(req) + require.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusUnauthorized, resp.StatusCode) + }) + + t.Run("instance id mismatch", func(t *testing.T) { + req, err := http.NewRequest("GET", server.URL, nil) + require.NoError(t, err) + + clientCfg := &v1.Config{ + Instance: "wrong-instance", + Multihost: &v1.Multihost{ + Identity: clientPrivKey, + }, + } + authHeader, err := createAuthHeader(clientCfg) + require.NoError(t, err) + req.Header.Set(authTokenHeader, authHeader) + + resp, err := http.DefaultClient.Do(req) + require.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusUnauthorized, resp.StatusCode) + }) + + t.Run("signature too old", func(t *testing.T) { + req, err := http.NewRequest("GET", server.URL, nil) + require.NoError(t, err) + + clientCfg := &v1.Config{ + Instance: "client-instance", + Multihost: &v1.Multihost{ + Identity: clientPrivKey, + }, + } + + privKey, err := cryptoutil.NewPrivateKey(clientPrivKey) + require.NoError(t, err) + + // create a signed message with an old timestamp + signedMessage, err := createSignedMessage([]byte(clientCfg.Instance), privKey) + require.NoError(t, err) + signedMessage.TimestampMillis = time.Now().Add(-2 * maxSignatureAge).UnixMilli() + + // create the auth token + authToken, err := createAuthToken(signedMessage, privKey) + require.NoError(t, err) + + req.Header.Set(authTokenHeader, authToken) + + resp, err := http.DefaultClient.Do(req) + require.NoError(t, err) + defer resp.Body.Close() + + assert.Equal(t, http.StatusUnauthorized, resp.StatusCode) + }) +} + +func createAuthToken(signedMessage *v1.SignedMessage, privKey *cryptoutil.PrivateKey) (string, error) { + authToken := &v1sync.AuthorizationToken{ + InstanceId: signedMessage, + PublicKey: privKey.PublicKeyProto(), + } + + tokenBytes, err := proto.Marshal(authToken) + if err != nil { + return "", fmt.Errorf("marshal auth token: %w", err) + } + + return base64.StdEncoding.EncodeToString(tokenBytes), nil +} diff --git a/internal/api/syncapi/client.go b/internal/api/syncapi/client.go new file mode 100644 index 00000000..6e444134 --- /dev/null +++ b/internal/api/syncapi/client.go @@ -0,0 +1,18 @@ +package syncapi + +import ( + v1 "github.com/garethgeorge/backrest/gen/go/v1" + "github.com/garethgeorge/backrest/gen/go/v1sync/v1syncconnect" +) + +type syncClient struct { + client *v1syncconnect.SyncPeerServiceClient + peer *v1.Multihost_Peer +} + +func newSyncClient(client *v1syncconnect.SyncPeerServiceClient, peer *v1.Multihost_Peer) *syncClient { + return &syncClient{ + client: client, + peer: peer, + } +} diff --git a/internal/api/syncapi/cmdstreamutil.go b/internal/api/syncapi/cmdstreamutil.go deleted file mode 100644 index c90ab653..00000000 --- a/internal/api/syncapi/cmdstreamutil.go +++ /dev/null @@ -1,111 +0,0 @@ -package syncapi - -import ( - "context" - "errors" - "fmt" - "io" - "time" - - "connectrpc.com/connect" - v1 "github.com/garethgeorge/backrest/gen/go/v1" -) - -type syncCommandStreamTrait interface { - Send(item *v1.SyncStreamItem) error - Receive() (*v1.SyncStreamItem, error) -} - -var _ syncCommandStreamTrait = (*connect.BidiStream[v1.SyncStreamItem, v1.SyncStreamItem])(nil) // Ensure that connect.BidiStream implements syncCommandStreamTrait -var _ syncCommandStreamTrait = (*connect.BidiStreamForClient[v1.SyncStreamItem, v1.SyncStreamItem])(nil) // Ensure that connect.BidiStreamForClient implements syncCommandStreamTrait - -type bidiSyncCommandStream struct { - sendChan chan *v1.SyncStreamItem - recvChan chan *v1.SyncStreamItem - terminateWithErrChan chan error -} - -func newBidiSyncCommandStream() *bidiSyncCommandStream { - return &bidiSyncCommandStream{ - sendChan: make(chan *v1.SyncStreamItem, 64), // Buffered channel to allow sending items without blocking - recvChan: make(chan *v1.SyncStreamItem, 1), - terminateWithErrChan: make(chan error, 1), - } -} - -func (s *bidiSyncCommandStream) Send(item *v1.SyncStreamItem) { - select { - case s.sendChan <- item: - default: - // Try again with a timeout, if it fails, send an error to terminate the stream - select { - case s.sendChan <- item: - case <-time.After(100 * time.Millisecond): - s.SendErrorAndTerminate(NewSyncErrorDisconnected(errors.New("send channel is full, cannot send item"))) - } - } -} - -// SendErrorAndTerminate sends an error to the termination channel. -// If the error is nil, it terminates only. -func (s *bidiSyncCommandStream) SendErrorAndTerminate(err error) { - if err == nil { - return - } - select { - case s.terminateWithErrChan <- err: - default: - // If the channel is full, we can't send the error, so we just ignore it. - // This is a best-effort termination. - } -} - -func (s *bidiSyncCommandStream) ReadChannel() chan *v1.SyncStreamItem { - return s.recvChan -} - -func (s *bidiSyncCommandStream) ReceiveWithinDuration(d time.Duration) *v1.SyncStreamItem { - select { - case item := <-s.recvChan: - return item - case <-time.After(d): - return nil // Return nil if no item is received within the duration - } -} - -func (s *bidiSyncCommandStream) ConnectStream(ctx context.Context, stream syncCommandStreamTrait) error { - ctx, cancel := context.WithCancel(ctx) - defer cancel() - go func() { - for ctx.Err() == nil { - if val, err := stream.Receive(); err != nil { - s.SendErrorAndTerminate(NewSyncErrorDisconnected(fmt.Errorf("receiving item: %w", err))) - break - } else { - s.recvChan <- val - } - } - close(s.recvChan) - }() - - for { - select { - case item := <-s.sendChan: - if item == nil { - continue - } - if err := stream.Send(item); err != nil { - if errors.Is(err, io.EOF) { - err = fmt.Errorf("connection failed or dropped: %w", err) - } - s.SendErrorAndTerminate(err) - return err - } - case err := <-s.terminateWithErrChan: - return err // Terminate the stream with the error or nil if no error was sent - case <-ctx.Done(): - // Context is done, we should stop processing. - return ctx.Err() - } - } -} diff --git a/internal/api/syncapi/errors.go b/internal/api/syncapi/errors.go index d847a7de..85530b3d 100644 --- a/internal/api/syncapi/errors.go +++ b/internal/api/syncapi/errors.go @@ -3,11 +3,11 @@ package syncapi import ( "fmt" - v1 "github.com/garethgeorge/backrest/gen/go/v1" + "github.com/garethgeorge/backrest/gen/go/v1sync" ) type SyncError struct { - State v1.SyncConnectionState + State v1sync.ConnectionState Message error } @@ -23,56 +23,56 @@ func (e *SyncError) Unwrap() error { func NewSyncErrorDisconnected(message error) *SyncError { return &SyncError{ - State: v1.SyncConnectionState_CONNECTION_STATE_DISCONNECTED, + State: v1sync.ConnectionState_CONNECTION_STATE_DISCONNECTED, Message: message, } } func NewSyncErrorUnknown(message error) *SyncError { return &SyncError{ - State: v1.SyncConnectionState_CONNECTION_STATE_UNKNOWN, + State: v1sync.ConnectionState_CONNECTION_STATE_UNKNOWN, Message: message, } } func NewSyncErrorPending(message error) *SyncError { return &SyncError{ - State: v1.SyncConnectionState_CONNECTION_STATE_PENDING, + State: v1sync.ConnectionState_CONNECTION_STATE_PENDING, Message: message, } } func NewSyncErrorConnected(message error) *SyncError { return &SyncError{ - State: v1.SyncConnectionState_CONNECTION_STATE_CONNECTED, + State: v1sync.ConnectionState_CONNECTION_STATE_CONNECTED, Message: message, } } func NewSyncErrorRetryWait(message error) *SyncError { return &SyncError{ - State: v1.SyncConnectionState_CONNECTION_STATE_RETRY_WAIT, + State: v1sync.ConnectionState_CONNECTION_STATE_RETRY_WAIT, Message: message, } } func NewSyncErrorAuth(message error) *SyncError { return &SyncError{ - State: v1.SyncConnectionState_CONNECTION_STATE_ERROR_AUTH, + State: v1sync.ConnectionState_CONNECTION_STATE_ERROR_AUTH, Message: message, } } func NewSyncErrorProtocol(message error) *SyncError { return &SyncError{ - State: v1.SyncConnectionState_CONNECTION_STATE_ERROR_PROTOCOL, + State: v1sync.ConnectionState_CONNECTION_STATE_ERROR_PROTOCOL, Message: message, } } func NewSyncErrorInternal(message error) *SyncError { return &SyncError{ - State: v1.SyncConnectionState_CONNECTION_STATE_ERROR_INTERNAL, + State: v1sync.ConnectionState_CONNECTION_STATE_ERROR_INTERNAL, Message: message, } } diff --git a/internal/api/syncapi/handler.go b/internal/api/syncapi/handler.go new file mode 100644 index 00000000..2629e230 --- /dev/null +++ b/internal/api/syncapi/handler.go @@ -0,0 +1,280 @@ +package syncapi + +import ( + "context" + "errors" + "fmt" + "io" + "sync" + "time" + "unique" + + "connectrpc.com/connect" + "github.com/garethgeorge/backrest/gen/go/types" + v1 "github.com/garethgeorge/backrest/gen/go/v1" + "github.com/garethgeorge/backrest/gen/go/v1sync" + "github.com/garethgeorge/backrest/gen/go/v1sync/v1syncconnect" + "github.com/garethgeorge/backrest/internal/logstore" + "github.com/garethgeorge/backrest/internal/oplog" + "github.com/garethgeorge/backrest/internal/protoutil" + lru "github.com/hashicorp/golang-lru/v2" + "go.uber.org/zap" + "google.golang.org/protobuf/proto" + "google.golang.org/protobuf/types/known/emptypb" +) + +type opIdCacheKey struct { + OriginalInstanceKeyid unique.Handle[string] + ID int64 +} + +// syncHandler provides server component functionality for the sync API. +type syncHandler struct { + v1syncconnect.UnimplementedSyncPeerServiceHandler + oplog *oplog.OpLog + logStore *logstore.LogStore + + opCacheMu sync.Mutex + flowIDLru *lru.Cache[opIdCacheKey, int64] + opIDLru *lru.Cache[opIdCacheKey, int64] + + logger *zap.Logger +} + +func NewSyncHandler(oplog *oplog.OpLog, logStore *logstore.LogStore) *syncHandler { + // Both caches want to be reasonably large to avoid db lookups. + flowIDLru, _ := lru.New[opIdCacheKey, int64](4 * 1024) + opIDLru, _ := lru.New[opIdCacheKey, int64](16 * 1024) + + return &syncHandler{ + oplog: oplog, + logStore: logStore, + flowIDLru: flowIDLru, + opIDLru: opIDLru, + logger: zap.NewNop(), + } +} + +var _ v1syncconnect.SyncPeerServiceHandler = (*syncHandler)(nil) + +// translateSingleID translates a single ID (either opID or flowID) using the provided cache and query +func (sh *syncHandler) translateSingleID( + originalInstanceKeyid string, + originalID int64, + cache *lru.Cache[opIdCacheKey, int64], + query oplog.Query, +) (int64, error) { + if originalID == 0 { + return 0, nil + } + + cacheKey := opIdCacheKey{ + OriginalInstanceKeyid: unique.Make(originalInstanceKeyid), + ID: originalID, + } + + // Check cache first + if translatedID, ok := cache.Get(cacheKey); ok { + return translatedID, nil + } + + // Cache miss - query the database + op, err := sh.oplog.FindOneMetadata(query) + if err != nil { + if errors.Is(err, oplog.ErrNoResults) { + return 0, nil // No results means the ID is not found + } + return 0, err // Other errors should be propagated + } + + // Cache the result and return + translatedID := op.FlowID + cache.Add(cacheKey, translatedID) + return translatedID, nil +} + +func (sh *syncHandler) translateOpIdAndFlowID(originalInstanceKeyid string, originalOpId int64, originalFlowId int64) (int64, int64, error) { + sh.opCacheMu.Lock() + defer sh.opCacheMu.Unlock() + + // Translate opID + opID, err := sh.translateSingleID( + originalInstanceKeyid, + originalOpId, + sh.opIDLru, + oplog.Query{ + OriginalInstanceKeyid: &originalInstanceKeyid, + OriginalID: &originalOpId, + }, + ) + if err != nil { + return 0, 0, err + } + + // Translate flowID + flowID, err := sh.translateSingleID( + originalInstanceKeyid, + originalFlowId, + sh.flowIDLru, + oplog.Query{ + OriginalInstanceKeyid: &originalInstanceKeyid, + OriginalFlowID: &originalFlowId, + }, + ) + if err != nil { + return 0, 0, err + } + + return opID, flowID, nil +} + +func (sh *syncHandler) GetOperationMetadata(ctx context.Context, req *connect.Request[v1.OpSelector]) (*connect.Response[v1sync.GetOperationMetadataResponse], error) { + peer := PeerFromContext(ctx) + if peer == nil { + return nil, connect.NewError(connect.CodeUnauthenticated, errors.New("no peer found in context")) + } + + // Check if the peer can read the selector + if req.Msg.OriginalInstanceKeyid == nil || *req.Msg.OriginalInstanceKeyid != peer.Keyid { + return nil, connect.NewError(connect.CodePermissionDenied, errors.New("GetOperationMetadata: peer must specify original instance keyid")) + } + + sel, err := protoutil.OpSelectorToQuery(req.Msg) + if err != nil { + return nil, connect.NewError(connect.CodeInvalidArgument, err) + } + + var opIDs []int64 + var modNos []int64 + + if err := sh.oplog.QueryMetadata(sel, func(op oplog.OpMetadata) error { + opIDs = append(opIDs, op.OriginalID) + modNos = append(modNos, op.Modno) + return nil + }); err != nil { + return nil, connect.NewError(connect.CodeInternal, err) + } + + return connect.NewResponse(&v1sync.GetOperationMetadataResponse{ + OpIds: opIDs, + Modnos: modNos, + }), nil +} + +func (sh *syncHandler) SendOperations(ctx context.Context, stream *connect.ClientStream[v1.Operation]) (*connect.Response[emptypb.Empty], error) { + peer := PeerFromContext(ctx) + if peer == nil { + return nil, connect.NewError(connect.CodeUnauthenticated, errors.New("no peer found in context")) + } + + for stream.Receive() { + op := stream.Msg() + if op == nil { + return nil, connect.NewError(connect.CodeInvalidArgument, errors.New("received nil operation")) + } + + id, flowID, err := sh.translateOpIdAndFlowID(peer.Keyid, op.Id, op.FlowId) + if err != nil { + return nil, connect.NewError(connect.CodeInternal, err) + } + + // Update the operation with the translated IDs + opCopy := proto.Clone(op).(*v1.Operation) + opCopy.Id = id + opCopy.FlowId = flowID + + // Set the operation in the oplog + if err := sh.oplog.Set(opCopy); err != nil { + return nil, connect.NewError(connect.CodeInternal, fmt.Errorf("set operation: %w", err)) + } + } + + return nil, connect.NewError(connect.CodeUnimplemented, errors.New("v1sync.SyncPeerService.SendOperations is not implemented")) +} + +func (sh *syncHandler) GetLog(ctx context.Context, req *connect.Request[types.StringValue], stream *connect.ServerStream[v1sync.LogDataEntry]) error { + peer := PeerFromContext(ctx) + if peer == nil { + return connect.NewError(connect.CodeUnauthenticated, errors.New("no peer found in context")) + } + logID := req.Msg.Value + + metadata, err := sh.logStore.GetMetadata(logID) + if err != nil { + if errors.Is(err, logstore.ErrLogNotFound) { + return connect.NewError(connect.CodeNotFound, fmt.Errorf("log with ID %s not found", logID)) + } + return connect.NewError(connect.CodeInternal, fmt.Errorf("get log metadata: %w", err)) + } + + log, err := sh.logStore.Open(logID) + if err != nil { + return connect.NewError(connect.CodeInternal, fmt.Errorf("get log: %w", err)) + } else if log == nil { + return connect.NewError(connect.CodeNotFound, fmt.Errorf("log with ID %s not found", logID)) + } + defer log.Close() + + // Send first entry with log ID and owner operation ID + entry := &v1sync.LogDataEntry{ + LogId: logID, + OwnerOpid: metadata.OwnerOpID, + } + if metadata.ExpirationTime != (time.Time{}) { + entry.ExpirationTsUnix = metadata.ExpirationTime.Unix() + } + if err := stream.Send(entry); err != nil { + if errors.Is(err, io.EOF) { + return nil // Client closed the stream + } + return connect.NewError(connect.CodeInternal, fmt.Errorf("send log entry: %w", err)) + } + + // Read the log in chunks and send each chunk as a LogDataEntry + buf := make([]byte, 0, 32*1024) + for { + n, err := log.Read(buf) + if err != nil { + if errors.Is(err, io.EOF) { + break // End of log + } + return connect.NewError(connect.CodeInternal, fmt.Errorf("read log: %w", err)) + } + if n == 0 { + break + } + bytes := buf[:n] + entry := &v1sync.LogDataEntry{ + Chunk: bytes, + } + if err := stream.Send(entry); err != nil { + if errors.Is(err, io.EOF) { + break // Client closed the stream + } + return connect.NewError(connect.CodeInternal, fmt.Errorf("send log entry: %w", err)) + } + } + if err := log.Close(); err != nil { + return connect.NewError(connect.CodeInternal, fmt.Errorf("close log: %w", err)) + } + return nil +} + +func (sh *syncHandler) SetAvailableResources(ctx context.Context, req *connect.Request[v1sync.SetAvailableResourcesRequest]) (*connect.Response[emptypb.Empty], error) { + peer := PeerFromContext(ctx) + if peer == nil { + return nil, connect.NewError(connect.CodeUnauthenticated, errors.New("no peer found in context")) + } + return nil, connect.NewError(connect.CodeUnimplemented, errors.New("v1sync.SyncPeerService.SetAvailableResources is not implemented")) +} + +func (sh *syncHandler) SetConfig(context.Context, *connect.Request[v1sync.SetConfigRequest]) (*connect.Response[emptypb.Empty], error) { + return nil, connect.NewError(connect.CodeUnimplemented, errors.New("v1sync.SyncPeerService.SetConfig is not implemented")) +} + +func (sh *syncHandler) GetConfig(context.Context, *connect.Request[emptypb.Empty]) (*connect.Response[v1sync.RemoteConfig], error) { + return nil, connect.NewError(connect.CodeUnimplemented, errors.New("v1sync.SyncPeerService.GetConfig is not implemented")) +} + +type syncHandlerClient struct { +} diff --git a/internal/api/syncapi/peerstate_test.go b/internal/api/syncapi/old/peerstate_test.go similarity index 100% rename from internal/api/syncapi/peerstate_test.go rename to internal/api/syncapi/old/peerstate_test.go diff --git a/internal/api/syncapi/syncclient.go b/internal/api/syncapi/old/syncclient.go similarity index 98% rename from internal/api/syncapi/syncclient.go rename to internal/api/syncapi/old/syncclient.go index 6083d80f..5dabe35b 100644 --- a/internal/api/syncapi/syncclient.go +++ b/internal/api/syncapi/old/syncclient.go @@ -14,6 +14,7 @@ import ( "github.com/garethgeorge/backrest/gen/go/types" v1 "github.com/garethgeorge/backrest/gen/go/v1" "github.com/garethgeorge/backrest/gen/go/v1/v1connect" + "github.com/garethgeorge/backrest/gen/go/v1sync" "github.com/garethgeorge/backrest/internal/api/syncapi/permissions" "github.com/garethgeorge/backrest/internal/env" "github.com/garethgeorge/backrest/internal/oplog" @@ -125,7 +126,7 @@ func (c *SyncClient) RunSync(ctx context.Context) { state.ConnectionState = syncErr.State state.ConnectionStateMessage = syncErr.Message.Error() } else { - state.ConnectionState = v1.SyncConnectionState_CONNECTION_STATE_ERROR_INTERNAL + state.ConnectionState = v1sync.ConnectionState_CONNECTION_STATE_ERROR_INTERNAL state.ConnectionStateMessage = err.Error() } c.mgr.peerStateManager.SetPeerState(c.peer.Keyid, state) @@ -210,7 +211,7 @@ func (c *syncSessionHandlerClient) OnConnectionEstablished(ctx context.Context, if peerState == nil { peerState = newPeerState(c.peer.InstanceId, peer.Keyid) } - peerState.ConnectionState = v1.SyncConnectionState_CONNECTION_STATE_CONNECTED + peerState.ConnectionState = v1sync.ConnectionState_CONNECTION_STATE_CONNECTED peerState.ConnectionStateMessage = "connected" peerState.LastHeartbeat = time.Now() c.mgr.peerStateManager.SetPeerState(peer.Keyid, peerState) @@ -222,7 +223,7 @@ func (c *syncSessionHandlerClient) OnConnectionEstablished(ctx context.Context, // start by forwarding the configuration and the resource lists the peer is allowed to see. { - remoteConfig := &v1.RemoteConfig{ + remoteConfig := &v1sync.RemoteConfig{ Version: localConfig.Version, Modno: localConfig.Modno, } diff --git a/internal/api/syncapi/synccommon.go b/internal/api/syncapi/old/synccommon.go similarity index 100% rename from internal/api/syncapi/synccommon.go rename to internal/api/syncapi/old/synccommon.go diff --git a/internal/api/syncapi/synchandler.go b/internal/api/syncapi/old/synchandler.go similarity index 98% rename from internal/api/syncapi/synchandler.go rename to internal/api/syncapi/old/synchandler.go index 279a1f09..31ba4649 100644 --- a/internal/api/syncapi/synchandler.go +++ b/internal/api/syncapi/old/synchandler.go @@ -11,6 +11,7 @@ import ( "connectrpc.com/connect" v1 "github.com/garethgeorge/backrest/gen/go/v1" "github.com/garethgeorge/backrest/gen/go/v1/v1connect" + "github.com/garethgeorge/backrest/gen/go/v1sync" "github.com/garethgeorge/backrest/internal/api/syncapi/permissions" "github.com/garethgeorge/backrest/internal/env" "github.com/garethgeorge/backrest/internal/oplog" @@ -72,9 +73,9 @@ func (h *BackrestSyncHandler) Sync(ctx context.Context, stream *connect.BidiStre h.mgr.peerStateManager.SetPeerState(sessionHandler.peer.Keyid, peerState) } switch syncErr.State { - case v1.SyncConnectionState_CONNECTION_STATE_ERROR_AUTH: + case v1sync.ConnectionState_CONNECTION_STATE_ERROR_AUTH: return connect.NewError(connect.CodePermissionDenied, syncErr.Message) - case v1.SyncConnectionState_CONNECTION_STATE_ERROR_PROTOCOL: + case v1sync.ConnectionState_CONNECTION_STATE_ERROR_PROTOCOL: return connect.NewError(connect.CodeInvalidArgument, syncErr.Message) default: return connect.NewError(connect.CodeInternal, syncErr.Message) @@ -141,7 +142,7 @@ func (h *syncSessionHandlerServer) OnConnectionEstablished(ctx context.Context, // Configure the state for the connected peer. peerState := newPeerState(peer.InstanceId, h.peer.Keyid) peerState.ConnectionStateMessage = "connected" - peerState.ConnectionState = v1.SyncConnectionState_CONNECTION_STATE_CONNECTED + peerState.ConnectionState = v1sync.ConnectionState_CONNECTION_STATE_CONNECTED peerState.LastHeartbeat = time.Now() h.mgr.peerStateManager.SetPeerState(h.peer.Keyid, peerState) @@ -420,7 +421,7 @@ func (h *syncSessionHandlerServer) deleteByOriginalID(originalID int64) error { } func (h *syncSessionHandlerServer) sendConfigToClient(stream *bidiSyncCommandStream, config *v1.Config) error { - remoteConfig := &v1.RemoteConfig{ + remoteConfig := &v1sync.RemoteConfig{ Version: config.Version, Modno: config.Modno, } diff --git a/internal/api/syncapi/peerstate.go b/internal/api/syncapi/peerstate.go index 1ba157ee..10d8bd4f 100644 --- a/internal/api/syncapi/peerstate.go +++ b/internal/api/syncapi/peerstate.go @@ -7,7 +7,7 @@ import ( "sync" "time" - v1 "github.com/garethgeorge/backrest/gen/go/v1" + "github.com/garethgeorge/backrest/gen/go/v1sync" "github.com/garethgeorge/backrest/internal/eventemitter" "github.com/garethgeorge/backrest/internal/kvstore" "go.uber.org/zap" @@ -20,15 +20,15 @@ type PeerState struct { LastHeartbeat time.Time - ConnectionState v1.SyncConnectionState + ConnectionState v1sync.ConnectionState ConnectionStateMessage string // Plans and repos available on this peer - KnownRepos map[string]struct{} - KnownPlans map[string]struct{} + KnownRepos map[string]*v1sync.RepoMetadata + KnownPlans map[string]*v1sync.PlanMetadata // Partial configuration available for this peer - Config *v1.RemoteConfig + Config *v1sync.RemoteConfig } func newPeerState(instanceID, keyID string) *PeerState { @@ -36,10 +36,10 @@ func newPeerState(instanceID, keyID string) *PeerState { InstanceID: instanceID, KeyID: keyID, LastHeartbeat: time.Now(), - ConnectionState: v1.SyncConnectionState_CONNECTION_STATE_DISCONNECTED, + ConnectionState: v1sync.ConnectionState_CONNECTION_STATE_DISCONNECTED, ConnectionStateMessage: "disconnected", - KnownRepos: make(map[string]struct{}), - KnownPlans: make(map[string]struct{}), + KnownRepos: make(map[string]*v1sync.RepoMetadata), + KnownPlans: make(map[string]*v1sync.PlanMetadata), Config: nil, // Will be set when the config is received } } @@ -53,15 +53,15 @@ func (ps *PeerState) Clone() *PeerState { *clone = *ps // Shallow copy of the PeerState clone.KnownRepos = maps.Clone(ps.KnownRepos) // Clone maps to ensure deep copy clone.KnownPlans = maps.Clone(ps.KnownPlans) - clone.Config = proto.Clone(ps.Config).(*v1.RemoteConfig) // Clone the protobuf Config + clone.Config = proto.Clone(ps.Config).(*v1sync.RemoteConfig) // Clone the protobuf Config return clone } -func peerStateToProto(state *PeerState) *v1.PeerState { +func peerStateToProto(state *PeerState) *v1sync.PeerState { if state == nil { - return &v1.PeerState{} + return &v1sync.PeerState{} } - return &v1.PeerState{ + return &v1sync.PeerState{ PeerInstanceId: state.InstanceID, PeerKeyid: state.KeyID, LastHeartbeatMillis: state.LastHeartbeat.UnixMilli(), @@ -73,19 +73,18 @@ func peerStateToProto(state *PeerState) *v1.PeerState { } } -func peerStateFromProto(state *v1.PeerState) *PeerState { +func peerStateFromProto(state *v1sync.PeerState) *PeerState { if state.PeerInstanceId == "" || state.PeerKeyid == "" { return nil } - knownRepos := make(map[string]struct{}, len(state.KnownRepos)) + knownRepos := make(map[string]*v1sync.RepoMetadata, len(state.KnownRepos)) for _, repo := range state.KnownRepos { - knownRepos[repo] = struct{}{} + knownRepos[repo.Id] = repo } - knownPlans := make(map[string]struct{}, len(state.KnownPlans)) + knownPlans := make(map[string]*v1sync.PlanMetadata, len(state.KnownPlans)) for _, plan := range state.KnownPlans { - knownPlans[plan] = struct{}{} + knownPlans[plan.Id] = plan } - return &PeerState{ InstanceID: state.PeerInstanceId, KeyID: state.PeerKeyid, @@ -199,7 +198,7 @@ func (m *SqlitePeerStateManager) GetPeerState(keyID string) *PeerState { return nil } - var stateProto v1.PeerState + var stateProto v1sync.PeerState if err := proto.Unmarshal(stateBytes, &stateProto); err != nil { zap.S().Warnf("error unmarshalling peer state for key %s: %v", keyID, err) return nil @@ -214,7 +213,7 @@ func (m *SqlitePeerStateManager) GetAll() []*PeerState { states := make([]*PeerState, 0) m.kvstore.ForEach("", func(key string, value []byte) error { - var stateProto v1.PeerState + var stateProto v1sync.PeerState if err := proto.Unmarshal(value, &stateProto); err != nil { zap.S().Warnf("error unmarshalling peer state for key %s: %v", key, err) return nil // Skip this entry diff --git a/internal/api/syncapi/syncapi_test.go b/internal/api/syncapi/syncapi_test.nogo similarity index 97% rename from internal/api/syncapi/syncapi_test.go rename to internal/api/syncapi/syncapi_test.nogo index 121a9d75..606dec0b 100644 --- a/internal/api/syncapi/syncapi_test.go +++ b/internal/api/syncapi/syncapi_test.nogo @@ -14,6 +14,7 @@ import ( v1 "github.com/garethgeorge/backrest/gen/go/v1" "github.com/garethgeorge/backrest/gen/go/v1/v1connect" + "github.com/garethgeorge/backrest/gen/go/v1sync" "github.com/garethgeorge/backrest/internal/config" "github.com/garethgeorge/backrest/internal/cryptoutil" "github.com/garethgeorge/backrest/internal/logstore" @@ -156,7 +157,7 @@ func TestConnectionBadKeyRejected(t *testing.T) { startRunningSyncAPI(t, peerHost, peerHostAddr) startRunningSyncAPI(t, peerClient, peerClientAddr) - waitForConnectionState(t, ctx, peerClient, peerClientConfig.Multihost.KnownHosts[0], v1.SyncConnectionState_CONNECTION_STATE_ERROR_AUTH) + waitForConnectionState(t, ctx, peerClient, peerClientConfig.Multihost.KnownHosts[0], v1sync.ConnectionState_CONNECTION_STATE_ERROR_AUTH) } func TestSyncConfigChange(t *testing.T) { @@ -229,7 +230,7 @@ func TestSyncConfigChange(t *testing.T) { tryConnect(t, ctx, peerClient, peerClientConfig.Multihost.KnownHosts[0]) // wait for the initial config to propagate - tryExpectConfigFromHost(t, ctx, peerClient, peerClientConfig.Multihost.KnownHosts[0], &v1.RemoteConfig{ + tryExpectConfigFromHost(t, ctx, peerClient, peerClientConfig.Multihost.KnownHosts[0], &v1sync.RemoteConfig{ Repos: []*v1.Repo{ { Id: defaultRepoID, @@ -241,7 +242,7 @@ func TestSyncConfigChange(t *testing.T) { hostConfigChanged.Repos[0].Env = []string{"SOME_ENV=VALUE"} peerHost.configMgr.Update(hostConfigChanged) - tryExpectConfigFromHost(t, ctx, peerClient, peerClientConfig.Multihost.KnownHosts[0], &v1.RemoteConfig{ + tryExpectConfigFromHost(t, ctx, peerClient, peerClientConfig.Multihost.KnownHosts[0], &v1sync.RemoteConfig{ Repos: []*v1.Repo{ { Id: defaultRepoID, @@ -602,7 +603,7 @@ func tryExpectOperationsSynced(t *testing.T, ctx context.Context, peer1 *peerUnd } } -func tryExpectConfigFromHost(t *testing.T, ctx context.Context, peer *peerUnderTest, hostPeer *v1.Multihost_Peer, wantCfg *v1.RemoteConfig) { +func tryExpectConfigFromHost(t *testing.T, ctx context.Context, peer *peerUnderTest, hostPeer *v1.Multihost_Peer, wantCfg *v1sync.RemoteConfig) { testutil.Try(t, ctx, func() error { state := peer.manager.peerStateManager.GetPeerState(hostPeer.Keyid) if state == nil { @@ -615,7 +616,7 @@ func tryExpectConfigFromHost(t *testing.T, ctx context.Context, peer *peerUnderT }) } -func waitForConnectionState(t *testing.T, ctx context.Context, peer *peerUnderTest, hostPeer *v1.Multihost_Peer, wantState v1.SyncConnectionState) { +func waitForConnectionState(t *testing.T, ctx context.Context, peer *peerUnderTest, hostPeer *v1.Multihost_Peer, wantState v1sync.ConnectionState) { ctx, cancel := testutil.WithDeadlineFromTest(t, ctx) defer cancel() @@ -625,7 +626,7 @@ func waitForConnectionState(t *testing.T, ctx context.Context, peer *peerUnderTe // First check if the peer is already connected. state := peer.manager.peerStateManager.GetPeerState(hostPeer.Keyid) - if state != nil && state.ConnectionState == v1.SyncConnectionState_CONNECTION_STATE_CONNECTED { + if state != nil && state.ConnectionState == v1sync.ConnectionState_CONNECTION_STATE_CONNECTED { return // Already connected, nothing to do } @@ -641,7 +642,7 @@ func waitForConnectionState(t *testing.T, ctx context.Context, peer *peerUnderTe } if state.KeyID == hostPeer.Keyid && state.InstanceID == hostPeer.InstanceId { lastState = state - if state.ConnectionState == v1.SyncConnectionState_CONNECTION_STATE_CONNECTED { + if state.ConnectionState == v1sync.ConnectionState_CONNECTION_STATE_CONNECTED { stop = true continue } @@ -653,13 +654,13 @@ func waitForConnectionState(t *testing.T, ctx context.Context, peer *peerUnderTe } if lastState == nil { t.Fatalf("timeout waiting for connection to host peer %s", hostPeer.InstanceId) - } else if lastState.ConnectionState != v1.SyncConnectionState_CONNECTION_STATE_CONNECTED { + } else if lastState.ConnectionState != v1sync.ConnectionState_CONNECTION_STATE_CONNECTED { t.Fatalf("expected connection state to be CONNECTED, got %v (reason: %q)", lastState.ConnectionState, lastState.ConnectionStateMessage) } } func tryConnect(t *testing.T, ctx context.Context, peer *peerUnderTest, hostPeer *v1.Multihost_Peer) { - waitForConnectionState(t, ctx, peer, hostPeer, v1.SyncConnectionState_CONNECTION_STATE_CONNECTED) + waitForConnectionState(t, ctx, peer, hostPeer, v1sync.ConnectionState_CONNECTION_STATE_CONNECTED) } func runSyncAPIWithCtx(ctx context.Context, peer *peerUnderTest, bindAddr string) { diff --git a/internal/api/syncapi/syncmanager.go b/internal/api/syncapi/syncmanager.nogo similarity index 97% rename from internal/api/syncapi/syncmanager.go rename to internal/api/syncapi/syncmanager.nogo index 550156c0..c373596a 100644 --- a/internal/api/syncapi/syncmanager.go +++ b/internal/api/syncapi/syncmanager.nogo @@ -9,6 +9,7 @@ import ( "time" v1 "github.com/garethgeorge/backrest/gen/go/v1" + "github.com/garethgeorge/backrest/gen/go/v1sync" "github.com/garethgeorge/backrest/internal/config" "github.com/garethgeorge/backrest/internal/cryptoutil" "github.com/garethgeorge/backrest/internal/oplog" @@ -42,7 +43,7 @@ func NewSyncManager(configMgr *config.ConfigManager, oplog *oplog.OpLog, orchest if state == nil { state = newPeerState(knownHostPeer.InstanceId, knownHostPeer.Keyid) } - state.ConnectionState = v1.SyncConnectionState_CONNECTION_STATE_DISCONNECTED + state.ConnectionState = v1sync.ConnectionState_CONNECTION_STATE_DISCONNECTED state.ConnectionStateMessage = "disconnected" peerStateManager.SetPeerState(knownHostPeer.Keyid, state) } @@ -51,7 +52,7 @@ func NewSyncManager(configMgr *config.ConfigManager, oplog *oplog.OpLog, orchest if state == nil { state = newPeerState(authorizedClient.InstanceId, authorizedClient.Keyid) } - state.ConnectionState = v1.SyncConnectionState_CONNECTION_STATE_DISCONNECTED + state.ConnectionState = v1sync.ConnectionState_CONNECTION_STATE_DISCONNECTED state.ConnectionStateMessage = "disconnected" peerStateManager.SetPeerState(authorizedClient.Keyid, state) } diff --git a/internal/api/syncapi/syncstatehandler.go b/internal/api/syncapi/syncstatehandler.nogo similarity index 81% rename from internal/api/syncapi/syncstatehandler.go rename to internal/api/syncapi/syncstatehandler.nogo index 5973db27..8a19d007 100644 --- a/internal/api/syncapi/syncstatehandler.go +++ b/internal/api/syncapi/syncstatehandler.nogo @@ -6,15 +6,16 @@ import ( "connectrpc.com/connect" v1 "github.com/garethgeorge/backrest/gen/go/v1" - "github.com/garethgeorge/backrest/gen/go/v1/v1connect" + "github.com/garethgeorge/backrest/gen/go/v1sync" + "github.com/garethgeorge/backrest/gen/go/v1sync/v1syncconnect" ) type BackrestSyncStateHandler struct { - v1connect.UnimplementedBackrestSyncStateServiceHandler + v1syncconnect.UnimplementedBackrestSyncStateServiceHandler mgr *SyncManager } -var _ v1connect.BackrestSyncStateServiceHandler = &BackrestSyncStateHandler{} +var _ v1syncconnect.BackrestSyncStateServiceHandler = &BackrestSyncStateHandler{} func NewBackrestSyncStateHandler(mgr *SyncManager) *BackrestSyncStateHandler { return &BackrestSyncStateHandler{ @@ -22,14 +23,14 @@ func NewBackrestSyncStateHandler(mgr *SyncManager) *BackrestSyncStateHandler { } } -func (h *BackrestSyncStateHandler) GetPeerSyncStatesStream(ctx context.Context, req *connect.Request[v1.SyncStateStreamRequest], stream *connect.ServerStream[v1.PeerState]) error { +func (h *BackrestSyncStateHandler) GetPeerSyncStatesStream(ctx context.Context, req *connect.Request[v1sync.SyncStateStreamRequest], stream *connect.ServerStream[v1.PeerState]) error { ctx, cancel := context.WithCancelCause(ctx) defer cancel(nil) // Subscribe to the peer state changes onStateChangeChan := h.mgr.peerStateManager.OnStateChanged().Subscribe() - messagesToSend := make(chan *v1.PeerState, 100) // Buffered channel to allow sending items without blocking + messagesToSend := make(chan *v1sync.PeerState, 100) // Buffered channel to allow sending items without blocking sendAllInList := func(peers []*v1.Multihost_Peer) { for _, peerState := range peers {