mirror of
https://github.com/garethgeorge/backrest.git
synced 2026-09-18 22:15:34 +00:00
initial test coverage
This commit is contained in:
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
@@ -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()
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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
|
||||
|
||||
@@ -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) {
|
||||
@@ -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)
|
||||
}
|
||||
+6
-5
@@ -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 {
|
||||
Reference in New Issue
Block a user