From db4f309a7797eeb4ba12a227bd7397173d861f5d Mon Sep 17 00:00:00 2001 From: mlsmaycon Date: Fri, 11 Sep 2026 19:06:29 +0200 Subject: [PATCH] [management] Allow concurrent service domain authorization Service writes only need shared locks on registrations covering their hostname. Hold those locks until commit to preserve the deletion guard while allowing concurrent writes on the same or unrelated domains. --- .../reverseproxy/domain/manager/manager.go | 6 +- .../server/store/sql_store_domain_deletion.go | 21 ++++-- .../store/sql_store_domain_deletion_test.go | 69 ++++++++++++++++++- management/server/store/store.go | 3 +- management/server/store/store_mock.go | 9 +-- 5 files changed, 94 insertions(+), 14 deletions(-) diff --git a/management/internals/modules/reverseproxy/domain/manager/manager.go b/management/internals/modules/reverseproxy/domain/manager/manager.go index da07bb368..c5ead6493 100644 --- a/management/internals/modules/reverseproxy/domain/manager/manager.go +++ b/management/internals/modules/reverseproxy/domain/manager/manager.go @@ -362,7 +362,11 @@ func (m Manager) ValidateServiceDomain(ctx context.Context, tx nbstore.Store, ac if _, ok := ExtractClusterFromFreeDomain(serviceDomain, []string{cluster}); ok { return nil } - customDomains, err := tx.LockCustomDomains(ctx, accountID) + name, err := nbdomain.FromString(serviceDomain) + if err != nil { + return fmt.Errorf("parse service domain: %w", err) + } + customDomains, err := tx.LockCustomDomains(ctx, accountID, name) if err != nil { return err } diff --git a/management/server/store/sql_store_domain_deletion.go b/management/server/store/sql_store_domain_deletion.go index dcad69503..9de50c048 100644 --- a/management/server/store/sql_store_domain_deletion.go +++ b/management/server/store/sql_store_domain_deletion.go @@ -4,19 +4,28 @@ import ( "context" "errors" "fmt" + "strings" "gorm.io/gorm" "gorm.io/gorm/clause" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain" + nbdomain "github.com/netbirdio/netbird/shared/management/domain" "github.com/netbirdio/netbird/shared/management/status" ) -// LockCustomDomains locks an account's registrations until the caller's transaction ends. -func (s *SqlStore) LockCustomDomains(ctx context.Context, accountID string) ([]*domain.Domain, error) { +// LockCustomDomains holds shared locks on registrations covering a service until commit. +func (s *SqlStore) LockCustomDomains(ctx context.Context, accountID string, serviceDomain nbdomain.Domain) ([]*domain.Domain, error) { + var names []string + for name := serviceDomain.PunycodeString(); name != ""; { + names = append(names, name) + _, name, _ = strings.Cut(name, ".") + } + var domains []*domain.Domain - if err := s.db.WithContext(ctx).Clauses(clause.Locking{Strength: string(LockingStrengthUpdate)}). - Where(accountIDCondition, accountID).Order("id").Find(&domains).Error; err != nil { + if err := s.db.WithContext(ctx).Clauses(clause.Locking{Strength: string(LockingStrengthShare)}). + Where(accountIDCondition, accountID).Where("domain IN ?", names). + Order("id").Find(&domains).Error; err != nil { return nil, fmt.Errorf("lock custom domains: %w", err) } return domains, nil @@ -26,8 +35,8 @@ func (s *SqlStore) LockCustomDomains(ctx context.Context, accountID string) ([]* func (s *SqlStore) DeleteCustomDomain(ctx context.Context, accountID string, domainID string) error { return s.db.WithContext(ctx).Transaction(func(tx *gorm.DB) error { var d domain.Domain - // Service writes take the same lock before checking validation, so neither - // operation can commit against the other's outdated view of the domain. + // Service writes hold a shared lock on this row through commit, so neither + // operation can proceed against the other's outdated view of the domain. if err := tx.Clauses(clause.Locking{Strength: string(LockingStrengthUpdate)}). Take(&d, accountAndIDQueryCondition, accountID, domainID).Error; err != nil { if errors.Is(err, gorm.ErrRecordNotFound) { diff --git a/management/server/store/sql_store_domain_deletion_test.go b/management/server/store/sql_store_domain_deletion_test.go index 9f0638bea..cc0c94b80 100644 --- a/management/server/store/sql_store_domain_deletion_test.go +++ b/management/server/store/sql_store_domain_deletion_test.go @@ -10,9 +10,74 @@ import ( "github.com/stretchr/testify/require" rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + "github.com/netbirdio/netbird/management/server/types" + nbdomain "github.com/netbirdio/netbird/shared/management/domain" "github.com/netbirdio/netbird/shared/management/status" ) +func TestLockCustomDomains_ConcurrentServices(t *testing.T) { + runTestForAllEngines(t, "", func(t *testing.T, store Store) { + if store.GetStoreEngine() == types.SqliteStoreEngine { + t.Skip("SQLite serializes transactions on one connection") + } + ctx, cancel := context.WithTimeout(context.Background(), 15*time.Second) + defer cancel() + require.NoError(t, store.SaveAccount(ctx, newAccountWithId(ctx, "owner", "admin", ""))) + _, err := store.CreateCustomDomain(ctx, "owner", "one.example.com", "cluster", true) + require.NoError(t, err) + _, err = store.CreateCustomDomain(ctx, "owner", "two.example.com", "cluster", true) + require.NoError(t, err) + + locked := make(chan error, 1) + release := make(chan struct{}) + done := make(chan error, 1) + go func() { + done <- store.ExecuteInTransaction(ctx, func(tx Store) error { + _, err := tx.LockCustomDomains(ctx, "owner", "app.one.example.com") + locked <- err + if err != nil { + return err + } + select { + case <-release: + return nil + case <-ctx.Done(): + return ctx.Err() + } + }) + }() + var lockErr error + select { + case lockErr = <-locked: + case err := <-done: + t.Fatalf("transaction ended before locking: %v", err) + } + writeCtx, writeCancel := context.WithTimeout(ctx, 3*time.Second) + defer writeCancel() + var writeErr error + for _, name := range []nbdomain.Domain{"app.one.example.com", "app.two.example.com"} { + writeErr = store.ExecuteInTransaction(writeCtx, func(tx Store) error { + if _, err := tx.LockCustomDomains(writeCtx, "owner", name); err != nil { + return err + } + return tx.CreateService(writeCtx, &rpservice.Service{ + ID: name.PunycodeString(), AccountID: "owner", Domain: name.PunycodeString(), + }) + }) + if writeErr != nil { + break + } + } + close(release) + require.NoError(t, <-done) + require.NoError(t, lockErr) + require.NoError(t, writeErr, "domain authorization locks must allow concurrent service writes") + services, err := store.GetAccountServices(ctx, LockingStrengthNone, "owner") + require.NoError(t, err) + assert.Len(t, services, 2, "both services must commit while the first domain is locked") + }) +} + func TestDeleteCustomDomain_ServiceDependencies(t *testing.T) { runTestForAllEngines(t, "", func(t *testing.T, store Store) { ctx := context.Background() @@ -52,14 +117,14 @@ func TestDeleteCustomDomain_ConcurrentServiceCreation(t *testing.T) { for i := range 10 { d, err := store.CreateCustomDomain(ctx, "owner", fmt.Sprintf("app%d.example.com", i), "cluster", true) require.NoError(t, err) - svc := &rpservice.Service{ID: fmt.Sprintf("service-%d", i), AccountID: "owner", Domain: d.Domain} + svc := &rpservice.Service{ID: fmt.Sprintf("service-%d", i), AccountID: "owner", Domain: "nested." + d.Domain} start := make(chan struct{}) created := make(chan error, 1) deleted := make(chan error, 1) go func() { <-start created <- store.ExecuteInTransaction(ctx, func(tx Store) error { - domains, err := tx.LockCustomDomains(ctx, "owner") + domains, err := tx.LockCustomDomains(ctx, "owner", nbdomain.Domain(svc.Domain)) if err != nil { return err } diff --git a/management/server/store/store.go b/management/server/store/store.go index fc050f3ba..9761fb4b4 100644 --- a/management/server/store/store.go +++ b/management/server/store/store.go @@ -34,6 +34,7 @@ import ( "github.com/netbirdio/netbird/management/server/telemetry" "github.com/netbirdio/netbird/management/server/testutil" "github.com/netbirdio/netbird/management/server/types" + nbdomain "github.com/netbirdio/netbird/shared/management/domain" "github.com/netbirdio/netbird/util" "github.com/netbirdio/netbird/util/crypt" @@ -302,7 +303,7 @@ type Store interface { GetCustomDomain(ctx context.Context, accountID string, domainID string) (*domain.Domain, error) ListFreeDomains(ctx context.Context, accountID string) ([]string, error) ListCustomDomains(ctx context.Context, accountID string) ([]*domain.Domain, error) - LockCustomDomains(ctx context.Context, accountID string) ([]*domain.Domain, error) + LockCustomDomains(ctx context.Context, accountID string, serviceDomain nbdomain.Domain) ([]*domain.Domain, error) GetCustomDomainByName(ctx context.Context, domainName string) (*domain.Domain, error) CreateCustomDomain(ctx context.Context, accountID string, domainName string, targetCluster string, validated bool) (*domain.Domain, error) UpdateCustomDomain(ctx context.Context, accountID string, d *domain.Domain) (*domain.Domain, error) diff --git a/management/server/store/store_mock.go b/management/server/store/store_mock.go index 066e148f0..a10d94c71 100644 --- a/management/server/store/store_mock.go +++ b/management/server/store/store_mock.go @@ -31,6 +31,7 @@ import ( posture "github.com/netbirdio/netbird/management/server/posture" types3 "github.com/netbirdio/netbird/management/server/types" route "github.com/netbirdio/netbird/route" + domain0 "github.com/netbirdio/netbird/shared/management/domain" crypt "github.com/netbirdio/netbird/util/crypt" gomock "go.uber.org/mock/gomock" ) @@ -3213,18 +3214,18 @@ func (mr *MockStoreMockRecorder) ListFreeDomains(ctx, accountID any) *gomock.Cal } // LockCustomDomains mocks base method. -func (m *MockStore) LockCustomDomains(ctx context.Context, accountID string) ([]*domain.Domain, error) { +func (m *MockStore) LockCustomDomains(ctx context.Context, accountID string, serviceDomain domain0.Domain) ([]*domain.Domain, error) { m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "LockCustomDomains", ctx, accountID) + ret := m.ctrl.Call(m, "LockCustomDomains", ctx, accountID, serviceDomain) ret0, _ := ret[0].([]*domain.Domain) ret1, _ := ret[1].(error) return ret0, ret1 } // LockCustomDomains indicates an expected call of LockCustomDomains. -func (mr *MockStoreMockRecorder) LockCustomDomains(ctx, accountID any) *gomock.Call { +func (mr *MockStoreMockRecorder) LockCustomDomains(ctx, accountID, serviceDomain any) *gomock.Call { mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "LockCustomDomains", reflect.TypeOf((*MockStore)(nil).LockCustomDomains), ctx, accountID) + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "LockCustomDomains", reflect.TypeOf((*MockStore)(nil).LockCustomDomains), ctx, accountID, serviceDomain) } // MarkAccountPrimary mocks base method.