initial test coverage

This commit is contained in:
Gareth George
2025-10-31 18:24:19 -07:00
committed by Gareth
parent 764693ba47
commit 1e4524fe01
14 changed files with 798 additions and 164 deletions
+245
View File
@@ -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
}
+198
View File
@@ -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
}
+18
View File
@@ -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,
}
}
-111
View File
@@ -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()
}
}
}
+10 -10
View File
@@ -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,
}
}
+280
View File
@@ -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,
}
+19 -20
View File
@@ -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,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 {