From b789ffbb9fc77e4ef9575023a6603bbdc7105f86 Mon Sep 17 00:00:00 2001 From: Maycon Santos Date: Fri, 11 Sep 2026 17:57:58 +0200 Subject: [PATCH] [management] Expire unvalidated custom domain registrations (#7497) Prevent unvalidated registrations from reserving domain names indefinitely. Give pending registrations a 48-hour validation window and clean up expired entries at startup and every 60 minutes. Emit CustomDomainValidationExpired for each deletion and preserve registrations referenced by services. Reject validation after expiry and prevent concurrent validation from recreating deleted registrations. Normalize domain names with the shared parser before registration. Migrate existing pending registrations to receive a fresh 48-hour validation window. --- docs/custom-domain-validation.md | 28 ++ .../modules/reverseproxy/domain/domain.go | 26 +- .../reverseproxy/domain/manager/expiration.go | 73 +++++ .../domain/manager/expiration_test.go | 274 ++++++++++++++++++ .../reverseproxy/domain/manager/manager.go | 52 +++- .../domain/manager/manager_realstore_test.go | 11 +- .../domain/manager/manager_test.go | 9 + .../domain/manager/normalization_test.go | 83 ++++++ management/internals/server/server.go | 23 +- management/server/activity/codes.go | 10 +- management/server/activity/store/sql_store.go | 12 +- .../server/activity/store/sql_store_test.go | 38 +++ management/server/event.go | 41 ++- management/server/event_test.go | 42 +++ .../events_handler_integration_test.go | 13 +- .../migration/migration_custom_domain.go | 22 ++ .../migration/migration_custom_domain_test.go | 44 +++ management/server/store/sql_store.go | 21 +- .../store/sql_store_domain_expiration.go | 60 ++++ .../store/sql_store_domain_expiration_test.go | 78 +++++ .../store/sql_store_get_account_test.go | 6 +- management/server/store/store.go | 5 + management/server/store/store_mock.go | 30 ++ 23 files changed, 941 insertions(+), 60 deletions(-) create mode 100644 docs/custom-domain-validation.md create mode 100644 management/internals/modules/reverseproxy/domain/manager/expiration.go create mode 100644 management/internals/modules/reverseproxy/domain/manager/expiration_test.go create mode 100644 management/internals/modules/reverseproxy/domain/manager/normalization_test.go create mode 100644 management/server/migration/migration_custom_domain.go create mode 100644 management/server/migration/migration_custom_domain_test.go create mode 100644 management/server/store/sql_store_domain_expiration.go create mode 100644 management/server/store/sql_store_domain_expiration_test.go diff --git a/docs/custom-domain-validation.md b/docs/custom-domain-validation.md new file mode 100644 index 000000000..a70388e4d --- /dev/null +++ b/docs/custom-domain-validation.md @@ -0,0 +1,28 @@ +A new custom domain name is converted to lowercase ASCII (punycode), with a +trailing dot removed, before availability and DNS validation checks. Invalid +names and wildcard registrations are rejected before storage. + +A custom domain registration must complete validation within 48 hours of +creation. Retrying validation does not extend this window. Once validation +succeeds, the registration is exempt from this expiration policy. + +Management removes expired, unvalidated registrations at startup and every +60 minutes. While management is running, removal normally occurs between +48 and 49 hours after registration. Validation is refused after the 48-hour +deadline even if cleanup has not yet removed the registration. + +Removal releases the name for a new registration. The new registration must +complete its own validation. Its account does not inherit validation or +services from the expired registration. + +The original account receives a system activity event named +`CustomDomainValidationExpired`, displayed as "Unvalidated domain registration +expired". The event includes the domain name, original registration ID, and +validation deadline. + +On upgrade, existing unvalidated registrations receive a 48-hour validation +window. Restarting management does not extend a previously assigned deadline. + +Registrations with existing services, including services using subdomains, are +retained for operator review. Management logs their account and domain IDs so +an operator can identify and resolve those dependencies before cleanup. diff --git a/management/internals/modules/reverseproxy/domain/domain.go b/management/internals/modules/reverseproxy/domain/domain.go index 08d7ad19b..74583c55f 100644 --- a/management/internals/modules/reverseproxy/domain/domain.go +++ b/management/internals/modules/reverseproxy/domain/domain.go @@ -1,5 +1,13 @@ package domain +import "time" + +// ValidationTTL is the time available to validate a custom domain registration. +const ValidationTTL = 48 * time.Hour + +// ID identifies a custom domain registration. +type ID string + type Type string const ( @@ -8,12 +16,13 @@ const ( ) type Domain struct { - ID string `gorm:"unique;primaryKey;autoIncrement"` - Domain string `gorm:"unique"` // Domain records must be unique, this avoids domain reuse across accounts. - AccountID string `gorm:"index"` - TargetCluster string // The proxy cluster this domain should be validated against - Type Type `gorm:"-"` - Validated bool + ID string `gorm:"unique;primaryKey;autoIncrement"` + Domain string `gorm:"unique"` // Domain records must be unique, this avoids domain reuse across accounts. + AccountID string `gorm:"index"` + TargetCluster string // The proxy cluster this domain should be validated against + Type Type `gorm:"-"` + Validated bool + ValidationExpiresAt *time.Time `gorm:"index"` // SupportsCustomPorts is populated at query time for free domains from the // proxy cluster capabilities. Not persisted. SupportsCustomPorts *bool `gorm:"-"` @@ -36,7 +45,12 @@ func (d *Domain) EventMeta() map[string]any { } } +// Copy returns a copy with an independent validation deadline. func (d *Domain) Copy() *Domain { dCopy := *d + if d.ValidationExpiresAt != nil { + expiresAt := *d.ValidationExpiresAt + dCopy.ValidationExpiresAt = &expiresAt + } return &dCopy } diff --git a/management/internals/modules/reverseproxy/domain/manager/expiration.go b/management/internals/modules/reverseproxy/domain/manager/expiration.go new file mode 100644 index 000000000..381e3ca51 --- /dev/null +++ b/management/internals/modules/reverseproxy/domain/manager/expiration.go @@ -0,0 +1,73 @@ +package manager + +import ( + "context" + "time" + + log "github.com/sirupsen/logrus" + + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain" + "github.com/netbirdio/netbird/management/server/activity" +) + +const ( + validationCleanupInterval = 60 * time.Minute + validationCleanupBatch = 100 +) + +// RunValidationCleanup removes expired registrations on startup and hourly until cancellation. +func (m Manager) RunValidationCleanup(ctx context.Context) { + ticker := time.NewTicker(validationCleanupInterval) + defer ticker.Stop() + for { + m.cleanupExpiredDomains(ctx, time.Now().UTC()) + select { + case <-ctx.Done(): + return + case <-ticker.C: + } + } +} + +func (m Manager) cleanupExpiredDomains(ctx context.Context, now time.Time) { + var afterID domain.ID + for ctx.Err() == nil { + domains, err := m.store.GetExpiredCustomDomains(ctx, now, afterID, validationCleanupBatch) + if err != nil { + if ctx.Err() == nil { + log.WithContext(ctx).WithError(err).Error("list expired custom domain registrations") + } + return + } + for _, d := range domains { + if ctx.Err() != nil { + return + } + m.deleteExpiredDomain(ctx, d, now) + afterID = domain.ID(d.ID) + } + if len(domains) < validationCleanupBatch { + return + } + } +} + +func (m Manager) deleteExpiredDomain(ctx context.Context, d *domain.Domain, now time.Time) { + deleted, err := m.store.DeleteExpiredCustomDomain(ctx, d, now) + if err != nil { + if ctx.Err() == nil { + log.WithContext(ctx).WithFields(log.Fields{"accountID": d.AccountID, "domainID": d.ID}). + WithError(err).Warn("could not expire custom domain registration") + } + return + } + if !deleted { + return + } + meta := d.EventMeta() + if d.ValidationExpiresAt != nil { + meta["validation_expires_at"] = d.ValidationExpiresAt.UTC().Format(time.RFC3339) + } + m.accountManager.StoreEvent(ctx, activity.SystemInitiator, d.ID, d.AccountID, + activity.CustomDomainValidationExpired, meta) +} diff --git a/management/internals/modules/reverseproxy/domain/manager/expiration_test.go b/management/internals/modules/reverseproxy/domain/manager/expiration_test.go new file mode 100644 index 000000000..e3b05efbb --- /dev/null +++ b/management/internals/modules/reverseproxy/domain/manager/expiration_test.go @@ -0,0 +1,274 @@ +package manager + +import ( + "context" + "fmt" + "sync" + "testing" + "testing/synctest" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain" + rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + "github.com/netbirdio/netbird/management/server/activity" + "github.com/netbirdio/netbird/management/server/mock_server" + nbstore "github.com/netbirdio/netbird/management/server/store" +) + +func TestValidateDomain_ExpiredRegistration(t *testing.T) { + env := setupDomainTest(t) + ctx := context.Background() + d, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "expired.example.com", testCluster) + require.NoError(t, err) + expiresAt := time.Now().Add(-time.Second) + db := env.store.(*nbstore.SqlStore).GetDB() + require.NoError(t, db.Model(&domain.Domain{}).Where("id = ?", d.ID). + Update("validation_expires_at", expiresAt).Error) + env.resolver.set("validation.expired.example.com", testCluster) + + env.manager.ValidateDomain(ctx, accountA, accountAUser, d.ID) + + stored := storedDomain(t, env.store, accountA, d.Domain) + require.NotNil(t, stored) + assert.False(t, stored.Validated, "an expired registration must not become usable before cleanup runs") +} + +func TestCreateDomain_ValidationDeadline(t *testing.T) { + env := setupClockDomainTest(t) + synctest.Test(t, func(t *testing.T) { + ctx := context.Background() + createdAt := time.Now().UTC() + d, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "pending.example.com", testCluster) + require.NoError(t, err) + require.NotNil(t, d.ValidationExpiresAt) + assert.Equal(t, createdAt.Add(48*time.Hour), *d.ValidationExpiresAt, "new registrations get 48 hours") + + time.Sleep(time.Hour) + env.manager.ValidateDomain(ctx, accountA, accountAUser, d.ID) + stored := storedDomain(t, env.store, accountA, d.Domain) + require.NotNil(t, stored) + require.NotNil(t, stored.ValidationExpiresAt) + assert.WithinDuration(t, *d.ValidationExpiresAt, *stored.ValidationExpiresAt, 0, "failed validation must not extend the deadline") + }) +} + +func TestCleanupExpiredDomains_Boundaries(t *testing.T) { + env := setupDomainTest(t) + events := captureDomainEvents(env) + ctx := context.Background() + now := time.Now().UTC().Truncate(time.Second) + tests := []struct { + name string + expiresAt time.Time + validated bool + deleted bool + }{ + {"expired", now.Add(-time.Second), false, true}, + {"deadline", now, false, true}, + {"pending", now.Add(time.Second), false, false}, + {"validated", now.Add(-time.Hour), true, false}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + d := createExpiringDomain(t, env, tt.name+".example.com", tt.expiresAt) + if tt.validated { + require.NoError(t, env.store.(*nbstore.SqlStore).GetDB().Model(d).Update("validated", true).Error) + } + env.manager.cleanupExpiredDomains(ctx, now) + stored := storedDomain(t, env.store, accountA, d.Domain) + if !tt.deleted { + assert.NotNil(t, stored, "pending and validated registrations must survive cleanup") + return + } + assert.Nil(t, stored, "expired unused registrations must be removed") + replacement, err := env.manager.CreateDomain(ctx, accountB, accountBUser, d.Domain, testCluster) + require.NoError(t, err) + assert.NotEqual(t, d.ID, replacement.ID, "the released name must receive a fresh registration") + assert.False(t, replacement.Validated, "the new account must validate its own registration") + }) + } + got := events.get() + require.Len(t, got, 2, "only successful expiration deletions emit events") + for _, event := range got { + assert.Equal(t, activity.CustomDomainValidationExpired, event.Activity, "use the requested expiration event") + assert.Equal(t, activity.SystemInitiator, event.InitiatorID, "cleanup is attributed to the system") + assert.Equal(t, accountA, event.AccountID, "expiration belongs to the original account") + assert.NotEmpty(t, event.TargetID, "retain the deleted domain ID") + assert.NotEmpty(t, event.Meta["domain"], "retain the deleted domain name") + assert.NotEmpty(t, event.Meta["validation_expires_at"], "include the validation deadline") + } +} + +func TestCleanupExpiredDomains_ContinuesPastProtectedBatch(t *testing.T) { + env := setupDomainTest(t) + ctx := context.Background() + now := time.Now().UTC() + for i := range validationCleanupBatch { + d := createExpiringDomain(t, env, fmt.Sprintf("protected-%d.example.com", i), now.Add(-time.Hour)) + require.NoError(t, env.store.CreateService(ctx, &rpservice.Service{ + ID: fmt.Sprintf("service-%d", i), AccountID: accountA, Domain: "app." + d.Domain, + })) + } + unprotected := createExpiringDomain(t, env, "unused.example.com", now.Add(-time.Hour)) + env.manager.cleanupExpiredDomains(ctx, now) + assert.Nil(t, storedDomain(t, env.store, accountA, unprotected.Domain), "protected registrations must not starve later batches") + remaining, err := env.store.ListCustomDomains(ctx, accountA) + require.NoError(t, err) + assert.Len(t, remaining, validationCleanupBatch, "all registrations with dependent services must survive") +} + +func TestCleanupExpiredDomains_ConcurrentWorkers(t *testing.T) { + env := setupDomainTest(t) + events := captureDomainEvents(env) + now := time.Now().UTC() + d := createExpiringDomain(t, env, "concurrent.example.com", now.Add(-time.Hour)) + var workers sync.WaitGroup + for range 2 { + workers.Go(func() { env.manager.cleanupExpiredDomains(context.Background(), now) }) + } + workers.Wait() + assert.Nil(t, storedDomain(t, env.store, accountA, d.Domain), "one worker must remove the expired registration") + assert.Len(t, events.get(), 1, "only the worker that deletes the row may emit the event") +} + +func TestRunValidationCleanup_HourlyAndRestart(t *testing.T) { + env := setupClockDomainTest(t) + synctest.Test(t, func(t *testing.T) { + events := captureDomainEvents(env) + now := time.Now().UTC() + startup := createExpiringDomain(t, env, "startup.example.com", now.Add(-time.Hour)) + hourly := createExpiringDomain(t, env, "hourly.example.com", now.Add(time.Minute)) + ctx, cancel := context.WithCancel(context.Background()) + done := make(chan struct{}) + go func() { + defer close(done) + env.manager.RunValidationCleanup(ctx) + }() + synctest.Wait() + assert.Nil(t, storedDomain(t, env.store, accountA, startup.Domain), "startup must collect overdue registrations") + time.Sleep(59 * time.Minute) + synctest.Wait() + assert.NotNil(t, storedDomain(t, env.store, accountA, hourly.Domain), "cleanup must wait for the 60-minute interval") + time.Sleep(time.Minute) + synctest.Wait() + assert.Nil(t, storedDomain(t, env.store, accountA, hourly.Domain), "the hourly scan must collect expired registrations") + cancel() + <-done + + offline := createExpiringDomain(t, env, "offline.example.com", time.Now().UTC().Add(time.Minute)) + time.Sleep(2 * time.Hour) + assert.NotNil(t, storedDomain(t, env.store, accountA, offline.Domain), "a stopped worker must not continue deleting") + ctx, cancel = context.WithCancel(context.Background()) + done = make(chan struct{}) + go func() { + defer close(done) + env.manager.RunValidationCleanup(ctx) + }() + synctest.Wait() + assert.Nil(t, storedDomain(t, env.store, accountA, offline.Domain), "restart must use the persisted deadline") + cancel() + <-done + assert.Len(t, events.get(), 3, "each deletion should emit an expiration event") + }) +} + +type blockingDomainResolver struct { + started chan struct{} + release chan struct{} +} + +func (r blockingDomainResolver) LookupCNAME(context.Context, string) (string, error) { + close(r.started) + <-r.release + return testCluster + ".", nil +} + +func TestValidateDomain_DeadlinePassesDuringLookup(t *testing.T) { + for _, cleanup := range []bool{false, true} { + t.Run(fmt.Sprintf("cleanup=%t", cleanup), func(t *testing.T) { + env := setupClockDomainTest(t) + synctest.Test(t, func(t *testing.T) { + events := captureDomainEvents(env) + ctx := context.Background() + d, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "late.example.com", testCluster) + require.NoError(t, err) + resolver := blockingDomainResolver{started: make(chan struct{}), release: make(chan struct{})} + env.manager.validator.Resolver = resolver + done := make(chan struct{}) + go func() { + defer close(done) + env.manager.ValidateDomain(ctx, accountA, accountAUser, d.ID) + }() + <-resolver.started + time.Sleep(48 * time.Hour) + if cleanup { + env.manager.cleanupExpiredDomains(ctx, time.Now().UTC()) + _, err = env.store.CreateCustomDomain(ctx, accountB, d.Domain, testCluster, false) + require.NoError(t, err) + } + close(resolver.release) + <-done + owner := accountA + if cleanup { + assert.Nil(t, storedDomain(t, env.store, accountA, d.Domain), "late validation must not restore the old claim") + owner = accountB + } + stored := storedDomain(t, env.store, owner, d.Domain) + require.NotNil(t, stored) + assert.False(t, stored.Validated, "late validation must not validate either claim") + for _, event := range events.get() { + assert.NotEqual(t, activity.DomainValidated, event.Activity, "a rejected write must not emit a validation event") + } + }) + }) + } +} + +func setupClockDomainTest(t *testing.T) *domainTestEnv { + t.Helper() + // Network driver watchers cannot share cancellation channels across synctest bubbles. + // Store boundary and concurrency tests still exercise the selected database engine. + t.Setenv("NETBIRD_STORE_ENGINE", "sqlite") + return setupDomainTest(t) +} + +func createExpiringDomain(t *testing.T, env *domainTestEnv, name string, expiresAt time.Time) *domain.Domain { + t.Helper() + d, err := env.store.CreateCustomDomain(context.Background(), accountA, name, testCluster, false) + require.NoError(t, err) + require.NoError(t, env.store.(*nbstore.SqlStore).GetDB().Model(d).Update("validation_expires_at", expiresAt).Error) + d.ValidationExpiresAt = &expiresAt + return d +} + +type domainEvents struct { + mu sync.Mutex + events []*activity.Event +} + +func captureDomainEvents(env *domainTestEnv) *domainEvents { + events := &domainEvents{} + env.manager.accountManager = &mock_server.MockAccountManager{ + StoreEventFunc: func(_ context.Context, initiator, target, account string, code activity.ActivityDescriber, meta map[string]any) { + if code == activity.DomainAdded { + return + } + events.mu.Lock() + defer events.mu.Unlock() + events.events = append(events.events, &activity.Event{ + InitiatorID: initiator, TargetID: target, AccountID: account, + Activity: code.(activity.Activity), Meta: meta, + }) + }, + } + return events +} + +func (e *domainEvents) get() []*activity.Event { + e.mu.Lock() + defer e.mu.Unlock() + return append([]*activity.Event(nil), e.events...) +} diff --git a/management/internals/modules/reverseproxy/domain/manager/manager.go b/management/internals/modules/reverseproxy/domain/manager/manager.go index 46e4ced83..c0fb12e9c 100644 --- a/management/internals/modules/reverseproxy/domain/manager/manager.go +++ b/management/internals/modules/reverseproxy/domain/manager/manager.go @@ -6,6 +6,7 @@ import ( "fmt" "net" "strings" + "time" log "github.com/sirupsen/logrus" @@ -18,6 +19,7 @@ import ( "github.com/netbirdio/netbird/management/server/permissions/operations" nbstore "github.com/netbirdio/netbird/management/server/store" "github.com/netbirdio/netbird/management/server/types" + nbdomain "github.com/netbirdio/netbird/shared/management/domain" "github.com/netbirdio/netbird/shared/management/status" ) @@ -32,6 +34,8 @@ type store interface { 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) DeleteCustomDomain(ctx context.Context, accountID string, domainID string) error + GetExpiredCustomDomains(ctx context.Context, now time.Time, afterID domain.ID, limit int) ([]*domain.Domain, error) + DeleteExpiredCustomDomain(ctx context.Context, d *domain.Domain, now time.Time) (bool, error) } type proxyManager interface { @@ -106,12 +110,13 @@ func (m Manager) GetDomains(ctx context.Context, accountID, userID string) ([]*d // Add custom domains. for _, d := range domains { cd := &domain.Domain{ - ID: d.ID, - Domain: d.Domain, - AccountID: accountID, - TargetCluster: d.TargetCluster, - Type: domain.TypeCustom, - Validated: d.Validated, + ID: d.ID, + Domain: d.Domain, + AccountID: accountID, + TargetCluster: d.TargetCluster, + Type: domain.TypeCustom, + Validated: d.Validated, + ValidationExpiresAt: d.ValidationExpiresAt, } if d.TargetCluster != "" { cd.SupportsCustomPorts = m.proxyManager.ClusterSupportsCustomPorts(ctx, d.TargetCluster) @@ -126,6 +131,7 @@ func (m Manager) GetDomains(ctx context.Context, accountID, userID string) ([]*d return ret, nil } +// CreateDomain registers a normalized custom domain and attempts DNS validation. func (m Manager) CreateDomain(ctx context.Context, accountID, userID, domainName, targetCluster string) (*domain.Domain, error) { ok, ctx, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, modules.Services, operations.Create) if err != nil { @@ -135,6 +141,15 @@ func (m Manager) CreateDomain(ctx context.Context, accountID, userID, domainName return nil, status.NewPermissionDeniedError() } + parsed, err := nbdomain.FromString(strings.TrimSuffix(domainName, ".")) + if err != nil { + return nil, status.Errorf(status.InvalidArgument, "invalid domain: %v", err) + } + domainName = parsed.PunycodeString() + if !nbdomain.IsValidDomainNoWildcard(domainName) { + return nil, status.Errorf(status.InvalidArgument, "invalid domain format") + } + // Verify the target cluster is in the available clusters for this account allowList, err := m.getClusterAllowList(ctx, accountID) if err != nil { @@ -243,6 +258,14 @@ func (m Manager) ValidateDomain(ctx context.Context, accountID, userID, domainID }).WithError(err).Error("get custom domain from store") return } + if d.Validated { + return + } + if d.ValidationExpiresAt == nil || !time.Now().Before(*d.ValidationExpiresAt) { + log.WithFields(log.Fields{"accountID": accountID, "domainID": domainID}). + Debug("custom domain validation window has expired") + return + } // Validate only against the domain's target cluster targetCluster := d.TargetCluster @@ -263,20 +286,21 @@ func (m Manager) ValidateDomain(ctx context.Context, accountID, userID, domainID }).Info("validating domain against target cluster") if m.validator.IsValid(context.Background(), d.Domain, []string{targetCluster}) { - log.WithFields(log.Fields{ - "accountID": accountID, - "domainID": domainID, - "domain": d.Domain, - }).Info("domain validated successfully") d.Validated = true if _, err := m.store.UpdateCustomDomain(context.Background(), accountID, d); err != nil { - log.WithFields(log.Fields{ + entry := log.WithFields(log.Fields{ "accountID": accountID, "domainID": domainID, - "domain": d.Domain, - }).WithError(err).Error("update custom domain in store") + }).WithError(err) + if sErr, ok := status.FromError(err); ok && sErr.Type() == status.PreconditionFailed { + entry.Debug("custom domain registration is no longer pending validation") + return + } + entry.Error("update custom domain in store") return } + log.WithFields(log.Fields{"accountID": accountID, "domainID": domainID}). + Info("custom domain validated successfully") m.accountManager.StoreEvent(context.Background(), userID, domainID, accountID, activity.DomainValidated, d.EventMeta()) } else { diff --git a/management/internals/modules/reverseproxy/domain/manager/manager_realstore_test.go b/management/internals/modules/reverseproxy/domain/manager/manager_realstore_test.go index 8a0b56171..5c973c40e 100644 --- a/management/internals/modules/reverseproxy/domain/manager/manager_realstore_test.go +++ b/management/internals/modules/reverseproxy/domain/manager/manager_realstore_test.go @@ -296,11 +296,8 @@ func TestValidateDomain_PermissionDeniedDoesNotValidate(t *testing.T) { assert.Error(t, err, "the domain must still be unservable") } -// Validation runs asynchronously, so it can finish after the domain was -// deleted and then write a stale row back. gorm's Save falls back to an insert -// when an update affects no rows, which would resurrect the domain as -// validated; UpdateCustomDomain avoids that by selecting explicit columns. -// This pins that behaviour, since dropping the Select would reintroduce it. +// A validation finishing after deletion must reject the stale write, without +// restoring the registration or reporting successful validation. func TestUpdateCustomDomain_DoesNotResurrectDeletedDomain(t *testing.T) { ctx := context.Background() env := setupDomainTest(t) @@ -315,11 +312,9 @@ func TestUpdateCustomDomain_DoesNotResurrectDeletedDomain(t *testing.T) { require.Nil(t, storedDomain(t, env.store, accountA, "racy.example.com"), "the domain should be gone") // What an in-flight validation would write once its CNAME check succeeded. - // The write has to succeed for the assertion below to mean anything: a - // rejected write would leave the domain absent for the wrong reason. stale.Validated = true _, err = env.store.UpdateCustomDomain(ctx, accountA, stale) - require.NoError(t, err, "the update itself must succeed, so absence is not just a failed write") + require.Error(t, err, "a deleted registration must reject a late validation") assert.Nil(t, storedDomain(t, env.store, accountA, "racy.example.com"), "a late validation write must not recreate a deleted domain") diff --git a/management/internals/modules/reverseproxy/domain/manager/manager_test.go b/management/internals/modules/reverseproxy/domain/manager/manager_test.go index 519f5efeb..1ca690130 100644 --- a/management/internals/modules/reverseproxy/domain/manager/manager_test.go +++ b/management/internals/modules/reverseproxy/domain/manager/manager_test.go @@ -4,6 +4,7 @@ import ( "context" "errors" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -208,6 +209,14 @@ func (s *stubStore) DeleteCustomDomain(context.Context, string, string) error { panic("not used in allow-list tests") } +func (s *stubStore) GetExpiredCustomDomains(context.Context, time.Time, domain.ID, int) ([]*domain.Domain, error) { + panic("not used in allow-list tests") +} + +func (s *stubStore) DeleteExpiredCustomDomain(context.Context, *domain.Domain, time.Time) (bool, error) { + panic("not used in allow-list tests") +} + // TestGetClusterAllowList_DedicatedGatewayAddressExcluded pins invariant (B)'s // chokepoint: a self-addressed settings pin reserves the account's gateway // address, so it is dropped from the allow list — which, because the diff --git a/management/internals/modules/reverseproxy/domain/manager/normalization_test.go b/management/internals/modules/reverseproxy/domain/manager/normalization_test.go new file mode 100644 index 000000000..da5e4bdd1 --- /dev/null +++ b/management/internals/modules/reverseproxy/domain/manager/normalization_test.go @@ -0,0 +1,83 @@ +package manager + +import ( + "context" + "strings" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/shared/management/status" +) + +func TestCreateDomain_NormalizesName(t *testing.T) { + for _, tt := range []struct { + name string + input string + canonical string + }{ + {"mixed case", "Apps.Example.COM", "apps.example.com"}, + {"unicode", "münchen.example.com", "xn--mnchen-3ya.example.com"}, + {"trailing dot", "apps.example.com.", "apps.example.com"}, + {"underscore", "My_App.example.com", "my_app.example.com"}, + } { + t.Run(tt.name, func(t *testing.T) { + ctx := context.Background() + env := setupDomainTest(t) + env.resolver.set("validation."+tt.canonical, testCluster) + + created, err := env.manager.CreateDomain(ctx, accountA, accountAUser, tt.input, testCluster) + require.NoError(t, err) + assert.Equal(t, tt.canonical, created.Domain, "the response must use the normalized name") + assert.True(t, created.Validated, "the CNAME lookup must use the normalized name") + stored, err := env.store.GetCustomDomain(ctx, accountA, created.ID) + require.NoError(t, err) + assert.Equal(t, tt.canonical, stored.Domain, "the database must retain the normalized name") + + _, err = env.manager.CreateDomain(ctx, accountB, accountBUser, tt.canonical, testCluster) + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok, "an equivalent name must return a typed conflict") + assert.Equal(t, status.AlreadyExists, sErr.Type(), "normalization must precede the availability check") + }) + } +} + +func TestCreateDomain_NormalizedNameCanValidateLater(t *testing.T) { + ctx := context.Background() + env := setupDomainTest(t) + created, err := env.manager.CreateDomain(ctx, accountA, accountAUser, "Apps.Example.COM.", testCluster) + require.NoError(t, err) + require.False(t, created.Validated, "a missing CNAME must leave the normalized registration pending") + + env.resolver.set("validation.apps.example.com", testCluster) + env.manager.ValidateDomain(ctx, accountA, accountAUser, created.ID) + stored, err := env.store.GetCustomDomain(ctx, accountA, created.ID) + require.NoError(t, err) + assert.Equal(t, "apps.example.com", stored.Domain, "retrying validation must retain the normalized name") + assert.True(t, stored.Validated, "later validation must look up the normalized name") +} + +func TestCreateDomain_RejectsInvalidName(t *testing.T) { + ctx := context.Background() + env := setupDomainTest(t) + for _, name := range []string{ + "", ".", "app..example.com", "app.example.com..", "-app.example.com", + "app%.example.com", "app!.example.com", "*.example.com", "app example.com", + "https://example.com", strings.Repeat("a", 64) + ".example.com", + } { + t.Run(name, func(t *testing.T) { + // A matching DNS response must not make a malformed name acceptable. + env.resolver.set("validation."+name, testCluster) + _, err := env.manager.CreateDomain(ctx, accountA, accountAUser, name, testCluster) + require.Error(t, err) + sErr, ok := status.FromError(err) + require.True(t, ok, "invalid names must return a typed client error") + assert.Equal(t, status.InvalidArgument, sErr.Type(), "malformed names must be rejected before storage") + }) + } + stored, err := env.store.ListCustomDomains(ctx, accountA) + require.NoError(t, err) + assert.Empty(t, stored, "invalid registration attempts must not reserve any names") +} diff --git a/management/internals/server/server.go b/management/internals/server/server.go index 9709d1099..a1b58fdf1 100644 --- a/management/internals/server/server.go +++ b/management/internals/server/server.go @@ -66,7 +66,8 @@ type BaseServer struct { disableLegacyManagementPort bool autoResolveDomains bool - proxyAuthClose func() + proxyAuthClose func() + domainCleanupStop func() // grpcExtensions holds additional gRPC services, interceptors, and shutdown // hooks registered by external modules via RegisterGRPCExtension. Populated @@ -227,9 +228,26 @@ func (s *BaseServer) Start(ctx context.Context) error { s.update.SetOnUpdateListener(func() { log.WithContext(ctx).Infof("your management version, \"%s\", is outdated, a new management version is available. Learn more here: https://github.com/netbirdio/netbird/releases", version.NetbirdVersion()) }) + s.startDomainCleanup(srvCtx) return nil } +func (s *BaseServer) startDomainCleanup(ctx context.Context) { + if s.domainCleanupStop != nil { + return + } + mgr := s.ReverseProxyDomainManager() + ctx, cancel := context.WithCancel(ctx) + done := make(chan struct{}) + s.domainCleanupStop = func() { + cancel() + <-done + } + go func() { + defer close(done) + mgr.RunValidationCleanup(ctx) + }() +} // setupTLS resolves the listener's TLS source: an injected config wins over the HttpConfig certificate settings func (s *BaseServer) setupTLS(ctx context.Context) (bool, error) { @@ -260,6 +278,9 @@ func (s *BaseServer) setupTLS(ctx context.Context) (bool, error) { func (s *BaseServer) Stop() error { ctx, cancel := context.WithTimeout(context.Background(), 5*time.Second) defer cancel() + if s.domainCleanupStop != nil { + s.domainCleanupStop() + } s.IntegratedValidator().Stop(ctx) if s.GeoLocationManager() != nil { diff --git a/management/server/activity/codes.go b/management/server/activity/codes.go index 4d50fee3e..dca043027 100644 --- a/management/server/activity/codes.go +++ b/management/server/activity/codes.go @@ -284,6 +284,9 @@ const ( // AgentNetworkSettingsDeleted indicates that a user deleted the Agent Network account settings, releasing the endpoint AgentNetworkSettingsDeleted Activity = 142 + // CustomDomainValidationExpired indicates that an unvalidated domain registration expired. + CustomDomainValidationExpired Activity = 143 + AccountDeleted Activity = 99999 ) @@ -461,9 +464,10 @@ var activityMap = map[Activity]Code{ AccountMetricsPushEnabled: {"Account metrics push enabled", "account.setting.metrics.push.enable"}, AccountMetricsPushDisabled: {"Account metrics push disabled", "account.setting.metrics.push.disable"}, - DomainAdded: {"Domain added", "domain.add"}, - DomainDeleted: {"Domain deleted", "domain.delete"}, - DomainValidated: {"Domain validated", "domain.validate"}, + DomainAdded: {"Domain added", "domain.add"}, + DomainDeleted: {"Domain deleted", "domain.delete"}, + DomainValidated: {"Domain validated", "domain.validate"}, + CustomDomainValidationExpired: {"Unvalidated domain registration expired", "domain.validation.expire"}, } // StringCode returns a string code of the activity diff --git a/management/server/activity/store/sql_store.go b/management/server/activity/store/sql_store.go index 73e8e295c..5b0f23dd0 100644 --- a/management/server/activity/store/sql_store.go +++ b/management/server/activity/store/sql_store.go @@ -165,16 +165,16 @@ func (store *Store) Get(ctx context.Context, accountID string, offset, limit int return store.processResult(ctx, events) } -// Save an event in the SQLite events table end encrypt the "email" element in meta map -func (store *Store) Save(_ context.Context, event *activity.Event) (*activity.Event, error) { +// Save persists an activity event and encrypts deleted user details using the caller's context. +func (store *Store) Save(ctx context.Context, event *activity.Event) (*activity.Event, error) { eventCopy := event.Copy() - meta, err := store.saveDeletedUserEmailAndNameInEncrypted(eventCopy) + meta, err := store.saveDeletedUserEmailAndNameInEncrypted(ctx, eventCopy) if err != nil { return nil, err } eventCopy.Meta = meta - if err = store.db.Create(eventCopy).Error; err != nil { + if err = store.db.WithContext(ctx).Create(eventCopy).Error; err != nil { return nil, err } @@ -183,7 +183,7 @@ func (store *Store) Save(_ context.Context, event *activity.Event) (*activity.Ev // saveDeletedUserEmailAndNameInEncrypted if the meta contains email and name then store it in encrypted way and delete // this item from meta map -func (store *Store) saveDeletedUserEmailAndNameInEncrypted(event *activity.Event) (map[string]any, error) { +func (store *Store) saveDeletedUserEmailAndNameInEncrypted(ctx context.Context, event *activity.Event) (map[string]any, error) { email, ok := event.Meta["email"] if !ok { return event.Meta, nil @@ -211,7 +211,7 @@ func (store *Store) saveDeletedUserEmailAndNameInEncrypted(event *activity.Event } deletedUser.Name = encryptedName - err = store.db.Clauses(clause.OnConflict{ + err = store.db.WithContext(ctx).Clauses(clause.OnConflict{ Columns: []clause.Column{{Name: "id"}}, DoUpdates: clause.AssignmentColumns([]string{"email", "name"}), }).Create(deletedUser).Error diff --git a/management/server/activity/store/sql_store_test.go b/management/server/activity/store/sql_store_test.go index d723f1623..0a68858e4 100644 --- a/management/server/activity/store/sql_store_test.go +++ b/management/server/activity/store/sql_store_test.go @@ -7,11 +7,49 @@ import ( "time" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "github.com/netbirdio/netbird/management/server/activity" "github.com/netbirdio/netbird/util/crypt" ) +func TestSave_CancellationWhileWaitingForConnection(t *testing.T) { + t.Setenv(storeEngineEnv, "sqlite") + key, err := crypt.GenerateKey() + require.NoError(t, err) + store, err := NewSqlStore(context.Background(), t.TempDir(), key) + require.NoError(t, err) + t.Cleanup(func() { assert.NoError(t, store.Close(context.Background())) }) + db, err := store.db.DB() + require.NoError(t, err) + conn, err := db.Conn(context.Background()) + require.NoError(t, err) + t.Cleanup(func() { _ = conn.Close() }) + ctx, cancel := context.WithTimeout(context.Background(), 50*time.Millisecond) + defer cancel() + result := make(chan error, 1) + go func() { + _, err := store.Save(ctx, &activity.Event{ + Timestamp: time.Now().UTC(), Activity: activity.CustomDomainValidationExpired, + AccountID: "account-id", TargetID: "domain-id", InitiatorID: activity.SystemInitiator, + }) + result <- err + }() + select { + case err := <-result: + assert.ErrorIs(t, err, context.DeadlineExceeded) + require.NoError(t, conn.Close()) + case <-time.After(time.Second): + // Release the connection so a regression cannot leave the writer running. + require.NoError(t, conn.Close()) + assert.ErrorIs(t, <-result, context.DeadlineExceeded) + t.Error("activity writes must stop waiting when their deadline expires") + } + events, err := store.Get(context.Background(), "account-id", 0, 10, true) + require.NoError(t, err) + assert.Empty(t, events, "a timed-out write must not persist after the connection is released") +} + func TestNewSqlStore(t *testing.T) { dataDir := t.TempDir() key, _ := crypt.GenerateKey() diff --git a/management/server/event.go b/management/server/event.go index 4211f2dda..2b1ac4987 100644 --- a/management/server/event.go +++ b/management/server/event.go @@ -61,23 +61,34 @@ func (am *DefaultAccountManager) GetEvents(ctx context.Context, accountID, userI return filtered, nil } +// StoreEvent records an activity, waiting for expiration events before cleanup can stop. func (am *DefaultAccountManager) StoreEvent(ctx context.Context, initiatorID, targetID, accountID string, activityID activity.ActivityDescriber, meta map[string]any) { - if isEnabled() { - go func() { - _, err := am.eventStore.Save(ctx, &activity.Event{ - Timestamp: time.Now().UTC(), - Activity: activityID.(activity.Activity), - InitiatorID: initiatorID, - TargetID: targetID, - AccountID: accountID, - Meta: meta, - }) - if err != nil { - // todo add metric - log.WithContext(ctx).Errorf("received an error while storing an activity event, error: %s", err) - } - }() + if !isEnabled() { + return } + eventStore := am.eventStore + save := func(ctx context.Context) { + _, err := eventStore.Save(ctx, &activity.Event{ + Timestamp: time.Now().UTC(), + Activity: activityID.(activity.Activity), + InitiatorID: initiatorID, + TargetID: targetID, + AccountID: accountID, + Meta: meta, + }) + if err != nil { + log.WithContext(ctx).Errorf("received an error while storing an activity event, error: %s", err) + } + } + if activityID == activity.CustomDomainValidationExpired { + // The domain is already deleted; shutdown must allow its audit write to finish. + ctx, cancel := context.WithTimeout(context.WithoutCancel(ctx), 5*time.Second) + defer cancel() + save(ctx) + return + } + // Request cancellation must not discard the audit record of a completed operation. + go save(context.WithoutCancel(ctx)) } type eventUserInfo struct { diff --git a/management/server/event_test.go b/management/server/event_test.go index 420e69866..6fb3f437f 100644 --- a/management/server/event_test.go +++ b/management/server/event_test.go @@ -6,10 +6,52 @@ import ( "time" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "github.com/netbirdio/netbird/management/server/activity" + activitystore "github.com/netbirdio/netbird/management/server/activity/store" + "github.com/netbirdio/netbird/util/crypt" ) +func TestStoreEvent_CanceledContext(t *testing.T) { + t.Setenv("NB_EVENT_ACTIVITY_LOG_ENABLED", "true") + t.Setenv("NB_ACTIVITY_EVENT_STORE_ENGINE", "sqlite") + for _, code := range []activity.Activity{activity.CustomDomainValidationExpired, activity.DomainAdded} { + t.Run(code.StringCode(), func(t *testing.T) { + dir := t.TempDir() + key, err := crypt.GenerateKey() + require.NoError(t, err) + eventStore, err := activitystore.NewSqlStore(context.Background(), dir, key) + require.NoError(t, err) + t.Cleanup(func() { assert.NoError(t, eventStore.Close(context.Background())) }) + manager := &DefaultAccountManager{eventStore: eventStore} + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + // The operation already succeeded when shutdown or the request cancels its context. + manager.StoreEvent(ctx, activity.SystemInitiator, "domain-id", "account-id", + code, map[string]any{"domain": "expired.example.com"}) + if code != activity.CustomDomainValidationExpired { + require.Eventually(t, func() bool { + events, err := eventStore.Get(context.Background(), "account-id", 0, 10, true) + return err == nil && len(events) == 1 + }, time.Second, time.Millisecond, "asynchronous events must survive request cancellation") + } + require.NoError(t, eventStore.Close(context.Background())) + + reopened, err := activitystore.NewSqlStore(context.Background(), dir, key) + require.NoError(t, err) + t.Cleanup(func() { assert.NoError(t, reopened.Close(context.Background())) }) + events, err := reopened.Get(context.Background(), "account-id", 0, 10, true) + require.NoError(t, err) + require.Len(t, events, 1, "the event must be persisted before shutdown closes the store") + assert.Equal(t, code, events[0].Activity, "persist the requested activity") + assert.Equal(t, "domain-id", events[0].TargetID, "retain the registration ID") + assert.Equal(t, "expired.example.com", events[0].Meta["domain"], "retain the domain name") + }) + } +} + func generateAndStoreEvents(t *testing.T, manager *DefaultAccountManager, typ activity.Activity, initiatorID, targetID, accountID string, count int) { t.Helper() diff --git a/management/server/http/testing/integration/events_handler_integration_test.go b/management/server/http/testing/integration/events_handler_integration_test.go index 6611b60ee..a051aabcd 100644 --- a/management/server/http/testing/integration/events_handler_integration_test.go +++ b/management/server/http/testing/integration/events_handler_integration_test.go @@ -3,6 +3,7 @@ package integration import ( + "context" "encoding/json" "net/http" "net/http/httptest" @@ -10,6 +11,7 @@ import ( "time" "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" "github.com/netbirdio/netbird/management/server/http/testing/testing_tools" "github.com/netbirdio/netbird/management/server/http/testing/testing_tools/channel" @@ -34,7 +36,7 @@ func Test_Events_GetAll(t *testing.T) { for _, user := range users { t.Run(user.name+" - Get all events", func(t *testing.T) { - apiHandler, _, _ := channel.BuildApiBlackBoxWithDBState(t, "../testdata/events.sql", nil, false) + apiHandler, accountManager, _ := channel.BuildApiBlackBoxWithDBState(t, "../testdata/events.sql", nil, false) // First, perform a mutation to generate an event (create a group as admin) groupBody, err := json.Marshal(&api.GroupRequest{Name: "eventTestGroup"}) @@ -44,7 +46,14 @@ func Test_Events_GetAll(t *testing.T) { createReq := testing_tools.BuildRequest(t, groupBody, http.MethodPost, "/api/groups", testing_tools.TestAdminId) createRecorder := httptest.NewRecorder() apiHandler.ServeHTTP(createRecorder, createReq) - assert.Equal(t, http.StatusOK, createRecorder.Code, "Failed to create group to generate event") + require.Equal(t, http.StatusOK, createRecorder.Code, "Failed to create group to generate event") + + // Group creation returns before its asynchronous audit write finishes. + require.EventuallyWithT(t, func(c *assert.CollectT) { + events, err := accountManager.GetEvents(context.Background(), testing_tools.TestAccountId, testing_tools.TestAdminId) + assert.NoError(c, err) + assert.NotEmpty(c, events, "wait for the group creation event before checking permissions") + }, time.Second, 10*time.Millisecond) // Now query events req := testing_tools.BuildRequest(t, []byte{}, http.MethodGet, "/api/events", user.userId) diff --git a/management/server/migration/migration_custom_domain.go b/management/server/migration/migration_custom_domain.go new file mode 100644 index 000000000..bac784358 --- /dev/null +++ b/management/server/migration/migration_custom_domain.go @@ -0,0 +1,22 @@ +package migration + +import ( + "context" + "fmt" + "time" + + "gorm.io/gorm" + + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain" +) + +// MigrateCustomDomainValidationExpiry gives existing pending registrations a validation window. +func MigrateCustomDomainValidationExpiry(ctx context.Context, db *gorm.DB) error { + result := db.WithContext(ctx).Model(&domain.Domain{}). + Where("validated = ? AND validation_expires_at IS NULL", false). + Update("validation_expires_at", time.Now().UTC().Add(domain.ValidationTTL)) + if result.Error != nil { + return fmt.Errorf("backfill custom domain validation expiry: %w", result.Error) + } + return nil +} diff --git a/management/server/migration/migration_custom_domain_test.go b/management/server/migration/migration_custom_domain_test.go new file mode 100644 index 000000000..facf23c9f --- /dev/null +++ b/management/server/migration/migration_custom_domain_test.go @@ -0,0 +1,44 @@ +package migration_test + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain" + "github.com/netbirdio/netbird/management/server/migration" +) + +func TestMigrateCustomDomainValidationExpiry(t *testing.T) { + db := setupDatabase(t) + require.NoError(t, db.AutoMigrate(&domain.Domain{})) + t.Cleanup(func() { require.NoError(t, db.Migrator().DropTable(&domain.Domain{})) }) + ctx := context.Background() + existingDeadline := time.Now().UTC().Add(time.Hour).Truncate(time.Second) + rows := []domain.Domain{ + {ID: "legacy", Domain: "legacy.example.com"}, + {ID: "validated", Domain: "validated.example.com", Validated: true}, + {ID: "pending", Domain: "pending.example.com", ValidationExpiresAt: &existingDeadline}, + } + require.NoError(t, db.Create(&rows).Error) + before := time.Now().UTC() + require.NoError(t, migration.MigrateCustomDomainValidationExpiry(ctx, db)) + after := time.Now().UTC() + var migrated domain.Domain + require.NoError(t, db.First(&migrated, "id = ?", "legacy").Error) + require.NotNil(t, migrated.ValidationExpiresAt) + assert.WithinRange(t, *migrated.ValidationExpiresAt, before.Truncate(time.Millisecond).Add(48*time.Hour), after.Add(48*time.Hour+time.Millisecond), "legacy pending registrations get a full window") + deadline := *migrated.ValidationExpiresAt + require.NoError(t, migration.MigrateCustomDomainValidationExpiry(ctx, db)) + require.NoError(t, db.First(&migrated, "id = ?", "legacy").Error) + assert.Equal(t, deadline, *migrated.ValidationExpiresAt, "repeated migration must not extend the deadline") + var validated, pending domain.Domain + require.NoError(t, db.First(&validated, "id = ?", "validated").Error) + require.NoError(t, db.First(&pending, "id = ?", "pending").Error) + assert.Nil(t, validated.ValidationExpiresAt, "validated domains do not acquire an expiry") + require.NotNil(t, pending.ValidationExpiresAt) + assert.WithinDuration(t, existingDeadline, *pending.ValidationExpiresAt, 0, "existing deadlines must be preserved") +} diff --git a/management/server/store/sql_store.go b/management/server/store/sql_store.go index 33c723a8a..08ec45395 100644 --- a/management/server/store/sql_store.go +++ b/management/server/store/sql_store.go @@ -5712,6 +5712,10 @@ func (s *SqlStore) CreateCustomDomain(ctx context.Context, accountID string, dom Type: domain.TypeCustom, Validated: validated, } + if !validated { + expiresAt := time.Now().UTC().Add(domain.ValidationTTL) + newDomain.ValidationExpiresAt = &expiresAt + } result := s.db.Create(newDomain) if result.Error != nil { // The unique index is the last guard when two requests clear the @@ -5733,12 +5737,21 @@ func (s *SqlStore) CreateCustomDomain(ctx context.Context, accountID string, dom return newDomain, nil } +// UpdateCustomDomain completes validation only while the original registration is pending. func (s *SqlStore) UpdateCustomDomain(ctx context.Context, accountID string, d *domain.Domain) (*domain.Domain, error) { - d.AccountID = accountID - result := s.db.Select("*").Save(d) + if !d.Validated { + return nil, status.Errorf(status.InvalidArgument, "custom domain update must complete validation") + } + result := s.db.WithContext(ctx).Model(&domain.Domain{}). + Where(accountAndIDQueryCondition, accountID, d.ID). + Where("domain = ? AND target_cluster = ?", d.Domain, d.TargetCluster). + Where("validated = ? AND validation_expires_at > ?", false, time.Now().UTC()). + Update("validated", true) if result.Error != nil { - log.WithContext(ctx).Errorf("failed to update reverse proxy custom domain to store: %v", result.Error) - return nil, status.Errorf(status.Internal, "failed to update reverse proxy custom domain to store") + return nil, fmt.Errorf("validate custom domain in store: %w", result.Error) + } + if result.RowsAffected == 0 { + return nil, status.Errorf(status.PreconditionFailed, "custom domain registration is no longer pending validation") } return d, nil diff --git a/management/server/store/sql_store_domain_expiration.go b/management/server/store/sql_store_domain_expiration.go new file mode 100644 index 000000000..525c54084 --- /dev/null +++ b/management/server/store/sql_store_domain_expiration.go @@ -0,0 +1,60 @@ +package store + +import ( + "context" + "fmt" + "strings" + "time" + + "gorm.io/gorm" + + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain" + rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + "github.com/netbirdio/netbird/shared/management/status" +) + +// GetExpiredCustomDomains lists pending registrations in stable batches across accounts. +func (s *SqlStore) GetExpiredCustomDomains(ctx context.Context, now time.Time, afterID domain.ID, limit int) ([]*domain.Domain, error) { + var domains []*domain.Domain + result := s.db.WithContext(ctx). + Where("validated = ? AND validation_expires_at <= ? AND id > ?", false, now, string(afterID)). + Order("id").Limit(limit).Find(&domains) + if result.Error != nil { + return nil, fmt.Errorf("list expired custom domains: %w", result.Error) + } + return domains, nil +} + +// DeleteExpiredCustomDomain deletes an expired registration only if no service uses its namespace. +func (s *SqlStore) DeleteExpiredCustomDomain(ctx context.Context, d *domain.Domain, now time.Time) (bool, error) { + db := s.db.WithContext(ctx) + services := customDomainServices(db, d) + result := db.Where(accountAndIDQueryCondition, d.AccountID, d.ID). + Where("domain = ? AND validated = ? AND validation_expires_at <= ?", d.Domain, false, now). + Where("NOT EXISTS (?)", services.Select("1")).Delete(&domain.Domain{}) + if result.Error != nil { + return false, fmt.Errorf("delete expired custom domain: %w", result.Error) + } + if result.RowsAffected > 0 { + return true, nil + } + var count int64 + if err := customDomainServices(db, d).Count(&count).Error; err != nil { + return false, fmt.Errorf("check expired custom domain services: %w", err) + } + if count > 0 { + return false, status.Errorf(status.PreconditionFailed, "expired custom domain still has dependent services") + } + return false, nil +} + +func customDomainServices(db *gorm.DB, d *domain.Domain) *gorm.DB { + name := strings.ToLower(strings.TrimSuffix(d.Domain, ".")) + // Shared domain validation permits underscores, and older rows may contain + // other LIKE metacharacters. + escaped := strings.NewReplacer("!", "!!", "%", "!%", "_", "!_").Replace(name) + return db.Model(&rpservice.Service{}).Where( + "LOWER(domain) IN ? OR LOWER(domain) LIKE ? ESCAPE '!' OR LOWER(domain) LIKE ? ESCAPE '!'", + []string{name, name + "."}, "%."+escaped, "%."+escaped+".", + ) +} diff --git a/management/server/store/sql_store_domain_expiration_test.go b/management/server/store/sql_store_domain_expiration_test.go new file mode 100644 index 000000000..7dcfb0562 --- /dev/null +++ b/management/server/store/sql_store_domain_expiration_test.go @@ -0,0 +1,78 @@ +package store + +import ( + "context" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain" + rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" +) + +func TestDeleteExpiredCustomDomain_ServiceDependencies(t *testing.T) { + runTestForAllEngines(t, "", func(t *testing.T, store Store) { + ctx := context.Background() + now := time.Now().UTC() + db := store.(*SqlStore).db + require.NoError(t, store.SaveAccount(ctx, newAccountWithId(ctx, "owner", "admin", ""))) + for _, tt := range []struct { + name string + domainName string + serviceHost string + protected bool + }{ + {"exact", "example.com", "example.com", true}, + {"subdomain", "example.com", "deep.app.example.com", true}, + {"case", "example.com", "APP.EXAMPLE.COM.", true}, + {"suffix-boundary", "example.com", "notexample.com", false}, + {"literal underscore", "a_b.example.com", "app.a_b.example.com", true}, + {"underscore wildcard", "a_b.example.com", "app.axb.example.com", false}, + {"legacy percent wildcard", "a%b.example.com", "app.axxb.example.com", false}, + {"legacy escape character", "a!b.example.com", "app.ab.example.com", false}, + } { + t.Run(tt.name, func(t *testing.T) { + d, err := store.CreateCustomDomain(ctx, "owner", tt.domainName, "cluster", false) + require.NoError(t, err) + require.NoError(t, db.Model(d).Update("validation_expires_at", now.Add(-time.Hour)).Error) + svc := &rpservice.Service{ID: "legacy", AccountID: "owner", Domain: tt.serviceHost} + require.NoError(t, store.CreateService(ctx, svc)) + deleted, err := store.DeleteExpiredCustomDomain(ctx, d, now) + if tt.protected { + require.Error(t, err) + assert.False(t, deleted, "service namespaces must remain reserved") + } else { + require.NoError(t, err) + assert.True(t, deleted, "a hostname outside the namespace must not prevent cleanup") + } + require.NoError(t, db.Delete(svc).Error) + require.NoError(t, db.Delete(d).Error) + }) + } + }) +} + +func TestDeleteExpiredCustomDomain_RechecksValidation(t *testing.T) { + runTestForAllEngines(t, "", func(t *testing.T, store Store) { + ctx := context.Background() + require.NoError(t, store.SaveAccount(ctx, newAccountWithId(ctx, "owner", "admin", ""))) + d, err := store.CreateCustomDomain(ctx, "owner", "validated.example.com", "cluster", false) + require.NoError(t, err) + d, err = store.GetCustomDomain(ctx, "owner", d.ID) + require.NoError(t, err) + stale := d.Copy() + d.Validated = true + _, err = store.UpdateCustomDomain(ctx, "owner", d) + require.NoError(t, err) + deleted, err := store.DeleteExpiredCustomDomain(ctx, stale, time.Now().Add(domain.ValidationTTL)) + require.NoError(t, err) + assert.False(t, deleted, "a stale cleanup candidate must not delete a validated registration") + stored, err := store.GetCustomDomain(ctx, "owner", d.ID) + require.NoError(t, err) + assert.True(t, stored.Validated, "the validated registration must remain usable") + require.NotNil(t, stored.ValidationExpiresAt) + assert.Equal(t, stale.ValidationExpiresAt, stored.ValidationExpiresAt, "validation must preserve the original deadline") + }) +} diff --git a/management/server/store/sql_store_get_account_test.go b/management/server/store/sql_store_get_account_test.go index 686839b1f..aca8d8a12 100644 --- a/management/server/store/sql_store_get_account_test.go +++ b/management/server/store/sql_store_get_account_test.go @@ -64,7 +64,7 @@ func assertGetAccountLoadsCustomDomains(t *testing.T, store Store) { _, err := store.CreateCustomDomain(ctx, accountID, "example.com", "eu.proxy.netbird.io", true) require.NoError(t, err, "creating the first custom domain must succeed") - _, err = store.CreateCustomDomain(ctx, accountID, "apps.acme.io", "us.proxy.netbird.io", false) + pending, err := store.CreateCustomDomain(ctx, accountID, "apps.acme.io", "us.proxy.netbird.io", false) require.NoError(t, err, "creating the second custom domain must succeed") account, err := store.GetAccount(ctx, accountID) @@ -75,6 +75,10 @@ func assertGetAccountLoadsCustomDomains(t *testing.T, store Store) { for _, d := range account.Domains { require.NotNil(t, d) byDomain[d.Domain] = d.TargetCluster + if d.ID == pending.ID { + require.NotNil(t, d.ValidationExpiresAt) + assert.WithinDuration(t, *pending.ValidationExpiresAt, *d.ValidationExpiresAt, time.Millisecond, "both account loaders must preserve the validation deadline") + } } assert.Equal(t, "eu.proxy.netbird.io", byDomain["example.com"], "custom domain must carry its target cluster") assert.Equal(t, "us.proxy.netbird.io", byDomain["apps.acme.io"], "custom domain must carry its target cluster") diff --git a/management/server/store/store.go b/management/server/store/store.go index da2b3c6e0..6886536b9 100644 --- a/management/server/store/store.go +++ b/management/server/store/store.go @@ -305,6 +305,8 @@ type Store interface { 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) + GetExpiredCustomDomains(ctx context.Context, now time.Time, afterID domain.ID, limit int) ([]*domain.Domain, error) + DeleteExpiredCustomDomain(ctx context.Context, d *domain.Domain, now time.Time) (bool, error) DeleteCustomDomain(ctx context.Context, accountID string, domainID string) error CreateAccessLog(ctx context.Context, log *accesslogs.AccessLogEntry) error @@ -642,6 +644,9 @@ func migratePostAuto(ctx context.Context, db *gorm.DB) error { func getMigrationsPostAuto(ctx context.Context) []migrationFunc { return []migrationFunc{ + func(db *gorm.DB) error { + return migration.MigrateCustomDomainValidationExpiry(ctx, db) + }, func(db *gorm.DB) error { return migration.CreateIndexIfNotExists[nbpeer.Peer](ctx, db, "idx_account_ip", "account_id", "ip") }, diff --git a/management/server/store/store_mock.go b/management/server/store/store_mock.go index 9bf49f076..04f79d30a 100644 --- a/management/server/store/store_mock.go +++ b/management/server/store/store_mock.go @@ -555,6 +555,21 @@ func (mr *MockStoreMockRecorder) DeleteDNSRecord(ctx, accountID, zoneID, recordI return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteDNSRecord", reflect.TypeOf((*MockStore)(nil).DeleteDNSRecord), ctx, accountID, zoneID, recordID) } +// DeleteExpiredCustomDomain mocks base method. +func (m *MockStore) DeleteExpiredCustomDomain(ctx context.Context, d *domain.Domain, now time.Time) (bool, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "DeleteExpiredCustomDomain", ctx, d, now) + ret0, _ := ret[0].(bool) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// DeleteExpiredCustomDomain indicates an expected call of DeleteExpiredCustomDomain. +func (mr *MockStoreMockRecorder) DeleteExpiredCustomDomain(ctx, d, now any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteExpiredCustomDomain", reflect.TypeOf((*MockStore)(nil).DeleteExpiredCustomDomain), ctx, d, now) +} + // DeleteGroup mocks base method. func (m *MockStore) DeleteGroup(ctx context.Context, accountID, groupID string) error { m.ctrl.T.Helper() @@ -2002,6 +2017,21 @@ func (mr *MockStoreMockRecorder) GetEmbeddedProxyPeerIDsByCluster(ctx, accountID return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetEmbeddedProxyPeerIDsByCluster", reflect.TypeOf((*MockStore)(nil).GetEmbeddedProxyPeerIDsByCluster), ctx, accountID) } +// GetExpiredCustomDomains mocks base method. +func (m *MockStore) GetExpiredCustomDomains(ctx context.Context, now time.Time, afterID domain.ID, limit int) ([]*domain.Domain, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "GetExpiredCustomDomains", ctx, now, afterID, limit) + ret0, _ := ret[0].([]*domain.Domain) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// GetExpiredCustomDomains indicates an expected call of GetExpiredCustomDomains. +func (mr *MockStoreMockRecorder) GetExpiredCustomDomains(ctx, now, afterID, limit any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetExpiredCustomDomains", reflect.TypeOf((*MockStore)(nil).GetExpiredCustomDomains), ctx, now, afterID, limit) +} + // GetExpiredEphemeralServices mocks base method. func (m *MockStore) GetExpiredEphemeralServices(ctx context.Context, ttl time.Duration, limit int) ([]*service.Service, error) { m.ctrl.T.Helper()