From 782c9434106f68fe83e45d612d756896a8179237 Mon Sep 17 00:00:00 2001 From: Pascal Fischer <32096965+pascal-fischer@users.noreply.github.com> Date: Tue, 29 Sep 2026 00:45:55 +0200 Subject: [PATCH] [management] extract shared db conn + data repository (#7649) * extract shared db conn + data repository * extract repository interface * protect against nested transactions * fix mysql and db conn creation * fix context management * remove query warpper * remove withContext and withLock wrapper * remove context from function call * remove in memory mode * fix nested transaction handling * use db directly * remove leftover test * remove pool close on error during conn creation --- .../network_map_db/sqlite_test_store.go | 8 +- .../accesslogs/manager/manager.go | 11 +- .../accesslogs/manager/manager_test.go | 86 ++-- .../accesslogs/manager/repository.go} | 98 ++--- .../accesslogs/manager/repository_test.go | 125 ++++++ .../reverseproxy/accesslogs/repository.go | 18 + .../accesslogs/repository_mock.go | 102 +++++ management/internals/server/boot.go | 16 +- management/internals/shared/db/conn.go | 121 ++++++ management/internals/shared/db/conn_test.go | 148 +++++++ .../internals/shared/db/dbtest/dbtest.go | 23 + .../internals/shared/db/dbtest/dbtest_test.go | 31 ++ management/internals/shared/db/engine.go | 10 + management/internals/shared/db/lock.go | 12 + management/internals/shared/db/open.go | 160 +++++++ management/internals/shared/db/open_test.go | 12 + management/internals/shared/db/transaction.go | 105 +++++ .../testing/testing_tools/channel/channel.go | 10 +- management/server/store/sql_store.go | 397 ++++-------------- management/server/store/sql_store_account.go | 8 +- .../store/sql_store_account_onboarding.go | 2 +- .../sql_store_agent_network_access_log.go | 4 +- .../store/sql_store_agent_network_usage.go | 2 +- .../server/store/sql_store_agentnetwork.go | 2 +- management/server/store/sql_store_group.go | 6 +- .../server/store/sql_store_group_peer.go | 2 +- .../server/store/sql_store_idp_migration.go | 10 +- .../store/sql_store_name_server_group.go | 2 +- management/server/store/sql_store_network.go | 2 +- .../store/sql_store_network_resource.go | 2 +- .../server/store/sql_store_network_router.go | 2 +- management/server/store/sql_store_peer.go | 4 +- .../store/sql_store_personal_access_token.go | 2 +- management/server/store/sql_store_policy.go | 4 +- .../server/store/sql_store_policy_rule.go | 2 +- .../server/store/sql_store_posture_checks.go | 2 +- management/server/store/sql_store_route.go | 2 +- management/server/store/sql_store_service.go | 4 +- .../server/store/sql_store_service_target.go | 2 +- .../server/store/sql_store_setup_key.go | 2 +- management/server/store/sql_store_test.go | 88 +++- management/server/store/sql_store_user.go | 4 +- .../server/store/sqlstore_bench_test.go | 10 +- management/server/store/store.go | 84 ++-- management/server/store/store_mock.go | 52 +-- management/server/types/store.go | 10 +- 46 files changed, 1247 insertions(+), 562 deletions(-) rename management/{server/store/sql_store_access_log.go => internals/modules/reverseproxy/accesslogs/manager/repository.go} (52%) create mode 100644 management/internals/modules/reverseproxy/accesslogs/manager/repository_test.go create mode 100644 management/internals/modules/reverseproxy/accesslogs/repository.go create mode 100644 management/internals/modules/reverseproxy/accesslogs/repository_mock.go create mode 100644 management/internals/shared/db/conn.go create mode 100644 management/internals/shared/db/conn_test.go create mode 100644 management/internals/shared/db/dbtest/dbtest.go create mode 100644 management/internals/shared/db/dbtest/dbtest_test.go create mode 100644 management/internals/shared/db/engine.go create mode 100644 management/internals/shared/db/lock.go create mode 100644 management/internals/shared/db/open.go create mode 100644 management/internals/shared/db/open_test.go create mode 100644 management/internals/shared/db/transaction.go diff --git a/integration_tests/management/network_map_db/sqlite_test_store.go b/integration_tests/management/network_map_db/sqlite_test_store.go index 1c70c93d4..622ac1c83 100644 --- a/integration_tests/management/network_map_db/sqlite_test_store.go +++ b/integration_tests/management/network_map_db/sqlite_test_store.go @@ -9,8 +9,8 @@ import ( "strings" networkmap_sqlite "github.com/netbirdio/netbird/management/internals/network_map_db/sqlite" + nbdb "github.com/netbirdio/netbird/management/internals/shared/db" gormstore "github.com/netbirdio/netbird/management/server/store" - "github.com/netbirdio/netbird/management/server/types" log "github.com/sirupsen/logrus" "gorm.io/driver/sqlite" "gorm.io/gorm" @@ -28,7 +28,11 @@ func createSqliteTestStore(baseData string) (*networkmap_sqlite.SqliteStore, fun if err != nil { log.Fatalf("error initializing db: %s", err.Error()) } - _, err = gormstore.NewSqlStore(context.TODO(), db, types.SqliteStoreEngine, nil, false) + conn, err := nbdb.NewConn(context.TODO(), db, nbdb.SqliteStoreEngine, nil) + if err != nil { + log.Fatalf("error initializing db: %s", err.Error()) + } + _, err = gormstore.NewSqlStore(context.TODO(), conn, nil, false) if err != nil { log.Fatalf("error initializing db: %s", err.Error()) } diff --git a/management/internals/modules/reverseproxy/accesslogs/manager/manager.go b/management/internals/modules/reverseproxy/accesslogs/manager/manager.go index 4b24adaa1..d22cdb980 100644 --- a/management/internals/modules/reverseproxy/accesslogs/manager/manager.go +++ b/management/internals/modules/reverseproxy/accesslogs/manager/manager.go @@ -9,6 +9,7 @@ import ( "github.com/netbirdio/netbird/management/internals/modules/agentnetwork" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs" + "github.com/netbirdio/netbird/management/internals/shared/db" "github.com/netbirdio/netbird/management/server/geolocation" "github.com/netbirdio/netbird/management/server/permissions" "github.com/netbirdio/netbird/management/server/permissions/modules" @@ -18,14 +19,16 @@ import ( ) type managerImpl struct { + repo accesslogs.Repository store store.Store permissionsManager permissions.Manager geo geolocation.Geolocation cleanupCancel context.CancelFunc } -func NewManager(store store.Store, permissionsManager permissions.Manager, geo geolocation.Geolocation) accesslogs.Manager { +func NewManager(repo accesslogs.Repository, store store.Store, permissionsManager permissions.Manager, geo geolocation.Geolocation) accesslogs.Manager { return &managerImpl{ + repo: repo, store: store, permissionsManager: permissionsManager, geo: geo, @@ -54,7 +57,7 @@ func (m *managerImpl) SaveAccessLog(ctx context.Context, logEntry *accesslogs.Ac } } - if err := m.store.CreateAccessLog(ctx, logEntry); err != nil { + if err := m.repo.Create(ctx, logEntry); err != nil { log.WithContext(ctx).WithFields(log.Fields{ "service_id": logEntry.ServiceID, "method": logEntry.Method, @@ -82,7 +85,7 @@ func (m *managerImpl) GetAllAccessLogs(ctx context.Context, accountID, userID st log.WithContext(ctx).Warnf("failed to resolve user filters: %v", err) } - logs, totalCount, err := m.store.GetAccountAccessLogs(ctx, store.LockingStrengthNone, accountID, *filter) + logs, totalCount, err := m.repo.ListByAccount(ctx, db.LockingStrengthNone, accountID, *filter) if err != nil { return nil, 0, err } @@ -98,7 +101,7 @@ func (m *managerImpl) CleanupOldAccessLogs(ctx context.Context, retentionDays in } cutoffTime := time.Now().AddDate(0, 0, -retentionDays) - deletedCount, err := m.store.DeleteOldAccessLogs(ctx, cutoffTime) + deletedCount, err := m.repo.DeleteOlderThan(ctx, cutoffTime) if err != nil { log.WithContext(ctx).Errorf("failed to cleanup old access logs: %v", err) return 0, err diff --git a/management/internals/modules/reverseproxy/accesslogs/manager/manager_test.go b/management/internals/modules/reverseproxy/accesslogs/manager/manager_test.go index 8e941d7e5..83dab7df6 100644 --- a/management/internals/modules/reverseproxy/accesslogs/manager/manager_test.go +++ b/management/internals/modules/reverseproxy/accesslogs/manager/manager_test.go @@ -5,27 +5,27 @@ import ( "testing" "time" - "go.uber.org/mock/gomock" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "go.uber.org/mock/gomock" - "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs" ) func TestCleanupOldAccessLogs(t *testing.T) { tests := []struct { name string retentionDays int - setupMock func(*store.MockStore) + setupMock func(*accesslogs.MockRepository) expectedCount int64 expectedError bool }{ { name: "cleanup logs older than retention period", retentionDays: 30, - setupMock: func(mockStore *store.MockStore) { - mockStore.EXPECT(). - DeleteOldAccessLogs(gomock.Any(), gomock.Any()). + setupMock: func(mockRepo *accesslogs.MockRepository) { + mockRepo.EXPECT(). + DeleteOlderThan(gomock.Any(), gomock.Any()). DoAndReturn(func(ctx context.Context, olderThan time.Time) (int64, error) { expectedCutoff := time.Now().AddDate(0, 0, -30) timeDiff := olderThan.Sub(expectedCutoff) @@ -41,9 +41,9 @@ func TestCleanupOldAccessLogs(t *testing.T) { { name: "no logs to cleanup", retentionDays: 30, - setupMock: func(mockStore *store.MockStore) { - mockStore.EXPECT(). - DeleteOldAccessLogs(gomock.Any(), gomock.Any()). + setupMock: func(mockRepo *accesslogs.MockRepository) { + mockRepo.EXPECT(). + DeleteOlderThan(gomock.Any(), gomock.Any()). Return(int64(0), nil) }, expectedCount: 0, @@ -52,8 +52,8 @@ func TestCleanupOldAccessLogs(t *testing.T) { { name: "zero retention days skips cleanup", retentionDays: 0, - setupMock: func(mockStore *store.MockStore) { - // No expectations - DeleteOldAccessLogs should not be called + setupMock: func(mockRepo *accesslogs.MockRepository) { + // No expectations - DeleteOlderThan should not be called }, expectedCount: 0, expectedError: false, @@ -61,8 +61,8 @@ func TestCleanupOldAccessLogs(t *testing.T) { { name: "negative retention days skips cleanup", retentionDays: -10, - setupMock: func(mockStore *store.MockStore) { - // No expectations - DeleteOldAccessLogs should not be called + setupMock: func(mockRepo *accesslogs.MockRepository) { + // No expectations - DeleteOlderThan should not be called }, expectedCount: 0, expectedError: false, @@ -74,11 +74,11 @@ func TestCleanupOldAccessLogs(t *testing.T) { ctrl := gomock.NewController(t) defer ctrl.Finish() - mockStore := store.NewMockStore(ctrl) - tt.setupMock(mockStore) + mockRepo := accesslogs.NewMockRepository(ctrl) + tt.setupMock(mockRepo) manager := &managerImpl{ - store: mockStore, + repo: mockRepo, } ctx := context.Background() @@ -98,10 +98,10 @@ func TestCleanupWithExactBoundary(t *testing.T) { ctrl := gomock.NewController(t) defer ctrl.Finish() - mockStore := store.NewMockStore(ctrl) + mockRepo := accesslogs.NewMockRepository(ctrl) - mockStore.EXPECT(). - DeleteOldAccessLogs(gomock.Any(), gomock.Any()). + mockRepo.EXPECT(). + DeleteOlderThan(gomock.Any(), gomock.Any()). DoAndReturn(func(ctx context.Context, olderThan time.Time) (int64, error) { expectedCutoff := time.Now().AddDate(0, 0, -30) timeDiff := olderThan.Sub(expectedCutoff) @@ -110,7 +110,7 @@ func TestCleanupWithExactBoundary(t *testing.T) { }) manager := &managerImpl{ - store: mockStore, + repo: mockRepo, } ctx := context.Background() @@ -125,11 +125,11 @@ func TestStartPeriodicCleanup(t *testing.T) { ctrl := gomock.NewController(t) defer ctrl.Finish() - mockStore := store.NewMockStore(ctrl) + mockRepo := accesslogs.NewMockRepository(ctrl) // No expectations - cleanup should not run manager := &managerImpl{ - store: mockStore, + repo: mockRepo, } ctx, cancel := context.WithCancel(context.Background()) @@ -139,22 +139,22 @@ func TestStartPeriodicCleanup(t *testing.T) { time.Sleep(100 * time.Millisecond) - // If DeleteOldAccessLogs was called, the test will fail due to unexpected call + // If DeleteOlderThan was called, the test will fail due to unexpected call }) t.Run("periodic cleanup runs immediately on start", func(t *testing.T) { ctrl := gomock.NewController(t) defer ctrl.Finish() - mockStore := store.NewMockStore(ctrl) + mockRepo := accesslogs.NewMockRepository(ctrl) - mockStore.EXPECT(). - DeleteOldAccessLogs(gomock.Any(), gomock.Any()). + mockRepo.EXPECT(). + DeleteOlderThan(gomock.Any(), gomock.Any()). Return(int64(2), nil). Times(1) manager := &managerImpl{ - store: mockStore, + repo: mockRepo, } ctx, cancel := context.WithCancel(context.Background()) @@ -171,15 +171,15 @@ func TestStartPeriodicCleanup(t *testing.T) { ctrl := gomock.NewController(t) defer ctrl.Finish() - mockStore := store.NewMockStore(ctrl) + mockRepo := accesslogs.NewMockRepository(ctrl) - mockStore.EXPECT(). - DeleteOldAccessLogs(gomock.Any(), gomock.Any()). + mockRepo.EXPECT(). + DeleteOlderThan(gomock.Any(), gomock.Any()). Return(int64(1), nil). Times(1) manager := &managerImpl{ - store: mockStore, + repo: mockRepo, } ctx, cancel := context.WithCancel(context.Background()) @@ -198,15 +198,15 @@ func TestStartPeriodicCleanup(t *testing.T) { ctrl := gomock.NewController(t) defer ctrl.Finish() - mockStore := store.NewMockStore(ctrl) + mockRepo := accesslogs.NewMockRepository(ctrl) - mockStore.EXPECT(). - DeleteOldAccessLogs(gomock.Any(), gomock.Any()). + mockRepo.EXPECT(). + DeleteOlderThan(gomock.Any(), gomock.Any()). Return(int64(0), nil). Times(1) manager := &managerImpl{ - store: mockStore, + repo: mockRepo, } ctx, cancel := context.WithCancel(context.Background()) @@ -223,15 +223,15 @@ func TestStartPeriodicCleanup(t *testing.T) { ctrl := gomock.NewController(t) defer ctrl.Finish() - mockStore := store.NewMockStore(ctrl) + mockRepo := accesslogs.NewMockRepository(ctrl) - mockStore.EXPECT(). - DeleteOldAccessLogs(gomock.Any(), gomock.Any()). + mockRepo.EXPECT(). + DeleteOlderThan(gomock.Any(), gomock.Any()). Return(int64(3), nil). Times(1) manager := &managerImpl{ - store: mockStore, + repo: mockRepo, } ctx, cancel := context.WithCancel(context.Background()) @@ -249,15 +249,15 @@ func TestStopPeriodicCleanup(t *testing.T) { ctrl := gomock.NewController(t) defer ctrl.Finish() - mockStore := store.NewMockStore(ctrl) + mockRepo := accesslogs.NewMockRepository(ctrl) - mockStore.EXPECT(). - DeleteOldAccessLogs(gomock.Any(), gomock.Any()). + mockRepo.EXPECT(). + DeleteOlderThan(gomock.Any(), gomock.Any()). Return(int64(1), nil). Times(1) manager := &managerImpl{ - store: mockStore, + repo: mockRepo, } ctx := context.Background() diff --git a/management/server/store/sql_store_access_log.go b/management/internals/modules/reverseproxy/accesslogs/manager/repository.go similarity index 52% rename from management/server/store/sql_store_access_log.go rename to management/internals/modules/reverseproxy/accesslogs/manager/repository.go index 56092eb74..105696ab7 100644 --- a/management/server/store/sql_store_access_log.go +++ b/management/internals/modules/reverseproxy/accesslogs/manager/repository.go @@ -1,4 +1,4 @@ -package store +package manager import ( "context" @@ -10,91 +10,78 @@ import ( "gorm.io/gorm/clause" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs" + "github.com/netbirdio/netbird/management/internals/shared/db" "github.com/netbirdio/netbird/shared/management/status" ) -// CreateAccessLog creates a new access log entry in the database -func (s *SqlStore) CreateAccessLog(ctx context.Context, logEntry *accesslogs.AccessLogEntry) error { - result := s.db.Create(logEntry) - if result.Error != nil { +type sqlRepository struct { + conn *db.Conn + db *gorm.DB +} + +// NewRepository returns the access log repository backed by conn. +func NewRepository(conn *db.Conn) accesslogs.Repository { + return &sqlRepository{conn: conn, db: conn.DB(nil)} +} + +func (r *sqlRepository) WithTx(tx *db.Tx) accesslogs.Repository { + return &sqlRepository{conn: r.conn, db: r.conn.DB(tx)} +} + +func (r *sqlRepository) Create(ctx context.Context, entry *accesslogs.AccessLogEntry) error { + if err := r.db.Create(entry).Error; err != nil { log.WithContext(ctx).WithFields(log.Fields{ - "service_id": logEntry.ServiceID, - "method": logEntry.Method, - "host": logEntry.Host, - "path": logEntry.Path, - }).Errorf("failed to create access log entry in store: %v", result.Error) + "service_id": entry.ServiceID, + "method": entry.Method, + "host": entry.Host, + "path": entry.Path, + }).Errorf("failed to create access log entry in store: %v", err) return status.Errorf(status.Internal, "failed to create access log entry in store") } return nil } -// GetAccountAccessLogs retrieves access logs for a given account with pagination and filtering -func (s *SqlStore) GetAccountAccessLogs(ctx context.Context, lockStrength LockingStrength, accountID string, filter accesslogs.AccessLogFilter) ([]*accesslogs.AccessLogEntry, int64, error) { - var logs []*accesslogs.AccessLogEntry +// ListByAccount returns one page of an account's access logs together with the +// total number of entries matching the filter. +func (r *sqlRepository) ListByAccount(ctx context.Context, lockStrength db.LockingStrength, accountID string, filter accesslogs.AccessLogFilter) ([]*accesslogs.AccessLogEntry, int64, error) { var totalCount int64 - - baseQuery := s.db. - Model(&accesslogs.AccessLogEntry{}). - Where(accountIDCondition, accountID) - - baseQuery = s.applyAccessLogFilters(baseQuery, filter) - - if err := baseQuery.Count(&totalCount).Error; err != nil { + countQuery := applyFilters(r.db.Model(&accesslogs.AccessLogEntry{}).Where("account_id = ?", accountID), filter) + if err := countQuery.Count(&totalCount).Error; err != nil { log.WithContext(ctx).Errorf("failed to count access logs: %v", err) return nil, 0, status.Errorf(status.Internal, "failed to count access logs") } - query := s.db. - Where(accountIDCondition, accountID) - - query = s.applyAccessLogFilters(query, filter) - - sortColumns := filter.GetSortColumn() + query := applyFilters(r.db.Where("account_id = ?", accountID), filter) sortOrder := strings.ToUpper(filter.GetSortOrder()) - - var orderClauses []string - for _, col := range strings.Split(sortColumns, ",") { - col = strings.TrimSpace(col) - if col != "" { - orderClauses = append(orderClauses, col+" "+sortOrder) + for _, column := range strings.Split(filter.GetSortColumn(), ",") { + if column = strings.TrimSpace(column); column != "" { + query = query.Order(column + " " + sortOrder) } } - orderClause := strings.Join(orderClauses, ", ") - - query = query. - Order(orderClause). - Limit(filter.GetLimit()). - Offset(filter.GetOffset()) - - if lockStrength != LockingStrengthNone { + query = query.Limit(filter.GetLimit()).Offset(filter.GetOffset()) + if lockStrength != db.LockingStrengthNone { query = query.Clauses(clause.Locking{Strength: string(lockStrength)}) } - result := query.Find(&logs) - if result.Error != nil { - log.WithContext(ctx).Errorf("failed to get access logs from store: %v", result.Error) + var logs []*accesslogs.AccessLogEntry + if err := query.Find(&logs).Error; err != nil { + log.WithContext(ctx).Errorf("failed to get access logs from store: %v", err) return nil, 0, status.Errorf(status.Internal, "failed to get access logs from store") } return logs, totalCount, nil } -// DeleteOldAccessLogs deletes all access logs older than the specified time -func (s *SqlStore) DeleteOldAccessLogs(ctx context.Context, olderThan time.Time) (int64, error) { - result := s.db. - Where("timestamp < ?", olderThan). - Delete(&accesslogs.AccessLogEntry{}) - +func (r *sqlRepository) DeleteOlderThan(ctx context.Context, olderThan time.Time) (int64, error) { + result := r.db.Where("timestamp < ?", olderThan).Delete(&accesslogs.AccessLogEntry{}) if result.Error != nil { log.WithContext(ctx).Errorf("failed to delete old access logs: %v", result.Error) return 0, status.Errorf(status.Internal, "failed to delete old access logs") } - return result.RowsAffected, nil } -// applyAccessLogFilters applies filter conditions to the query -func (s *SqlStore) applyAccessLogFilters(query *gorm.DB, filter accesslogs.AccessLogFilter) *gorm.DB { +func applyFilters(query *gorm.DB, filter accesslogs.AccessLogFilter) *gorm.DB { if filter.Search != nil { searchPattern := "%" + *filter.Search + "%" query = query.Where( @@ -112,7 +99,6 @@ func (s *SqlStore) applyAccessLogFilters(query *gorm.DB, filter accesslogs.Acces } if filter.Path != nil { - // Support LIKE pattern for path filtering query = query.Where("path LIKE ?", "%"+*filter.Path+"%") } @@ -127,9 +113,9 @@ func (s *SqlStore) applyAccessLogFilters(query *gorm.DB, filter accesslogs.Acces if filter.Status != nil { switch *filter.Status { case "success": - query = query.Where("status_code >= ? AND status_code < ?", 200, 400) + query = query.Where("(status_code >= ? AND status_code < ?)", 200, 400) case "failed": - query = query.Where("status_code < ? OR status_code >= ?", 200, 400) + query = query.Where("((status_code >= ? AND status_code < ?) OR status_code >= ?)", 100, 200, 400) } } diff --git a/management/internals/modules/reverseproxy/accesslogs/manager/repository_test.go b/management/internals/modules/reverseproxy/accesslogs/manager/repository_test.go new file mode 100644 index 000000000..9931cc27c --- /dev/null +++ b/management/internals/modules/reverseproxy/accesslogs/manager/repository_test.go @@ -0,0 +1,125 @@ +package manager + +import ( + "context" + "errors" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs" + "github.com/netbirdio/netbird/management/internals/shared/db" + "github.com/netbirdio/netbird/management/internals/shared/db/dbtest" +) + +func newTestRepository(t *testing.T) (accesslogs.Repository, *db.Conn) { + conn := dbtest.NewConn(t, &accesslogs.AccessLogEntry{}) + return NewRepository(conn), conn +} + +func newEntry(id, accountID, method string, age time.Duration) *accesslogs.AccessLogEntry { + return &accesslogs.AccessLogEntry{ + ID: id, + AccountID: accountID, + Method: method, + Host: "app.example.com", + Path: "/", + StatusCode: 200, + Timestamp: time.Now().Add(-age), + } +} + +func TestSqlRepository_ListByAccount(t *testing.T) { + repo, _ := newTestRepository(t) + ctx := context.Background() + for _, entry := range []*accesslogs.AccessLogEntry{ + newEntry("a1", "acc-a", "GET", 3*time.Hour), + newEntry("a2", "acc-a", "POST", 2*time.Hour), + newEntry("a3", "acc-a", "GET", time.Hour), + newEntry("b1", "acc-b", "GET", time.Hour), + } { + require.NoError(t, repo.Create(ctx, entry)) + } + + logs, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc-a", accesslogs.AccessLogFilter{Page: 1, PageSize: 2}) + require.NoError(t, err) + assert.EqualValues(t, 3, total) + require.Len(t, logs, 2) + assert.Equal(t, "a3", logs[0].ID) + assert.Equal(t, "a2", logs[1].ID) + + method := "GET" + logs, total, err = repo.ListByAccount(ctx, db.LockingStrengthNone, "acc-a", accesslogs.AccessLogFilter{Page: 1, PageSize: 10, Method: &method, SortOrder: "asc"}) + require.NoError(t, err) + assert.EqualValues(t, 2, total) + require.Len(t, logs, 2) + assert.Equal(t, "a1", logs[0].ID) + assert.Equal(t, "a3", logs[1].ID) +} + +func TestSqlRepository_DeleteOlderThan(t *testing.T) { + repo, _ := newTestRepository(t) + ctx := context.Background() + require.NoError(t, repo.Create(ctx, newEntry("old", "acc", "GET", 48*time.Hour))) + require.NoError(t, repo.Create(ctx, newEntry("new", "acc", "GET", time.Hour))) + + deleted, err := repo.DeleteOlderThan(ctx, time.Now().Add(-24*time.Hour)) + require.NoError(t, err) + assert.EqualValues(t, 1, deleted) + + logs, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10}) + require.NoError(t, err) + assert.EqualValues(t, 1, total) + require.Len(t, logs, 1) + assert.Equal(t, "new", logs[0].ID) +} + +func TestSqlRepository_CreateInsideTransactionRollsBack(t *testing.T) { + repo, conn := newTestRepository(t) + ctx := context.Background() + failure := errors.New("abort") + + err := conn.RunInTx(ctx, func(tx *db.Tx) error { + txRepo := repo.WithTx(tx) + require.NoError(t, txRepo.Create(ctx, newEntry("tx", "acc", "GET", 0))) + _, total, err := txRepo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10}) + require.NoError(t, err) + assert.EqualValues(t, 1, total) + return failure + }) + require.ErrorIs(t, err, failure) + + _, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10}) + require.NoError(t, err) + assert.Zero(t, total) +} + +func TestSqlRepository_ListByAccount_StatusFilter(t *testing.T) { + repo, _ := newTestRepository(t) + ctx := context.Background() + statusCodes := map[string]int{"l4": 0, "info": 101, "ok": 200, "notfound": 404} + for id, code := range statusCodes { + entry := newEntry(id, "acc", "GET", time.Hour) + entry.StatusCode = code + require.NoError(t, repo.Create(ctx, entry)) + } + foreign := newEntry("foreign", "other", "GET", time.Hour) + foreign.StatusCode = 500 + require.NoError(t, repo.Create(ctx, foreign)) + + listIDs := func(status string) []string { + logs, total, err := repo.ListByAccount(ctx, db.LockingStrengthNone, "acc", accesslogs.AccessLogFilter{Page: 1, PageSize: 10, Status: &status, SortBy: "status_code", SortOrder: "asc"}) + require.NoError(t, err) + require.EqualValues(t, len(logs), total) + ids := make([]string, 0, len(logs)) + for _, entry := range logs { + ids = append(ids, entry.ID) + } + return ids + } + + assert.Equal(t, []string{"info", "notfound"}, listIDs("failed")) + assert.Equal(t, []string{"ok"}, listIDs("success")) +} diff --git a/management/internals/modules/reverseproxy/accesslogs/repository.go b/management/internals/modules/reverseproxy/accesslogs/repository.go new file mode 100644 index 000000000..5945f454c --- /dev/null +++ b/management/internals/modules/reverseproxy/accesslogs/repository.go @@ -0,0 +1,18 @@ +package accesslogs + +import ( + "context" + "time" + + "github.com/netbirdio/netbird/management/internals/shared/db" +) + +//go:generate go tool mockgen -package accesslogs -destination=repository_mock.go -source=./repository.go -build_flags=-mod=mod + +// Repository persists reverse proxy access log entries. +type Repository interface { + WithTx(tx *db.Tx) Repository + Create(ctx context.Context, entry *AccessLogEntry) error + ListByAccount(ctx context.Context, lockStrength db.LockingStrength, accountID string, filter AccessLogFilter) ([]*AccessLogEntry, int64, error) + DeleteOlderThan(ctx context.Context, olderThan time.Time) (int64, error) +} diff --git a/management/internals/modules/reverseproxy/accesslogs/repository_mock.go b/management/internals/modules/reverseproxy/accesslogs/repository_mock.go new file mode 100644 index 000000000..7fbb05b7f --- /dev/null +++ b/management/internals/modules/reverseproxy/accesslogs/repository_mock.go @@ -0,0 +1,102 @@ +// Code generated by MockGen. DO NOT EDIT. +// Source: ./repository.go +// +// Generated by this command: +// +// mockgen -package accesslogs -destination=repository_mock.go -source=./repository.go -build_flags=-mod=mod +// + +// Package accesslogs is a generated GoMock package. +package accesslogs + +import ( + context "context" + reflect "reflect" + time "time" + + db "github.com/netbirdio/netbird/management/internals/shared/db" + gomock "go.uber.org/mock/gomock" +) + +// MockRepository is a mock of Repository interface. +type MockRepository struct { + ctrl *gomock.Controller + recorder *MockRepositoryMockRecorder + isgomock struct{} +} + +// MockRepositoryMockRecorder is the mock recorder for MockRepository. +type MockRepositoryMockRecorder struct { + mock *MockRepository +} + +// NewMockRepository creates a new mock instance. +func NewMockRepository(ctrl *gomock.Controller) *MockRepository { + mock := &MockRepository{ctrl: ctrl} + mock.recorder = &MockRepositoryMockRecorder{mock} + return mock +} + +// EXPECT returns an object that allows the caller to indicate expected use. +func (m *MockRepository) EXPECT() *MockRepositoryMockRecorder { + return m.recorder +} + +// Create mocks base method. +func (m *MockRepository) Create(ctx context.Context, entry *AccessLogEntry) error { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "Create", ctx, entry) + ret0, _ := ret[0].(error) + return ret0 +} + +// Create indicates an expected call of Create. +func (mr *MockRepositoryMockRecorder) Create(ctx, entry any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Create", reflect.TypeOf((*MockRepository)(nil).Create), ctx, entry) +} + +// DeleteOlderThan mocks base method. +func (m *MockRepository) DeleteOlderThan(ctx context.Context, olderThan time.Time) (int64, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "DeleteOlderThan", ctx, olderThan) + ret0, _ := ret[0].(int64) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// DeleteOlderThan indicates an expected call of DeleteOlderThan. +func (mr *MockRepositoryMockRecorder) DeleteOlderThan(ctx, olderThan any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteOlderThan", reflect.TypeOf((*MockRepository)(nil).DeleteOlderThan), ctx, olderThan) +} + +// ListByAccount mocks base method. +func (m *MockRepository) ListByAccount(ctx context.Context, lockStrength db.LockingStrength, accountID string, filter AccessLogFilter) ([]*AccessLogEntry, int64, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "ListByAccount", ctx, lockStrength, accountID, filter) + ret0, _ := ret[0].([]*AccessLogEntry) + ret1, _ := ret[1].(int64) + ret2, _ := ret[2].(error) + return ret0, ret1, ret2 +} + +// ListByAccount indicates an expected call of ListByAccount. +func (mr *MockRepositoryMockRecorder) ListByAccount(ctx, lockStrength, accountID, filter any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ListByAccount", reflect.TypeOf((*MockRepository)(nil).ListByAccount), ctx, lockStrength, accountID, filter) +} + +// WithTx mocks base method. +func (m *MockRepository) WithTx(tx *db.Tx) Repository { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "WithTx", tx) + ret0, _ := ret[0].(Repository) + return ret0 +} + +// WithTx indicates an expected call of WithTx. +func (mr *MockRepositoryMockRecorder) WithTx(tx any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "WithTx", reflect.TypeOf((*MockRepository)(nil).WithTx), tx) +} diff --git a/management/internals/server/boot.go b/management/internals/server/boot.go index dbad2f0c8..9eb8b7748 100644 --- a/management/internals/server/boot.go +++ b/management/internals/server/boot.go @@ -32,6 +32,7 @@ import ( networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db" networkmapdbfactory "github.com/netbirdio/netbird/management/internals/network_map_db/factory" nbconfig "github.com/netbirdio/netbird/management/internals/server/config" + "github.com/netbirdio/netbird/management/internals/shared/db" nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc" "github.com/netbirdio/netbird/management/server/activity" activitystore "github.com/netbirdio/netbird/management/server/activity/store" @@ -84,9 +85,20 @@ func (s *BaseServer) CacheStore() nbcache.Store { }) } +// DBConn opens the database connection shared by the store and the domain repositories. +func (s *BaseServer) DBConn() *db.Conn { + return Create(s, func() *db.Conn { + conn, err := store.OpenConn(context.Background(), s.Config.StoreConfig.Engine, s.Config.Datadir) + if err != nil { + log.Fatalf("failed to open database connection: %v", err) + } + return conn + }) +} + func (s *BaseServer) Store() store.Store { return Create(s, func() store.Store { - store, err := store.NewStore(context.Background(), s.Config.StoreConfig.Engine, s.Config.Datadir, s.Metrics(), false) + store, err := store.NewSqlStore(context.Background(), s.DBConn(), s.Metrics(), false) if err != nil { log.Fatalf("failed to create store: %v", err) } @@ -308,7 +320,7 @@ func (s *BaseServer) ProxyActivityManager() proxyactivity.Manager { func (s *BaseServer) AccessLogsManager() accesslogs.Manager { return Create(s, func() accesslogs.Manager { - accessLogManager := accesslogsmanager.NewManager(s.Store(), s.PermissionsManager(), s.GeoLocationManager()) + accessLogManager := accesslogsmanager.NewManager(accesslogsmanager.NewRepository(s.DBConn()), s.Store(), s.PermissionsManager(), s.GeoLocationManager()) accessLogManager.StartPeriodicCleanup( context.Background(), s.Config.ReverseProxy.AccessLogRetentionDays, diff --git a/management/internals/shared/db/conn.go b/management/internals/shared/db/conn.go new file mode 100644 index 000000000..8edaa1e4e --- /dev/null +++ b/management/internals/shared/db/conn.go @@ -0,0 +1,121 @@ +package db + +import ( + "context" + "fmt" + "os" + "runtime" + "strconv" + "time" + + "github.com/jackc/pgx/v5/pgxpool" + log "github.com/sirupsen/logrus" + "gorm.io/gorm" +) + +const ( + defaultTransactionTimeout = 5 * time.Minute + connMaxLifetime = time.Hour + connMaxIdleTime = 3 * time.Minute +) + +// TxMetrics receives the duration of every committed top-level transaction. +type TxMetrics interface { + CountTransactionDuration(duration time.Duration) +} + +// Conn is the database connection shared by all repositories: one gorm handle, +// the pgx pool of a Postgres deployment and the engine they talk to. +type Conn struct { + db *gorm.DB + pool *pgxpool.Pool + engine Engine + txTimeout time.Duration + metrics TxMetrics +} + +// NewConn takes ownership of an open gorm handle and pool once it returns +// without error, applying the connection limits and transaction timeout +// configured through the environment. +func NewConn(ctx context.Context, gormDB *gorm.DB, engine Engine, pool *pgxpool.Pool) (*Conn, error) { + sqlDB, err := gormDB.DB() + if err != nil { + return nil, err + } + + txTimeout := defaultTransactionTimeout + if v := os.Getenv("NB_STORE_TRANSACTION_TIMEOUT"); v != "" { + if parsed, err := time.ParseDuration(v); err == nil { + txTimeout = parsed + } + } + log.WithContext(ctx).Infof("Setting transaction timeout to %v", txTimeout) + + conns := runtime.NumCPU() + configuredConns, err := strconv.Atoi(os.Getenv("NB_SQL_MAX_OPEN_CONNS")) + connsConfigured := err == nil + if connsConfigured { + conns = configuredConns + } + if engine == SqliteStoreEngine { + if connsConfigured { + log.WithContext(ctx).Warnf("setting NB_SQL_MAX_OPEN_CONNS is not supported for sqlite, using default value 1") + } + conns = 1 + } + + sqlDB.SetMaxOpenConns(conns) + sqlDB.SetMaxIdleConns(conns) + sqlDB.SetConnMaxLifetime(connMaxLifetime) + sqlDB.SetConnMaxIdleTime(connMaxIdleTime) + + log.WithContext(ctx).Infof("Set max open db connections to %d, max idle to %d, max lifetime to %v, max idle time to %v", + conns, conns, connMaxLifetime, connMaxIdleTime) + + return &Conn{db: gormDB, pool: pool, engine: engine, txTimeout: txTimeout}, nil +} + +// DB returns the handle a query must run on: the transaction when tx is set, +// otherwise the shared connection. +func (c *Conn) DB(tx *Tx) *gorm.DB { + if tx != nil { + return tx.db + } + return c.db +} + +// Pool returns the pgx pool for read paths that bypass gorm. It is nil on +// engines other than Postgres and inside a transaction, where the pool would +// not see the uncommitted writes. +func (c *Conn) Pool(tx *Tx) *pgxpool.Pool { + if tx != nil { + return nil + } + return c.pool +} + +func (c *Conn) Engine() Engine { + return c.engine +} + +// SetTxMetrics registers the sink that receives transaction durations. +func (c *Conn) SetTxMetrics(metrics TxMetrics) { + c.metrics = metrics +} + +// AutoMigrate creates or updates the tables of the given models. +func (c *Conn) AutoMigrate(models ...any) error { + return c.db.AutoMigrate(models...) +} + +// Close releases the gorm connection and the pgx pool. +func (c *Conn) Close() error { + if c.pool != nil { + c.pool.Close() + } + sqlDB, err := c.db.DB() + if err != nil { + return fmt.Errorf("get db: %w", err) + } + return sqlDB.Close() +} diff --git a/management/internals/shared/db/conn_test.go b/management/internals/shared/db/conn_test.go new file mode 100644 index 000000000..4f1a234b2 --- /dev/null +++ b/management/internals/shared/db/conn_test.go @@ -0,0 +1,148 @@ +package db + +import ( + "context" + "errors" + "path/filepath" + "testing" + "time" + + "github.com/jackc/pgx/v5/pgxpool" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/driver/sqlite" + "gorm.io/gorm" +) + +type testRow struct { + ID uint `gorm:"primaryKey"` + Name string +} + +func openTestConn(t *testing.T) *Conn { + t.Helper() + conn, err := OpenSqliteFile(context.Background(), t.TempDir(), SqliteFileName) + require.NoError(t, err) + t.Cleanup(func() { require.NoError(t, conn.Close()) }) + require.NoError(t, conn.AutoMigrate(&testRow{})) + return conn +} + +func countRows(t *testing.T, conn *Conn) int64 { + t.Helper() + var count int64 + require.NoError(t, conn.DB(nil).Model(&testRow{}).Count(&count).Error) + return count +} + +func TestNewConn_ReadsTransactionTimeoutFromEnv(t *testing.T) { + t.Setenv("NB_STORE_TRANSACTION_TIMEOUT", "1s") + conn := openTestConn(t) + assert.Equal(t, time.Second, conn.txTimeout) + assert.Equal(t, SqliteStoreEngine, conn.Engine()) +} + +func TestRunInTx_CommitsOnSuccess(t *testing.T) { + conn := openTestConn(t) + + err := conn.RunInTx(context.Background(), func(tx *Tx) error { + return conn.DB(tx).Create(&testRow{Name: "a"}).Error + }) + require.NoError(t, err) + assert.EqualValues(t, 1, countRows(t, conn)) +} + +func TestRunInTx_RollsBackOnError(t *testing.T) { + conn := openTestConn(t) + failure := errors.New("boom") + + err := conn.RunInTx(context.Background(), func(tx *Tx) error { + require.NoError(t, conn.DB(tx).Create(&testRow{Name: "a"}).Error) + return failure + }) + require.ErrorIs(t, err, failure) + assert.EqualValues(t, 0, countRows(t, conn)) +} + +func TestRunInTx_RollsBackOnPanic(t *testing.T) { + conn := openTestConn(t) + + require.Panics(t, func() { + _ = conn.RunInTx(context.Background(), func(tx *Tx) error { + require.NoError(t, conn.DB(tx).Create(&testRow{Name: "a"}).Error) + panic("boom") + }) + }) + assert.EqualValues(t, 0, countRows(t, conn)) +} + +func TestRunInTx_FailsWhenTimeoutExceeded(t *testing.T) { + t.Setenv("NB_STORE_TRANSACTION_TIMEOUT", "50ms") + conn := openTestConn(t) + + err := conn.RunInTx(context.Background(), func(tx *Tx) error { + time.Sleep(100 * time.Millisecond) + return conn.DB(tx).Create(&testRow{Name: "a"}).Error + }) + require.ErrorIs(t, err, context.DeadlineExceeded) + assert.EqualValues(t, 0, countRows(t, conn)) +} + +func TestRunInTx_ReportsDurationToMetrics(t *testing.T) { + conn := openTestConn(t) + metrics := &recordingMetrics{} + conn.SetTxMetrics(metrics) + + require.NoError(t, conn.RunInTx(context.Background(), func(*Tx) error { return nil })) + assert.Equal(t, 1, metrics.calls) +} + +func TestConn_DBSelectsTransactionHandle(t *testing.T) { + conn := openTestConn(t) + assert.Same(t, conn.db, conn.DB(nil)) + + err := conn.RunInTx(context.Background(), func(tx *Tx) error { + assert.Same(t, tx.db, conn.DB(tx)) + assert.NotSame(t, conn.db, conn.DB(tx)) + return nil + }) + require.NoError(t, err) +} + +func TestConn_PoolIsUnavailableInsideTransaction(t *testing.T) { + conn := openTestConn(t) + conn.pool = &pgxpool.Pool{} + defer func() { conn.pool = nil }() + + assert.Same(t, conn.pool, conn.Pool(nil)) + err := conn.RunInTx(context.Background(), func(tx *Tx) error { + assert.Nil(t, conn.Pool(tx)) + return nil + }) + require.NoError(t, err) +} + +type recordingMetrics struct { + calls int +} + +func (m *recordingMetrics) CountTransactionDuration(time.Duration) { + m.calls++ +} + +func TestNewConn_MaxOpenConnsFromEnv(t *testing.T) { + t.Setenv("NB_SQL_MAX_OPEN_CONNS", "7") + + gormDB, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "store.db")), GormConfig()) + require.NoError(t, err) + conn, err := NewConn(context.Background(), gormDB, PostgresStoreEngine, nil) + require.NoError(t, err) + t.Cleanup(func() { _ = conn.Close() }) + sqlDB, err := conn.DB(nil).DB() + require.NoError(t, err) + assert.Equal(t, 7, sqlDB.Stats().MaxOpenConnections) + + sqliteDB, err := openTestConn(t).DB(nil).DB() + require.NoError(t, err) + assert.Equal(t, 1, sqliteDB.Stats().MaxOpenConnections) +} diff --git a/management/internals/shared/db/dbtest/dbtest.go b/management/internals/shared/db/dbtest/dbtest.go new file mode 100644 index 000000000..0f0fe0356 --- /dev/null +++ b/management/internals/shared/db/dbtest/dbtest.go @@ -0,0 +1,23 @@ +package dbtest + +import ( + "context" + "testing" + + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/shared/db" +) + +// NewConn opens a fresh SQLite database in a temporary directory, migrates the +// given models and closes the connection when the test ends. It ignores +// NB_STORE_ENGINE_SQLITE_FILE, so a developer's configured database is never +// touched, and is safe to call from parallel tests. +func NewConn(t testing.TB, models ...any) *db.Conn { + t.Helper() + conn, err := db.OpenSqliteFile(context.Background(), t.TempDir(), db.SqliteFileName) + require.NoError(t, err) + t.Cleanup(func() { _ = conn.Close() }) + require.NoError(t, conn.AutoMigrate(models...)) + return conn +} diff --git a/management/internals/shared/db/dbtest/dbtest_test.go b/management/internals/shared/db/dbtest/dbtest_test.go new file mode 100644 index 000000000..3f8566876 --- /dev/null +++ b/management/internals/shared/db/dbtest/dbtest_test.go @@ -0,0 +1,31 @@ +package dbtest + +import ( + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + + "github.com/netbirdio/netbird/management/internals/shared/db" +) + +func TestNewConn_IgnoresSqliteFileOverride(t *testing.T) { + override := filepath.Join(t.TempDir(), "configured.db") + t.Setenv("NB_STORE_ENGINE_SQLITE_FILE", override) + + conn := NewConn(t) + + assert.Equal(t, db.SqliteStoreEngine, conn.Engine()) + _, err := os.Stat(override) + require.ErrorIs(t, err, os.ErrNotExist) +} + +func TestNewConn_Parallel(t *testing.T) { + t.Parallel() + + conn := NewConn(t) + + assert.Equal(t, db.SqliteStoreEngine, conn.Engine()) +} diff --git a/management/internals/shared/db/engine.go b/management/internals/shared/db/engine.go new file mode 100644 index 000000000..02a3f5570 --- /dev/null +++ b/management/internals/shared/db/engine.go @@ -0,0 +1,10 @@ +package db + +// Engine identifies the SQL engine behind a Conn. +type Engine string + +const ( + SqliteStoreEngine Engine = "sqlite" + PostgresStoreEngine Engine = "postgres" + MysqlStoreEngine Engine = "mysql" +) diff --git a/management/internals/shared/db/lock.go b/management/internals/shared/db/lock.go new file mode 100644 index 000000000..bf6fc9b5a --- /dev/null +++ b/management/internals/shared/db/lock.go @@ -0,0 +1,12 @@ +package db + +// LockingStrength is the row lock a query holds until its transaction ends. +type LockingStrength string + +const ( + LockingStrengthUpdate LockingStrength = "UPDATE" + LockingStrengthShare LockingStrength = "SHARE" + LockingStrengthNoKeyUpdate LockingStrength = "NO KEY UPDATE" + LockingStrengthKeyShare LockingStrength = "KEY SHARE" + LockingStrengthNone LockingStrength = "NONE" +) diff --git a/management/internals/shared/db/open.go b/management/internals/shared/db/open.go new file mode 100644 index 000000000..020aff422 --- /dev/null +++ b/management/internals/shared/db/open.go @@ -0,0 +1,160 @@ +package db + +import ( + "context" + "fmt" + "net/url" + "os" + "path/filepath" + "runtime" + "strings" + "time" + + "github.com/jackc/pgx/v5/pgxpool" + "gorm.io/driver/mysql" + "gorm.io/driver/postgres" + "gorm.io/driver/sqlite" + "gorm.io/gorm" + "gorm.io/gorm/logger" +) + +// SqliteFileName is the default SQLite database file inside the data directory. +const SqliteFileName = "store.db" + +// PoolConfig sizes the pgx pool a Postgres deployment uses for the read paths +// that bypass gorm. +type PoolConfig struct { + MaxConns int32 + MinConns int32 + MaxConnLifetime time.Duration + HealthCheckPeriod time.Duration +} + +var DefaultPoolConfig = PoolConfig{ + MaxConns: 30, + MinConns: 1, + MaxConnLifetime: 60 * time.Minute, + HealthCheckPeriod: time.Minute, +} + +// GormConfig is the configuration every engine is opened with. +func GormConfig() *gorm.Config { + return &gorm.Config{ + Logger: logger.Default.LogMode(logger.Silent), + CreateBatchSize: 400, + } +} + +// OpenSqlite opens the SQLite database in dataDir, or the file named by +// NB_STORE_ENGINE_SQLITE_FILE. +func OpenSqlite(ctx context.Context, dataDir string) (*Conn, error) { + storeFile := SqliteFileName + if envFile, ok := os.LookupEnv("NB_STORE_ENGINE_SQLITE_FILE"); ok && envFile != "" { + storeFile = envFile + } + return OpenSqliteFile(ctx, dataDir, storeFile) +} + +// OpenSqliteFile opens the SQLite database storeFile, resolved against dataDir +// when relative. storeFile may carry SQLite URI query parameters. +func OpenSqliteFile(ctx context.Context, dataDir, storeFile string) (*Conn, error) { + // Separate file path from any SQLite URI query parameters (e.g., "store.db?mode=rwc") + filePath, query, hasQuery := strings.Cut(storeFile, "?") + + connStr := filePath + if !filepath.IsAbs(filePath) { + connStr = filepath.Join(dataDir, filePath) + } + + // Compose query parameters. User-provided ?_busy_timeout (or its mattn alias + // ?_timeout) overrides our default; otherwise inject 30s so SQLite waits at + // most that long on a lock instead of blocking the only Go-side connection. + // mattn/go-sqlite3 applies PRAGMA from the DSN on every fresh connection, so + // the value survives ConnMaxIdleTime/ConnMaxLifetime recycling. cache=shared + // stays the default on non-Windows for the same reason as before. + parsed, _ := url.ParseQuery(query) + var defaults []string + if parsed.Get("_busy_timeout") == "" && parsed.Get("_timeout") == "" { + defaults = append(defaults, "_busy_timeout=30000") + } + if !hasQuery && runtime.GOOS != "windows" { + // To avoid `The process cannot access the file because it is being used by another process` on Windows + defaults = append(defaults, "cache=shared") + } + parts := defaults + if hasQuery { + parts = append(parts, query) + } + if len(parts) > 0 { + connStr += "?" + strings.Join(parts, "&") + } + + gormDB, err := gorm.Open(sqlite.Open(connStr), GormConfig()) + if err != nil { + return nil, err + } + return NewConn(ctx, gormDB, SqliteStoreEngine, nil) +} + +// OpenPostgres opens a Postgres database through gorm and a pgx pool sized by pool. +func OpenPostgres(ctx context.Context, dsn string, pool PoolConfig) (*Conn, error) { + gormDB, err := gorm.Open(postgres.Open(dsn), GormConfig()) + if err != nil { + return nil, err + } + pgxPool, err := newPgxPool(ctx, dsn, pool) + if err != nil { + closeGorm(gormDB) + return nil, err + } + return NewConn(ctx, gormDB, PostgresStoreEngine, pgxPool) +} + +// MysqlDSN adds the connection parameters every MySQL handle needs, keeping +// the options already present in dsn. +func MysqlDSN(dsn string) string { + separator := "?" + if strings.Contains(dsn, "?") { + separator = "&" + } + return dsn + separator + "charset=utf8&parseTime=True&loc=Local" +} + +// OpenMysql opens a MySQL database through gorm. +func OpenMysql(ctx context.Context, dsn string) (*Conn, error) { + gormDB, err := gorm.Open(mysql.Open(MysqlDSN(dsn)), GormConfig()) + if err != nil { + return nil, err + } + return NewConn(ctx, gormDB, MysqlStoreEngine, nil) +} + +func newPgxPool(ctx context.Context, dsn string, cfg PoolConfig) (*pgxpool.Pool, error) { + config, err := pgxpool.ParseConfig(dsn) + if err != nil { + return nil, fmt.Errorf("unable to parse database config: %w", err) + } + + config.MaxConns = cfg.MaxConns + config.MinConns = cfg.MinConns + config.MaxConnLifetime = cfg.MaxConnLifetime + config.HealthCheckPeriod = cfg.HealthCheckPeriod + + pool, err := pgxpool.NewWithConfig(ctx, config) + if err != nil { + return nil, fmt.Errorf("unable to create connection pool: %w", err) + } + + if err := pool.Ping(ctx); err != nil { + pool.Close() + return nil, fmt.Errorf("unable to ping database: %w", err) + } + + return pool, nil +} + +func closeGorm(gormDB *gorm.DB) { + if sqlDB, err := gormDB.DB(); err == nil { + _ = sqlDB.Close() + } +} diff --git a/management/internals/shared/db/open_test.go b/management/internals/shared/db/open_test.go new file mode 100644 index 000000000..e6b506b78 --- /dev/null +++ b/management/internals/shared/db/open_test.go @@ -0,0 +1,12 @@ +package db + +import ( + "testing" + + "github.com/stretchr/testify/assert" +) + +func TestMysqlDSN(t *testing.T) { + assert.Equal(t, "user:pw@tcp(host:3306)/db?charset=utf8&parseTime=True&loc=Local", MysqlDSN("user:pw@tcp(host:3306)/db")) + assert.Equal(t, "user:pw@tcp(host:3306)/db?tls=true&charset=utf8&parseTime=True&loc=Local", MysqlDSN("user:pw@tcp(host:3306)/db?tls=true")) +} diff --git a/management/internals/shared/db/transaction.go b/management/internals/shared/db/transaction.go new file mode 100644 index 000000000..9699aaee6 --- /dev/null +++ b/management/internals/shared/db/transaction.go @@ -0,0 +1,105 @@ +package db + +import ( + "context" + "errors" + "fmt" + "runtime/debug" + "time" + + log "github.com/sirupsen/logrus" + "gorm.io/gorm" +) + +// Tx is an open transaction handed to repository calls; nil means autocommit. +type Tx struct { + db *gorm.DB +} + +// RunInTx runs fn in one transaction that commits when fn returns nil and rolls +// back otherwise, bounded by the configured transaction timeout. +func (c *Conn) RunInTx(ctx context.Context, fn func(tx *Tx) error) error { + timeoutCtx, cancel := context.WithTimeout(ctx, c.txTimeout) + defer cancel() + + startTime := time.Now() + tx := c.db.WithContext(timeoutCtx).Begin() + if tx.Error != nil { + return tx.Error + } + defer func() { + if r := recover(); r != nil { + tx.Rollback() + panic(r) + } + }() + + if err := c.applyStatementTimeouts(tx); err != nil { + tx.Rollback() + return err + } + + err := c.withForeignKeyChecksDisabled(tx, func() error { + return fn(&Tx{db: tx}) + }) + if err != nil { + tx.Rollback() + c.logIfTimedOut(ctx, timeoutCtx, err, "transaction", startTime) + return err + } + + if err := tx.Commit().Error; err != nil { + c.logIfTimedOut(ctx, timeoutCtx, err, "transaction commit", startTime) + return err + } + + log.WithContext(ctx).Tracef("transaction took %v", time.Since(startTime)) + if c.metrics != nil { + c.metrics.CountTransactionDuration(time.Since(startTime)) + } + return nil +} + +func (c *Conn) applyStatementTimeouts(tx *gorm.DB) error { + if c.engine != PostgresStoreEngine { + return nil + } + if err := tx.Exec("SET LOCAL statement_timeout = '1min'").Error; err != nil { + return fmt.Errorf("failed to set statement timeout: %w", err) + } + if err := tx.Exec("SET LOCAL lock_timeout = '1min'").Error; err != nil { + return fmt.Errorf("failed to set lock timeout: %w", err) + } + return nil +} + +// withForeignKeyChecksDisabled runs fn with MySQL's FK checks off, which avoids +// deadlocks on MySQL and Aurora without needing SUPER privilege. The setting is +// session-scoped and survives a rollback, so it is turned back on whenever fn +// returns or panics; otherwise the pooled connection would keep it disabled. +func (c *Conn) withForeignKeyChecksDisabled(tx *gorm.DB, fn func() error) (err error) { + if c.engine != MysqlStoreEngine { + return fn() + } + if err := tx.Exec("SET FOREIGN_KEY_CHECKS = 0").Error; err != nil { + return fmt.Errorf("failed to disable FK checks: %w", err) + } + defer func() { + restoreErr := tx.Exec("SET FOREIGN_KEY_CHECKS = 1").Error + if restoreErr == nil { + return + } + if err == nil { + err = fmt.Errorf("failed to re-enable FK checks: %w", restoreErr) + return + } + log.WithContext(tx.Statement.Context).Warnf("failed to re-enable FK checks after failed transaction: %v", restoreErr) + }() + return fn() +} + +func (c *Conn) logIfTimedOut(ctx, timeoutCtx context.Context, err error, phase string, startTime time.Time) { + if errors.Is(err, context.DeadlineExceeded) || errors.Is(timeoutCtx.Err(), context.DeadlineExceeded) { + log.WithContext(ctx).Warnf("%s exceeded %s timeout after %v, stack: %s", phase, c.txTimeout, time.Since(startTime), debug.Stack()) + } +} diff --git a/management/server/http/testing/testing_tools/channel/channel.go b/management/server/http/testing/testing_tools/channel/channel.go index ab992b1ca..6ea26e634 100644 --- a/management/server/http/testing/testing_tools/channel/channel.go +++ b/management/server/http/testing/testing_tools/channel/channel.go @@ -46,14 +46,14 @@ import ( "github.com/netbirdio/netbird/management/server/networks/routers" "github.com/netbirdio/netbird/management/server/permissions" "github.com/netbirdio/netbird/management/server/settings" - "github.com/netbirdio/netbird/management/server/store" + nbstore "github.com/netbirdio/netbird/management/server/store" "github.com/netbirdio/netbird/management/server/telemetry" "github.com/netbirdio/netbird/management/server/users" "github.com/netbirdio/netbird/shared/auth" ) func BuildApiBlackBoxWithDBState(t testing_tools.TB, sqlFile string, expectedPeerUpdate *network_map.UpdateMessage, validateUpdate bool) (http.Handler, account.Manager, chan struct{}) { - store, cleanup, err := store.NewTestStoreFromSQL(context.Background(), sqlFile, t.TempDir()) + store, cleanup, err := nbstore.NewTestStoreFromSQL(context.Background(), sqlFile, t.TempDir()) if err != nil { t.Fatalf("Failed to create test store: %v", err) } @@ -108,7 +108,7 @@ func BuildApiBlackBoxWithDBState(t testing_tools.TB, sqlFile string, expectedPee t.Fatalf("Failed to create manager: %v", err) } - accessLogsManager := accesslogsmanager.NewManager(store, permissionsManager, nil) + accessLogsManager := accesslogsmanager.NewManager(accesslogsmanager.NewRepository(store.(*nbstore.SqlStore).Conn()), store, permissionsManager, nil) proxyTokenStore := nbgrpc.NewOneTimeTokenStore(ctx, cacheStore) pkceverifierStore := nbgrpc.NewPKCEVerifierStore(ctx, cacheStore) noopMeter := noop.NewMeterProvider().Meter("") @@ -204,7 +204,7 @@ func PeerShouldNotReceiveAnyUpdate(t testing_tools.TB, updateMessage <-chan *net // BuildApiBlackBoxWithDBStateAndPeerChannel creates the API handler and returns // the peer update channel directly so tests can verify updates inline. func BuildApiBlackBoxWithDBStateAndPeerChannel(t testing_tools.TB, sqlFile string) (http.Handler, account.Manager, <-chan *network_map.UpdateMessage) { - store, cleanup, err := store.NewTestStoreFromSQL(context.Background(), sqlFile, t.TempDir()) + store, cleanup, err := nbstore.NewTestStoreFromSQL(context.Background(), sqlFile, t.TempDir()) if err != nil { t.Fatalf("Failed to create test store: %v", err) } @@ -248,7 +248,7 @@ func BuildApiBlackBoxWithDBStateAndPeerChannel(t testing_tools.TB, sqlFile strin t.Fatalf("Failed to create manager: %v", err) } - accessLogsManager := accesslogsmanager.NewManager(store, permissionsManager, nil) + accessLogsManager := accesslogsmanager.NewManager(accesslogsmanager.NewRepository(store.(*nbstore.SqlStore).Conn()), store, permissionsManager, nil) proxyTokenStore := nbgrpc.NewOneTimeTokenStore(ctx, cacheStore) pkceverifierStore := nbgrpc.NewPKCEVerifierStore(ctx, cacheStore) noopMeter := noop.NewMeterProvider().Meter("") diff --git a/management/server/store/sql_store.go b/management/server/store/sql_store.go index 425c1a0cb..b705ba0a2 100644 --- a/management/server/store/sql_store.go +++ b/management/server/store/sql_store.go @@ -2,25 +2,13 @@ package store import ( "context" - "errors" "fmt" - "net/url" - "os" - "path/filepath" - "runtime" - "runtime/debug" - "strconv" - "strings" "sync" "time" "github.com/jackc/pgx/v5/pgxpool" log "github.com/sirupsen/logrus" - "gorm.io/driver/mysql" - "gorm.io/driver/postgres" - "gorm.io/driver/sqlite" "gorm.io/gorm" - "gorm.io/gorm/logger" nbdns "github.com/netbirdio/netbird/dns" agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" @@ -30,6 +18,7 @@ import ( rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" "github.com/netbirdio/netbird/management/internals/modules/zones" "github.com/netbirdio/netbird/management/internals/modules/zones/records" + "github.com/netbirdio/netbird/management/internals/shared/db" resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types" routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" networkTypes "github.com/netbirdio/netbird/management/server/networks/types" @@ -42,7 +31,6 @@ import ( ) const ( - storeSqliteFileName = "store.db" idQueryCondition = "id = ?" keyQueryCondition = "key = ?" mysqlKeyQueryCondition = "`key` = ?" @@ -52,71 +40,44 @@ const ( accountAndIDsQueryCondition = "account_id = ? AND id IN ?" accountIDCondition = "account_id = ?" peerNotFoundFMT = "peer %s not found" - - pgMaxConnections = 30 - pgMinConnections = 1 - pgMaxConnLifetime = 60 * time.Minute - pgHealthCheckPeriod = 1 * time.Minute ) +var testPoolConfig = db.PoolConfig{ + MaxConns: 5, + MinConns: 1, + MaxConnLifetime: 30 * time.Second, + HealthCheckPeriod: 10 * time.Second, +} + // SqlStore represents an account storage backed by a Sql DB persisted to disk type SqlStore struct { - db *gorm.DB - globalAccountLock sync.Mutex - metrics telemetry.AppMetrics - installationPK int - storeEngine types.Engine - pool *pgxpool.Pool - fieldEncrypt *crypt.FieldEncrypt - transactionTimeout time.Duration + conn *db.Conn + db *gorm.DB + tx *db.Tx + globalAccountLock sync.Mutex + metrics telemetry.AppMetrics + installationPK int + fieldEncrypt *crypt.FieldEncrypt } type migrationFunc func(*gorm.DB) error -// NewSqlStore creates a new SqlStore instance. -func NewSqlStore(ctx context.Context, db *gorm.DB, storeEngine types.Engine, metrics telemetry.AppMetrics, skipMigration bool) (*SqlStore, error) { - sql, err := db.DB() - if err != nil { - return nil, err +// NewSqlStore creates a new SqlStore instance on top of an open connection. +func NewSqlStore(ctx context.Context, conn *db.Conn, metrics telemetry.AppMetrics, skipMigration bool) (*SqlStore, error) { + if metrics != nil { + conn.SetTxMetrics(metrics.StoreMetrics()) } - - conns, err := strconv.Atoi(os.Getenv("NB_SQL_MAX_OPEN_CONNS")) - if err != nil { - conns = runtime.NumCPU() - } - - transactionTimeout := 5 * time.Minute - if v := os.Getenv("NB_STORE_TRANSACTION_TIMEOUT"); v != "" { - if parsed, err := time.ParseDuration(v); err == nil { - transactionTimeout = parsed - } - } - log.WithContext(ctx).Infof("Setting transaction timeout to %v", transactionTimeout) - - if storeEngine == types.SqliteStoreEngine { - if err == nil { - log.WithContext(ctx).Warnf("setting NB_SQL_MAX_OPEN_CONNS is not supported for sqlite, using default value 1") - } - conns = 1 - } - - sql.SetMaxOpenConns(conns) - sql.SetMaxIdleConns(conns) - sql.SetConnMaxLifetime(time.Hour) - sql.SetConnMaxIdleTime(3 * time.Minute) - - log.WithContext(ctx).Infof("Set max open db connections to %d, max idle to %d, max lifetime to %v, max idle time to %v", - conns, conns, time.Hour, 3*time.Minute) + store := &SqlStore{conn: conn, db: conn.DB(nil), metrics: metrics, installationPK: 1} if skipMigration { log.WithContext(ctx).Infof("skipping migration") - return &SqlStore{db: db, storeEngine: storeEngine, metrics: metrics, installationPK: 1, transactionTimeout: transactionTimeout}, nil + return store, nil } - if err := migratePreAuto(ctx, db); err != nil { + if err := migratePreAuto(ctx, store.db); err != nil { return nil, fmt.Errorf("migratePreAuto: %w", err) } - err = db.AutoMigrate( + err := conn.AutoMigrate( &types.SetupKey{}, &nbpeer.Peer{}, &types.User{}, &types.PersonalAccessToken{}, &types.ProxyAccessToken{}, &types.Group{}, &types.GroupPeer{}, &types.Account{}, &types.Policy{}, &types.PolicyRule{}, &route.Route{}, &nbdns.NameServerGroup{}, @@ -132,15 +93,34 @@ func NewSqlStore(ctx context.Context, db *gorm.DB, storeEngine types.Engine, met if err != nil { return nil, fmt.Errorf("auto migratePreAuto: %w", err) } - if err := migratePostAuto(ctx, db); err != nil { + if err := migratePostAuto(ctx, store.db); err != nil { return nil, fmt.Errorf("migratePostAuto: %w", err) } - return &SqlStore{db: db, storeEngine: storeEngine, metrics: metrics, installationPK: 1, transactionTimeout: transactionTimeout}, nil + return store, nil +} + +// newStore runs the migrations on conn and releases it when they fail. +func newStore(ctx context.Context, conn *db.Conn, metrics telemetry.AppMetrics, skipMigration bool) (*SqlStore, error) { + store, err := NewSqlStore(ctx, conn, metrics, skipMigration) + if err != nil { + _ = conn.Close() + return nil, err + } + return store, nil +} + +// Conn returns the shared connection so domain repositories can run alongside this store. +func (s *SqlStore) Conn() *db.Conn { + return s.conn +} + +func (s *SqlStore) pgxPool() *pgxpool.Pool { + return s.conn.Pool(s.tx) } func GetKeyQueryCondition(s *SqlStore) string { - if s.storeEngine == types.MysqlStoreEngine { + if s.conn.Engine() == db.MysqlStoreEngine { return mysqlKeyQueryCondition } return keyQueryCondition @@ -168,145 +148,39 @@ func (s *SqlStore) AcquireGlobalLock(ctx context.Context) (unlock func()) { // Close closes the underlying DB connection func (s *SqlStore) Close(_ context.Context) error { - sql, err := s.db.DB() - if err != nil { - return fmt.Errorf("get db: %w", err) - } - return sql.Close() + return s.conn.Close() } // GetStoreEngine returns underlying store engine func (s *SqlStore) GetStoreEngine() types.Engine { - return s.storeEngine + return s.conn.Engine() } // NewSqliteStore creates a new SQLite store. func NewSqliteStore(ctx context.Context, dataDir string, metrics telemetry.AppMetrics, skipMigration bool) (*SqlStore, error) { - storeFile := storeSqliteFileName - if envFile, ok := os.LookupEnv("NB_STORE_ENGINE_SQLITE_FILE"); ok && envFile != "" { - storeFile = envFile - } - - // Separate file path from any SQLite URI query parameters (e.g., "store.db?mode=rwc") - filePath, query, hasQuery := strings.Cut(storeFile, "?") - - connStr := filePath - if !filepath.IsAbs(filePath) { - connStr = filepath.Join(dataDir, filePath) - } - - // Compose query parameters. User-provided ?_busy_timeout (or its mattn alias - // ?_timeout) overrides our default; otherwise inject 30s so SQLite waits at - // most that long on a lock instead of blocking the only Go-side connection. - // mattn/go-sqlite3 applies PRAGMA from the DSN on every fresh connection, so - // the value survives ConnMaxIdleTime/ConnMaxLifetime recycling. cache=shared - // stays the default on non-Windows for the same reason as before. - parsed, _ := url.ParseQuery(query) - var defaults []string - if parsed.Get("_busy_timeout") == "" && parsed.Get("_timeout") == "" { - defaults = append(defaults, "_busy_timeout=30000") - } - if !hasQuery && runtime.GOOS != "windows" { - // To avoid `The process cannot access the file because it is being used by another process` on Windows - defaults = append(defaults, "cache=shared") - } - parts := defaults - if hasQuery { - parts = append(parts, query) - } - if len(parts) > 0 { - connStr += "?" + strings.Join(parts, "&") - } - - db, err := gorm.Open(sqlite.Open(connStr), getGormConfig()) + conn, err := db.OpenSqlite(ctx, dataDir) if err != nil { return nil, err } - - return NewSqlStore(ctx, db, types.SqliteStoreEngine, metrics, skipMigration) + return newStore(ctx, conn, metrics, skipMigration) } // NewPostgresqlStore creates a new Postgres store. func NewPostgresqlStore(ctx context.Context, dsn string, metrics telemetry.AppMetrics, skipMigration bool) (*SqlStore, error) { - db, err := gorm.Open(postgres.Open(dsn), getGormConfig()) + conn, err := db.OpenPostgres(ctx, dsn, db.DefaultPoolConfig) if err != nil { return nil, err } - pool, err := connectToPgDb(context.Background(), dsn) - if err != nil { - return nil, err - } - store, err := NewSqlStore(ctx, db, types.PostgresStoreEngine, metrics, skipMigration) - if err != nil { - pool.Close() - return nil, err - } - store.pool = pool - return store, nil -} - -func connectToPgDb(ctx context.Context, dsn string) (*pgxpool.Pool, error) { - config, err := pgxpool.ParseConfig(dsn) - if err != nil { - return nil, fmt.Errorf("unable to parse database config: %w", err) - } - - config.MaxConns = pgMaxConnections - config.MinConns = pgMinConnections - config.MaxConnLifetime = pgMaxConnLifetime - config.HealthCheckPeriod = pgHealthCheckPeriod - - pool, err := pgxpool.NewWithConfig(ctx, config) - if err != nil { - return nil, fmt.Errorf("unable to create connection pool: %w", err) - } - - if err := pool.Ping(ctx); err != nil { - pool.Close() - return nil, fmt.Errorf("unable to ping database: %w", err) - } - - return pool, nil + return newStore(ctx, conn, metrics, skipMigration) } // NewMysqlStore creates a new MySQL store. func NewMysqlStore(ctx context.Context, dsn string, metrics telemetry.AppMetrics, skipMigration bool) (*SqlStore, error) { - db, err := gorm.Open(mysql.Open(dsn+"?charset=utf8&parseTime=True&loc=Local"), getGormConfig()) + conn, err := db.OpenMysql(ctx, dsn) if err != nil { return nil, err } - - store, err := NewSqlStore(ctx, db, types.MysqlStoreEngine, metrics, skipMigration) - if err != nil { - closeGormDB(db) - return nil, err - } - return store, nil -} - -func getGormConfig() *gorm.Config { - return &gorm.Config{ - Logger: logger.Default.LogMode(logger.Silent), - CreateBatchSize: 400, - } -} - -// newPostgresStore initializes a new Postgres store. -func newPostgresStore(ctx context.Context, metrics telemetry.AppMetrics, skipMigration bool) (Store, error) { - dsn, ok := lookupDSNEnv(PostgresDsnEnv, PostgresDsnEnvLegacy) - if !ok { - return nil, fmt.Errorf("%s is not set", PostgresDsnEnv) - } - return NewPostgresqlStore(ctx, dsn, metrics, skipMigration) -} - -// newMysqlStore initializes a new MySQL store. -func newMysqlStore(ctx context.Context, metrics telemetry.AppMetrics, skipMigration bool) (Store, error) { - dsn, ok := lookupDSNEnv(mysqlDsnEnv, mysqlDsnEnvLegacy) - if !ok { - return nil, fmt.Errorf("%s is not set", mysqlDsnEnv) - } - return NewMysqlStore(ctx, dsn, metrics, skipMigration) + return newStore(ctx, conn, metrics, skipMigration) } // NewSqliteStoreFromFileStore restores a store from FileStore and stores SQLite DB in the file located in datadir. @@ -350,7 +224,7 @@ func newPostgresqlStoreFromSqlStore(ctx context.Context, sqliteStore *SqlStore, } if err := seedFromSqliteStore(ctx, store, sqliteStore); err != nil { - closeStore(ctx, store) + _ = store.Close(ctx) return nil, err } @@ -359,49 +233,11 @@ func newPostgresqlStoreFromSqlStore(ctx context.Context, sqliteStore *SqlStore, // used for tests only func NewPostgresqlStoreForTests(ctx context.Context, dsn string, metrics telemetry.AppMetrics, skipMigration bool) (*SqlStore, error) { - db, err := gorm.Open(postgres.Open(dsn), getGormConfig()) + conn, err := db.OpenPostgres(ctx, dsn, testPoolConfig) if err != nil { return nil, err } - pool, err := connectToPgDbForTests(context.Background(), dsn) - if err != nil { - closeGormDB(db) - return nil, err - } - store, err := NewSqlStore(ctx, db, types.PostgresStoreEngine, metrics, skipMigration) - if err != nil { - // Release the sessions, or the caller cannot drop the database. - pool.Close() - closeGormDB(db) - return nil, err - } - store.pool = pool - return store, nil -} - -// used for tests only -func connectToPgDbForTests(ctx context.Context, dsn string) (*pgxpool.Pool, error) { - config, err := pgxpool.ParseConfig(dsn) - if err != nil { - return nil, fmt.Errorf("unable to parse database config: %w", err) - } - - config.MaxConns = 5 - config.MinConns = 1 - config.MaxConnLifetime = 30 * time.Second - config.HealthCheckPeriod = 10 * time.Second - - pool, err := pgxpool.NewWithConfig(ctx, config) - if err != nil { - return nil, fmt.Errorf("unable to create connection pool: %w", err) - } - - if err := pool.Ping(ctx); err != nil { - pool.Close() - return nil, fmt.Errorf("unable to ping database: %w", err) - } - - return pool, nil + return newStore(ctx, conn, metrics, skipMigration) } // NewMysqlStoreFromSqlStore restores a store from SqlStore and stores MySQL DB. @@ -423,15 +259,6 @@ func seedFromSqliteStore(ctx context.Context, store, sqliteStore *SqlStore) erro return nil } -// closeStore releases a store that is not handed to the caller, so a failed -// seed does not leak its connection and pool. -func closeStore(ctx context.Context, store *SqlStore) { - store.Close(ctx) - if store.pool != nil { - store.pool.Close() - } -} - func newMysqlStoreFromSqlStore(ctx context.Context, sqliteStore *SqlStore, dsn string, metrics telemetry.AppMetrics, skipMigration bool) (*SqlStore, error) { store, err := NewMysqlStore(ctx, dsn, metrics, skipMigration) if err != nil { @@ -439,113 +266,41 @@ func newMysqlStoreFromSqlStore(ctx context.Context, sqliteStore *SqlStore, dsn s } if err := seedFromSqliteStore(ctx, store, sqliteStore); err != nil { - closeStore(ctx, store) + _ = store.Close(ctx) return nil, err } return store, nil } +// ExecuteInTransaction runs operation in a transaction. A store that is already +// bound to one joins it instead of opening a second, independent transaction. func (s *SqlStore) ExecuteInTransaction(ctx context.Context, operation func(store Store) error) error { - timeoutCtx, cancel := context.WithTimeout(ctx, s.transactionTimeout) - defer cancel() - - startTime := time.Now() - tx := s.db.WithContext(timeoutCtx).Begin() - if tx.Error != nil { - return tx.Error + if s.tx != nil { + return operation(s) } - defer func() { - if r := recover(); r != nil { - tx.Rollback() - panic(r) - } - }() - - if s.storeEngine == types.PostgresStoreEngine { - if err := tx.Exec("SET LOCAL statement_timeout = '1min'").Error; err != nil { - tx.Rollback() - return fmt.Errorf("failed to set statement timeout: %w", err) - } - if err := tx.Exec("SET LOCAL lock_timeout = '1min'").Error; err != nil { - tx.Rollback() - return fmt.Errorf("failed to set lock timeout: %w", err) - } - } - - // For MySQL, disable FK checks within this transaction to avoid deadlocks - // This is session-scoped and doesn't require SUPER privileges - if s.storeEngine == types.MysqlStoreEngine { - if err := tx.Exec("SET FOREIGN_KEY_CHECKS = 0").Error; err != nil { - tx.Rollback() - return fmt.Errorf("failed to disable FK checks: %w", err) - } - } - - repo := s.withTx(tx) - err := operation(repo) - if err != nil { - tx.Rollback() - if errors.Is(err, context.DeadlineExceeded) || errors.Is(timeoutCtx.Err(), context.DeadlineExceeded) { - log.WithContext(ctx).Warnf("transaction exceeded %s timeout after %v, stack: %s", s.transactionTimeout, time.Since(startTime), debug.Stack()) - } - return err - } - - // Re-enable FK checks before commit (optional, as transaction end resets it) - if s.storeEngine == types.MysqlStoreEngine { - if err := tx.Exec("SET FOREIGN_KEY_CHECKS = 1").Error; err != nil { - tx.Rollback() - return fmt.Errorf("failed to re-enable FK checks: %w", err) - } - } - - err = tx.Commit().Error - if err != nil { - if errors.Is(err, context.DeadlineExceeded) || errors.Is(timeoutCtx.Err(), context.DeadlineExceeded) { - log.WithContext(ctx).Warnf("transaction commit exceeded %s timeout after %v, stack: %s", s.transactionTimeout, time.Since(startTime), debug.Stack()) - } - return err - } - - log.WithContext(ctx).Tracef("transaction took %v", time.Since(startTime)) - if s.metrics != nil { - s.metrics.StoreMetrics().CountTransactionDuration(time.Since(startTime)) - } - - return nil + return s.conn.RunInTx(ctx, func(tx *db.Tx) error { + return operation(s.withTx(tx)) + }) } -func (s *SqlStore) withTx(tx *gorm.DB) Store { +func (s *SqlStore) withTx(tx *db.Tx) Store { return &SqlStore{ - db: tx, - storeEngine: s.storeEngine, + conn: s.conn, + db: s.conn.DB(tx), + tx: tx, fieldEncrypt: s.fieldEncrypt, } } -// transaction wraps a GORM transaction with MySQL-specific FK checks handling -// Use this instead of db.Transaction() directly to avoid deadlocks on MySQL/Aurora -func (s *SqlStore) transaction(fn func(*gorm.DB) error) error { - return s.db.Transaction(func(tx *gorm.DB) error { - // For MySQL, disable FK checks within this transaction to avoid deadlocks - // This is session-scoped and doesn't require SUPER privileges - if s.storeEngine == types.MysqlStoreEngine { - if err := tx.Exec("SET FOREIGN_KEY_CHECKS = 0").Error; err != nil { - return fmt.Errorf("failed to disable FK checks: %w", err) - } - } - - err := fn(tx) - - // Re-enable FK checks before commit (optional, as transaction end resets it) - if s.storeEngine == types.MysqlStoreEngine && err == nil { - if fkErr := tx.Exec("SET FOREIGN_KEY_CHECKS = 1").Error; fkErr != nil { - return fmt.Errorf("failed to re-enable FK checks: %w", fkErr) - } - } - - return err +// transaction runs fn as a savepoint of the bound transaction, or in a new +// transaction when the store is not bound to one. +func (s *SqlStore) transaction(ctx context.Context, fn func(tx *gorm.DB) error) error { + if s.tx != nil { + return s.db.Transaction(fn) + } + return s.conn.RunInTx(ctx, func(tx *db.Tx) error { + return fn(s.conn.DB(tx)) }) } diff --git a/management/server/store/sql_store_account.go b/management/server/store/sql_store_account.go index f4cb15a3c..8729c1f5a 100644 --- a/management/server/store/sql_store_account.go +++ b/management/server/store/sql_store_account.go @@ -51,7 +51,7 @@ func (s *SqlStore) SaveAccount(ctx context.Context, account *types.Account) erro group.StoreGroupPeers() } - err := s.transaction(func(tx *gorm.DB) error { + err := s.transaction(ctx, func(tx *gorm.DB) error { result := tx.Select(clause.Associations).Delete(account.Policies, "account_id = ?", account.Id) if result.Error != nil { return result.Error @@ -146,7 +146,7 @@ func (s *SqlStore) checkAccountDomainBeforeSave(ctx context.Context, accountID, func (s *SqlStore) DeleteAccount(ctx context.Context, account *types.Account) error { start := time.Now() - err := s.transaction(func(tx *gorm.DB) error { + err := s.transaction(ctx, func(tx *gorm.DB) error { result := tx.Select(clause.Associations).Delete(account.Policies, "account_id = ?", account.Id) if result.Error != nil { return result.Error @@ -281,7 +281,7 @@ func (s *SqlStore) GetAccountMeta(ctx context.Context, lockStrength LockingStren } func (s *SqlStore) GetAccount(ctx context.Context, accountID string) (*types.Account, error) { - if s.pool != nil { + if s.pgxPool() != nil { return s.getAccountPgx(ctx, accountID) } return s.getAccountGorm(ctx, accountID) @@ -761,7 +761,7 @@ func (s *SqlStore) getAccount(ctx context.Context, accountID string) (*types.Acc networkSerial sql.NullInt64 createdAt sql.NullTime ) - err := s.pool.QueryRow(ctx, accountQuery, accountID).Scan( + err := s.pgxPool().QueryRow(ctx, accountQuery, accountID).Scan( &account.Id, &account.CreatedBy, &createdAt, &account.Domain, &account.DomainCategory, &account.IsDomainPrimaryAccount, &networkIdentifier, &networkNet, &networkNetV6, &networkDns, &networkSerial, &dnsSettingsDisabledGroups, diff --git a/management/server/store/sql_store_account_onboarding.go b/management/server/store/sql_store_account_onboarding.go index 5872ea63e..73c8f14d0 100644 --- a/management/server/store/sql_store_account_onboarding.go +++ b/management/server/store/sql_store_account_onboarding.go @@ -44,7 +44,7 @@ func (s *SqlStore) getAccountOnboarding(ctx context.Context, accountID string, a const query = `SELECT account_id, onboarding_flow_pending, signup_form_pending, created_at, updated_at FROM account_onboardings WHERE account_id = $1` var onboardingFlowPending, signupFormPending sql.NullBool var createdAt, updatedAt sql.NullTime - err := s.pool.QueryRow(ctx, query, accountID).Scan( + err := s.pgxPool().QueryRow(ctx, query, accountID).Scan( &account.Onboarding.AccountID, &onboardingFlowPending, &signupFormPending, diff --git a/management/server/store/sql_store_agent_network_access_log.go b/management/server/store/sql_store_agent_network_access_log.go index fdaea1bf7..a1ae0150d 100644 --- a/management/server/store/sql_store_agent_network_access_log.go +++ b/management/server/store/sql_store_agent_network_access_log.go @@ -16,7 +16,7 @@ import ( // entry together with its authorising-group child rows in a single // transaction. func (s *SqlStore) CreateAgentNetworkAccessLog(ctx context.Context, entry *agentNetworkTypes.AgentNetworkAccessLog, groups []agentNetworkTypes.AgentNetworkAccessLogGroup) error { - err := s.db.Transaction(func(tx *gorm.DB) error { + err := s.transaction(ctx, func(tx *gorm.DB) error { // Idempotent on the log id / (log_id, group_id) so a proxy resend of the // same entry can't fail the request. if err := tx.Clauses(clause.OnConflict{DoNothing: true}).Create(entry).Error; err != nil { @@ -46,7 +46,7 @@ func (s *SqlStore) CreateAgentNetworkAccessLog(ctx context.Context, entry *agent // deleted. func (s *SqlStore) DeleteOldAgentNetworkAccessLogs(ctx context.Context, accountID string, olderThan time.Time) (int64, error) { var deleted int64 - err := s.db.Transaction(func(tx *gorm.DB) error { + err := s.transaction(ctx, func(tx *gorm.DB) error { // Remove group child rows for the soon-to-be-deleted logs first. if err := tx.Exec( "DELETE FROM agent_network_access_log_group WHERE account_id = ? AND log_id IN (SELECT id FROM agent_network_access_log WHERE account_id = ? AND timestamp < ?)", diff --git a/management/server/store/sql_store_agent_network_usage.go b/management/server/store/sql_store_agent_network_usage.go index c34b31e0e..dc557d2f4 100644 --- a/management/server/store/sql_store_agent_network_usage.go +++ b/management/server/store/sql_store_agent_network_usage.go @@ -14,7 +14,7 @@ import ( // CreateAgentNetworkUsage persists a stripped agent-network usage record // together with its authorising-group child rows in a single transaction. func (s *SqlStore) CreateAgentNetworkUsage(ctx context.Context, usage *agentNetworkTypes.AgentNetworkUsage, groups []agentNetworkTypes.AgentNetworkUsageGroup) error { - err := s.db.Transaction(func(tx *gorm.DB) error { + err := s.transaction(ctx, func(tx *gorm.DB) error { // Idempotent on the usage id / (usage_id, group_id) so a proxy resend of // the same entry can't fail the request. if err := tx.Clauses(clause.OnConflict{DoNothing: true}).Create(usage).Error; err != nil { diff --git a/management/server/store/sql_store_agentnetwork.go b/management/server/store/sql_store_agentnetwork.go index 8a92f7147..4fe77d994 100644 --- a/management/server/store/sql_store_agentnetwork.go +++ b/management/server/store/sql_store_agentnetwork.go @@ -619,7 +619,7 @@ func (s *SqlStore) IncrementAgentNetworkConsumptionBatch( } const tbl = "agent_network_consumption" - err := s.db.Transaction(func(tx *gorm.DB) error { + err := s.transaction(ctx, func(tx *gorm.DB) error { for _, k := range keys { if k.DimID == "" || k.WindowSeconds <= 0 { return status.Errorf(status.InvalidArgument, "dim_id and window_seconds must be set") diff --git a/management/server/store/sql_store_group.go b/management/server/store/sql_store_group.go index 16e812048..323a45734 100644 --- a/management/server/store/sql_store_group.go +++ b/management/server/store/sql_store_group.go @@ -37,7 +37,7 @@ func (s *SqlStore) CreateGroups(ctx context.Context, accountID string, groups [] return nil } - return s.db.Transaction(func(tx *gorm.DB) error { + return s.transaction(ctx, func(tx *gorm.DB) error { result := tx. Clauses( clause.OnConflict{ @@ -63,7 +63,7 @@ func (s *SqlStore) UpdateGroups(ctx context.Context, accountID string, groups [] return nil } - return s.db.Transaction(func(tx *gorm.DB) error { + return s.transaction(ctx, func(tx *gorm.DB) error { result := tx. Clauses( clause.OnConflict{ @@ -137,7 +137,7 @@ func (s *SqlStore) GetResourceGroups(ctx context.Context, lockStrength LockingSt func (s *SqlStore) getGroups(ctx context.Context, accountID string) ([]*types.Group, error) { const query = `SELECT id, account_id, public_id, name, issued, resources, integration_ref_id, integration_ref_integration_type FROM groups WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) + rows, err := s.pgxPool().Query(ctx, query, accountID) if err != nil { return nil, err } diff --git a/management/server/store/sql_store_group_peer.go b/management/server/store/sql_store_group_peer.go index 6b0339816..cc4c44ad1 100644 --- a/management/server/store/sql_store_group_peer.go +++ b/management/server/store/sql_store_group_peer.go @@ -19,7 +19,7 @@ func (s *SqlStore) getGroupPeers(ctx context.Context, groupIDs []string) ([]type return nil, nil } const query = `SELECT account_id, group_id, peer_id FROM group_peers WHERE group_id = ANY($1)` - rows, err := s.pool.Query(ctx, query, groupIDs) + rows, err := s.pgxPool().Query(ctx, query, groupIDs) if err != nil { return nil, err } diff --git a/management/server/store/sql_store_idp_migration.go b/management/server/store/sql_store_idp_migration.go index 64962845b..760d2967c 100644 --- a/management/server/store/sql_store_idp_migration.go +++ b/management/server/store/sql_store_idp_migration.go @@ -57,11 +57,11 @@ func (s *SqlStore) ListUsers(ctx context.Context) ([]*types.User, error) { // txDeferFKConstraints defers foreign key constraint checks for the duration of the transaction. // MySQL is already handled by s.transaction (SET FOREIGN_KEY_CHECKS = 0). func (s *SqlStore) txDeferFKConstraints(tx *gorm.DB) error { - if s.storeEngine == types.SqliteStoreEngine { + if s.conn.Engine() == types.SqliteStoreEngine { return tx.Exec("PRAGMA defer_foreign_keys = ON").Error } - if s.storeEngine != types.PostgresStoreEngine { + if s.conn.Engine() != types.PostgresStoreEngine { return nil } @@ -86,7 +86,7 @@ func (s *SqlStore) txDeferFKConstraints(tx *gorm.DB) error { // txRestoreFKConstraints reverts FK constraints back to NOT DEFERRABLE after the // deferred updates are done but before the transaction commits. func (s *SqlStore) txRestoreFKConstraints(tx *gorm.DB) error { - if s.storeEngine != types.PostgresStoreEngine { + if s.conn.Engine() != types.PostgresStoreEngine { return nil } @@ -138,7 +138,7 @@ func (s *SqlStore) UpdateUserID(ctx context.Context, accountID, oldUserID, newUs } log.Info("Updating user ID in the store") - err := s.transaction(func(tx *gorm.DB) error { + err := s.transaction(ctx, func(tx *gorm.DB) error { if err := s.txDeferFKConstraints(tx); err != nil { return err } @@ -161,7 +161,7 @@ func (s *SqlStore) UpdateUserID(ctx context.Context, accountID, oldUserID, newUs } log.Info("Restoring FK constraints") - err = s.transaction(func(tx *gorm.DB) error { + err = s.transaction(ctx, func(tx *gorm.DB) error { if err := s.txRestoreFKConstraints(tx); err != nil { return fmt.Errorf("restore FK constraints: %w", err) } diff --git a/management/server/store/sql_store_name_server_group.go b/management/server/store/sql_store_name_server_group.go index 4af831148..595921913 100644 --- a/management/server/store/sql_store_name_server_group.go +++ b/management/server/store/sql_store_name_server_group.go @@ -17,7 +17,7 @@ import ( func (s *SqlStore) getNameServerGroups(ctx context.Context, accountID string) ([]nbdns.NameServerGroup, error) { const query = `SELECT id, account_id, public_id, name, description, name_servers, groups, "primary", domains, enabled, search_domains_enabled FROM name_server_groups WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) + rows, err := s.pgxPool().Query(ctx, query, accountID) if err != nil { return nil, err } diff --git a/management/server/store/sql_store_network.go b/management/server/store/sql_store_network.go index 963d67e8f..66333a4fb 100644 --- a/management/server/store/sql_store_network.go +++ b/management/server/store/sql_store_network.go @@ -15,7 +15,7 @@ import ( func (s *SqlStore) getNetworks(ctx context.Context, accountID string) ([]*networkTypes.Network, error) { const query = `SELECT id, account_id, public_id, name, description FROM networks WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) + rows, err := s.pgxPool().Query(ctx, query, accountID) if err != nil { return nil, err } diff --git a/management/server/store/sql_store_network_resource.go b/management/server/store/sql_store_network_resource.go index 659382c61..dd4352ea0 100644 --- a/management/server/store/sql_store_network_resource.go +++ b/management/server/store/sql_store_network_resource.go @@ -17,7 +17,7 @@ import ( func (s *SqlStore) getNetworkResources(ctx context.Context, accountID string) ([]*resourceTypes.NetworkResource, error) { const query = `SELECT id, network_id, account_id, public_id, name, description, type, domain, prefix, enabled FROM network_resources WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) + rows, err := s.pgxPool().Query(ctx, query, accountID) if err != nil { return nil, err } diff --git a/management/server/store/sql_store_network_router.go b/management/server/store/sql_store_network_router.go index b2594d483..bb5cb6621 100644 --- a/management/server/store/sql_store_network_router.go +++ b/management/server/store/sql_store_network_router.go @@ -19,7 +19,7 @@ import ( func (s *SqlStore) getNetworkRouters(ctx context.Context, accountID string) ([]*routerTypes.NetworkRouter, error) { const query = `SELECT id, network_id, account_id, public_id, peer, peer_groups, masquerade, metric, enabled FROM network_routers WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) + rows, err := s.pgxPool().Query(ctx, query, accountID) if err != nil { return nil, err } diff --git a/management/server/store/sql_store_peer.go b/management/server/store/sql_store_peer.go index 90c742a72..e5086b6db 100644 --- a/management/server/store/sql_store_peer.go +++ b/management/server/store/sql_store_peer.go @@ -25,7 +25,7 @@ func (s *SqlStore) SavePeer(ctx context.Context, accountID string, peer *nbpeer. peerCopy := peer.Copy() peerCopy.AccountID = accountID - err := s.transaction(func(tx *gorm.DB) error { + err := s.transaction(ctx, func(tx *gorm.DB) error { // check if peer exists before saving var peerID string result := tx.Model(&nbpeer.Peer{}).Select("id").Take(&peerID, accountAndIDQueryCondition, accountID, peer.ID) @@ -190,7 +190,7 @@ func (s *SqlStore) getPeers(ctx context.Context, accountID string) ([]nbpeer.Pee peer_status_connected, peer_status_login_expired, peer_status_requires_approval, location_connection_ip, location_country_code, location_city_name, location_geo_name_id, proxy_meta_embedded, proxy_meta_cluster, ipv6, meta_sync_message_version FROM peers WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) + rows, err := s.pgxPool().Query(ctx, query, accountID) if err != nil { return nil, err } diff --git a/management/server/store/sql_store_personal_access_token.go b/management/server/store/sql_store_personal_access_token.go index 351d122fd..d5be3e327 100644 --- a/management/server/store/sql_store_personal_access_token.go +++ b/management/server/store/sql_store_personal_access_token.go @@ -45,7 +45,7 @@ func (s *SqlStore) getPersonalAccessTokens(ctx context.Context, userIDs []string return nil, nil } const query = `SELECT id, user_id, name, hashed_token, expiration_date, created_by, created_at, last_used FROM personal_access_tokens WHERE user_id = ANY($1)` - rows, err := s.pool.Query(ctx, query, userIDs) + rows, err := s.pgxPool().Query(ctx, query, userIDs) if err != nil { return nil, err } diff --git a/management/server/store/sql_store_policy.go b/management/server/store/sql_store_policy.go index 7725dede5..95e80e712 100644 --- a/management/server/store/sql_store_policy.go +++ b/management/server/store/sql_store_policy.go @@ -18,7 +18,7 @@ import ( func (s *SqlStore) getPolicies(ctx context.Context, accountID string) ([]*types.Policy, error) { const query = `SELECT id, account_id, public_id, name, description, enabled, source_posture_checks FROM policies WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) + rows, err := s.pgxPool().Query(ctx, query, accountID) if err != nil { return nil, err } @@ -128,7 +128,7 @@ func (s *SqlStore) SavePolicy(ctx context.Context, policy *types.Policy) error { } func (s *SqlStore) DeletePolicy(ctx context.Context, accountID, policyID string) error { - return s.transaction(func(tx *gorm.DB) error { + return s.transaction(ctx, func(tx *gorm.DB) error { if err := tx.Where("policy_id = ?", policyID).Delete(&types.PolicyRule{}).Error; err != nil { return fmt.Errorf("delete policy rules: %w", err) } diff --git a/management/server/store/sql_store_policy_rule.go b/management/server/store/sql_store_policy_rule.go index d1d6ef3d3..f822788ac 100644 --- a/management/server/store/sql_store_policy_rule.go +++ b/management/server/store/sql_store_policy_rule.go @@ -18,7 +18,7 @@ func (s *SqlStore) getPolicyRules(ctx context.Context, policyIDs []string) ([]*t return nil, nil } const query = `SELECT id, policy_id, name, description, enabled, action, destinations, destination_resource, sources, source_resource, bidirectional, protocol, ports, port_ranges, authorized_groups, authorized_user FROM policy_rules WHERE policy_id = ANY($1)` - rows, err := s.pool.Query(ctx, query, policyIDs) + rows, err := s.pgxPool().Query(ctx, query, policyIDs) if err != nil { return nil, err } diff --git a/management/server/store/sql_store_posture_checks.go b/management/server/store/sql_store_posture_checks.go index 71997ec73..d951086d4 100644 --- a/management/server/store/sql_store_posture_checks.go +++ b/management/server/store/sql_store_posture_checks.go @@ -16,7 +16,7 @@ import ( func (s *SqlStore) getPostureChecks(ctx context.Context, accountID string) ([]*posture.Checks, error) { const query = `SELECT id, account_id, public_id, name, description, checks FROM posture_checks WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) + rows, err := s.pgxPool().Query(ctx, query, accountID) if err != nil { return nil, err } diff --git a/management/server/store/sql_store_route.go b/management/server/store/sql_store_route.go index 0aa0b399a..ec9e130b1 100644 --- a/management/server/store/sql_store_route.go +++ b/management/server/store/sql_store_route.go @@ -17,7 +17,7 @@ import ( func (s *SqlStore) getRoutes(ctx context.Context, accountID string) ([]route.Route, error) { const query = `SELECT id, account_id, public_id, network, domains, keep_route, net_id, description, peer, peer_groups, network_type, masquerade, metric, enabled, groups, access_control_groups, skip_auto_apply FROM routes WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) + rows, err := s.pgxPool().Query(ctx, query, accountID) if err != nil { return nil, err } diff --git a/management/server/store/sql_store_service.go b/management/server/store/sql_store_service.go index 5eedbcb37..4f4f3546d 100644 --- a/management/server/store/sql_store_service.go +++ b/management/server/store/sql_store_service.go @@ -30,7 +30,7 @@ const serviceSelectColumns = `id, account_id, name, domain, enabled, auth, restr func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpservice.Service, error) { const serviceQuery = `SELECT ` + serviceSelectColumns + ` FROM services WHERE account_id = $1` - serviceRows, err := s.pool.Query(ctx, serviceQuery, accountID) + serviceRows, err := s.pgxPool().Query(ctx, serviceQuery, accountID) if err != nil { return nil, err } @@ -205,7 +205,7 @@ func (s *SqlStore) UpdateService(ctx context.Context, service *rpservice.Service targetType := &rpservice.Target{} // Use a transaction to ensure atomic updates of the service and its targets - err := s.db.Transaction(func(tx *gorm.DB) error { + err := s.transaction(ctx, func(tx *gorm.DB) error { // Delete existing targets if err := tx.Where("service_id = ?", serviceCopy.ID).Delete(targetType).Error; err != nil { return err diff --git a/management/server/store/sql_store_service_target.go b/management/server/store/sql_store_service_target.go index c4531a158..5c2b99510 100644 --- a/management/server/store/sql_store_service_target.go +++ b/management/server/store/sql_store_service_target.go @@ -26,7 +26,7 @@ const targetSelectColumns = `id, account_id, service_id, path, host, port, proto func (s *SqlStore) getServiceTargets(ctx context.Context, serviceIDs []string) ([]*rpservice.Target, error) { const targetsQuery = `SELECT ` + targetSelectColumns + ` FROM targets WHERE service_id = ANY($1)` - rows, err := s.pool.Query(ctx, targetsQuery, serviceIDs) + rows, err := s.pgxPool().Query(ctx, targetsQuery, serviceIDs) if err != nil { return nil, err } diff --git a/management/server/store/sql_store_setup_key.go b/management/server/store/sql_store_setup_key.go index 79fb406ad..3857f6e8c 100644 --- a/management/server/store/sql_store_setup_key.go +++ b/management/server/store/sql_store_setup_key.go @@ -37,7 +37,7 @@ func (s *SqlStore) GetAccountBySetupKey(ctx context.Context, setupKey string) (* func (s *SqlStore) getSetupKeys(ctx context.Context, accountID string) ([]types.SetupKey, error) { const query = `SELECT id, account_id, key, key_secret, name, type, created_at, expires_at, updated_at, revoked, used_times, last_used, auto_groups, usage_limit, ephemeral, allow_extra_dns_labels FROM setup_keys WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) + rows, err := s.pgxPool().Query(ctx, query, accountID) if err != nil { return nil, err } diff --git a/management/server/store/sql_store_test.go b/management/server/store/sql_store_test.go index 731b90ce9..f695a150d 100644 --- a/management/server/store/sql_store_test.go +++ b/management/server/store/sql_store_test.go @@ -15,6 +15,8 @@ import ( log "github.com/sirupsen/logrus" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "gorm.io/gorm" + "gorm.io/gorm/clause" nbdns "github.com/netbirdio/netbird/dns" nbpeer "github.com/netbirdio/netbird/management/server/peer" @@ -346,7 +348,6 @@ func TestSqlStore_ExecuteInTransaction_Timeout(t *testing.T) { sqlStore, ok := store.(*SqlStore) require.True(t, ok) - assert.Equal(t, 1*time.Second, sqlStore.transactionTimeout) ctx := context.Background() err = sqlStore.ExecuteInTransaction(ctx, func(transaction Store) error { @@ -411,3 +412,88 @@ func TestNewSqliteStore_BusyTimeoutRespectsUserOverride(t *testing.T) { }) } } + +func TestSqlStore_ExecuteInTransaction_RestoresForeignKeyChecksOnMysql(t *testing.T) { + runTestForAllEngines(t, "", func(t *testing.T, store Store) { + sqlStore := store.(*SqlStore) + if sqlStore.conn.Engine() != types.MysqlStoreEngine { + t.Skip("FOREIGN_KEY_CHECKS is MySQL specific") + } + sqlDB, err := sqlStore.GetDB().DB() + require.NoError(t, err) + sqlDB.SetMaxOpenConns(1) + ctx := context.Background() + + foreignKeyChecks := func() int { + var enabled int + require.NoError(t, sqlStore.GetDB().Raw("SELECT @@SESSION.foreign_key_checks").Scan(&enabled).Error) + return enabled + } + + err = store.ExecuteInTransaction(ctx, func(Store) error { return assert.AnError }) + require.ErrorIs(t, err, assert.AnError) + assert.Equal(t, 1, foreignKeyChecks()) + + require.Panics(t, func() { + _ = store.ExecuteInTransaction(ctx, func(Store) error { panic("boom") }) + }) + assert.Equal(t, 1, foreignKeyChecks()) + + err = sqlStore.transaction(ctx, func(*gorm.DB) error { return assert.AnError }) + require.ErrorIs(t, err, assert.AnError) + assert.Equal(t, 1, foreignKeyChecks()) + + err = store.ExecuteInTransaction(ctx, func(transaction Store) error { + bound := transaction.(*SqlStore) + require.NoError(t, bound.transaction(ctx, func(*gorm.DB) error { return nil })) + var enabled int + require.NoError(t, bound.GetDB().Raw("SELECT @@SESSION.foreign_key_checks").Scan(&enabled).Error) + assert.Equal(t, 0, enabled, "a savepoint must not re-enable FK checks for the rest of the transaction") + return nil + }) + require.NoError(t, err) + assert.Equal(t, 1, foreignKeyChecks()) + }) +} + +func TestSqlStore_Transaction_RollsBackOnError(t *testing.T) { + runTestForAllEngines(t, "../testdata/extended-store.sql", func(t *testing.T, store Store) { + ctx := context.Background() + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + group := &types.Group{ID: "rolled-back-group", AccountID: accountID, Name: "rolled back", Issued: "api"} + + err := store.(*SqlStore).transaction(ctx, func(tx *gorm.DB) error { + require.NoError(t, tx.Omit(clause.Associations).Create(group).Error) + return assert.AnError + }) + require.ErrorIs(t, err, assert.AnError) + + _, err = store.GetGroupByID(ctx, LockingStrengthNone, accountID, group.ID) + require.Error(t, err) + }) +} + +func TestSqlStore_Transaction_NestedIsSavepoint(t *testing.T) { + runTestForAllEngines(t, "../testdata/extended-store.sql", func(t *testing.T, store Store) { + ctx := context.Background() + accountID := "bf1c8084-ba50-4ce7-9439-34653001fc3b" + outer := &types.Group{ID: "outer-group", AccountID: accountID, Name: "outer", Issued: "api"} + inner := &types.Group{ID: "inner-group", AccountID: accountID, Name: "inner", Issued: "api"} + + err := store.ExecuteInTransaction(ctx, func(transaction Store) error { + require.NoError(t, transaction.CreateGroup(ctx, outer)) + err := transaction.(*SqlStore).transaction(ctx, func(tx *gorm.DB) error { + require.NoError(t, tx.Omit(clause.Associations).Create(inner).Error) + return assert.AnError + }) + require.ErrorIs(t, err, assert.AnError) + return nil + }) + require.NoError(t, err) + + _, err = store.GetGroupByID(ctx, LockingStrengthNone, accountID, outer.ID) + require.NoError(t, err) + _, err = store.GetGroupByID(ctx, LockingStrengthNone, accountID, inner.ID) + require.Error(t, err) + }) +} diff --git a/management/server/store/sql_store_user.go b/management/server/store/sql_store_user.go index 2ead2156a..4db03fb71 100644 --- a/management/server/store/sql_store_user.go +++ b/management/server/store/sql_store_user.go @@ -108,7 +108,7 @@ func (s *SqlStore) GetUserByUserID(ctx context.Context, lockStrength LockingStre } func (s *SqlStore) DeleteUser(ctx context.Context, accountID, userID string) error { - err := s.transaction(func(tx *gorm.DB) error { + err := s.transaction(ctx, func(tx *gorm.DB) error { result := tx.Delete(&types.PersonalAccessToken{}, "user_id = ?", userID) if result.Error != nil { return result.Error @@ -173,7 +173,7 @@ func (s *SqlStore) GetAccountOwner(ctx context.Context, lockStrength LockingStre func (s *SqlStore) getUsers(ctx context.Context, accountID string) ([]types.User, error) { const query = `SELECT id, account_id, role, is_service_user, non_deletable, service_user_name, auto_groups, blocked, pending_approval, last_login, created_at, issued, integration_ref_id, integration_ref_integration_type, email, name FROM users WHERE account_id = $1` - rows, err := s.pool.Query(ctx, query, accountID) + rows, err := s.pgxPool().Query(ctx, query, accountID) if err != nil { return nil, err } diff --git a/management/server/store/sqlstore_bench_test.go b/management/server/store/sqlstore_bench_test.go index a38b4a8c1..b1933b0d3 100644 --- a/management/server/store/sqlstore_bench_test.go +++ b/management/server/store/sqlstore_bench_test.go @@ -22,6 +22,7 @@ import ( nbdns "github.com/netbirdio/netbird/dns" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" + nbdb "github.com/netbirdio/netbird/management/internals/shared/db" resourceTypes "github.com/netbirdio/netbird/management/server/networks/resources/types" routerTypes "github.com/netbirdio/netbird/management/server/networks/routers/types" networkTypes "github.com/netbirdio/netbird/management/server/networks/types" @@ -281,10 +282,11 @@ func setupBenchmarkDB(b testing.TB) (*SqlStore, func(), string) { b.Fatalf("failed to migrate database: %v", err) } - store := &SqlStore{ - db: db, - pool: pool, + conn, err := nbdb.NewConn(context.Background(), db, nbdb.PostgresStoreEngine, pool) + if err != nil { + b.Fatalf("failed to create connection: %v", err) } + store := &SqlStore{conn: conn, db: conn.DB(nil)} const ( accountID = "benchmark-account-id" @@ -527,7 +529,7 @@ func BenchmarkGetAccount(b *testing.B) { } } }) - store.pool.Close() + _ = store.Close(ctx) } func TestAccountEquivalence(t *testing.T) { diff --git a/management/server/store/store.go b/management/server/store/store.go index ed53690a5..e935d2d88 100644 --- a/management/server/store/store.go +++ b/management/server/store/store.go @@ -27,12 +27,12 @@ import ( "gorm.io/gorm" "github.com/netbirdio/netbird/dns" - "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain" "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy" rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" "github.com/netbirdio/netbird/management/internals/modules/zones" "github.com/netbirdio/netbird/management/internals/modules/zones/records" + "github.com/netbirdio/netbird/management/internals/shared/db" "github.com/netbirdio/netbird/management/server/telemetry" "github.com/netbirdio/netbird/management/server/testutil" "github.com/netbirdio/netbird/management/server/types" @@ -50,14 +50,14 @@ import ( "github.com/netbirdio/netbird/route" ) -type LockingStrength string +type LockingStrength = db.LockingStrength const ( - LockingStrengthUpdate LockingStrength = "UPDATE" // Strongest lock, preventing any changes by other transactions until your transaction completes. - LockingStrengthShare LockingStrength = "SHARE" // Allows reading but prevents changes by other transactions. - LockingStrengthNoKeyUpdate LockingStrength = "NO KEY UPDATE" // Similar to UPDATE but allows changes to related rows. - LockingStrengthKeyShare LockingStrength = "KEY SHARE" // Protects against changes to primary/unique keys but allows other updates. - LockingStrengthNone LockingStrength = "NONE" // No locking, allowing all transactions to proceed without restrictions. + LockingStrengthUpdate = db.LockingStrengthUpdate + LockingStrengthShare = db.LockingStrengthShare + LockingStrengthNoKeyUpdate = db.LockingStrengthNoKeyUpdate + LockingStrengthKeyShare = db.LockingStrengthKeyShare + LockingStrengthNone = db.LockingStrengthNone ) type Store interface { @@ -316,9 +316,6 @@ type Store interface { 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 - GetAccountAccessLogs(ctx context.Context, lockStrength LockingStrength, accountID string, filter accesslogs.AccessLogFilter) ([]*accesslogs.AccessLogEntry, int64, error) - DeleteOldAccessLogs(ctx context.Context, olderThan time.Time) (int64, error) CreateAgentNetworkAccessLog(ctx context.Context, entry *agentNetworkTypes.AgentNetworkAccessLog, groups []agentNetworkTypes.AgentNetworkAccessLogGroup) error CreateAgentNetworkUsage(ctx context.Context, usage *agentNetworkTypes.AgentNetworkUsage, groups []agentNetworkTypes.AgentNetworkUsageGroup) error GetAgentNetworkAccessLogs(ctx context.Context, lockStrength LockingStrength, accountID string, filter agentNetworkTypes.AgentNetworkAccessLogFilter) ([]*agentNetworkTypes.AgentNetworkAccessLog, int64, error) @@ -492,7 +489,7 @@ func getStoreEngine(ctx context.Context, dataDir string, kind types.Engine) type // Migrate if it is the first run with a JSON file existing and no SQLite file present jsonStoreFile := filepath.Join(dataDir, storeFileName) - sqliteStoreFile := filepath.Join(dataDir, storeSqliteFileName) + sqliteStoreFile := filepath.Join(dataDir, db.SqliteFileName) if util.FileExists(jsonStoreFile) && !util.FileExists(sqliteStoreFile) { log.WithContext(ctx).Warnf("unsupported store engine specified, but found %s. Automatically migrating to SQLite.", jsonStoreFile) @@ -511,6 +508,16 @@ func getStoreEngine(ctx context.Context, dataDir string, kind types.Engine) type // NewStore creates a new store based on the provided engine type, data directory, and telemetry metrics func NewStore(ctx context.Context, kind types.Engine, dataDir string, metrics telemetry.AppMetrics, skipMigration bool) (Store, error) { + conn, err := OpenConn(ctx, kind, dataDir) + if err != nil { + return nil, err + } + return newStore(ctx, conn, metrics, skipMigration) +} + +// OpenConn resolves the configured engine and opens the connection that the +// store and the domain repositories share. +func OpenConn(ctx context.Context, kind types.Engine, dataDir string) (*db.Conn, error) { kind = getStoreEngine(ctx, dataDir, kind) if err := checkFileStoreEngine(kind, dataDir); err != nil { @@ -520,13 +527,21 @@ func NewStore(ctx context.Context, kind types.Engine, dataDir string, metrics te switch kind { case types.SqliteStoreEngine: log.WithContext(ctx).Info("using SQLite store engine") - return NewSqliteStore(ctx, dataDir, metrics, skipMigration) + return db.OpenSqlite(ctx, dataDir) case types.PostgresStoreEngine: log.WithContext(ctx).Info("using Postgres store engine") - return newPostgresStore(ctx, metrics, skipMigration) + dsn, ok := lookupDSNEnv(PostgresDsnEnv, PostgresDsnEnvLegacy) + if !ok { + return nil, fmt.Errorf("%s is not set", PostgresDsnEnv) + } + return db.OpenPostgres(ctx, dsn, db.DefaultPoolConfig) case types.MysqlStoreEngine: log.WithContext(ctx).Info("using MySQL store engine") - return newMysqlStore(ctx, metrics, skipMigration) + dsn, ok := lookupDSNEnv(mysqlDsnEnv, mysqlDsnEnvLegacy) + if !ok { + return nil, fmt.Errorf("%s is not set", mysqlDsnEnv) + } + return db.OpenMysql(ctx, dsn) default: return nil, fmt.Errorf("unsupported kind of store: %s", kind) } @@ -700,29 +715,34 @@ func NewTestStoreFromSQL(ctx context.Context, filename string, dataDir string) ( kind = types.SqliteStoreEngine } - storeStr := fmt.Sprintf("%s?cache=shared", storeSqliteFileName) + storeStr := fmt.Sprintf("%s?cache=shared", db.SqliteFileName) if runtime.GOOS == "windows" { // Vo avoid `The process cannot access the file because it is being used by another process` on Windows - storeStr = storeSqliteFileName + storeStr = db.SqliteFileName } file := filepath.Join(dataDir, storeStr) - db, err := gorm.Open(sqlite.Open(file), getGormConfig()) + gormDB, err := gorm.Open(sqlite.Open(file), db.GormConfig()) if err != nil { return nil, nil, err } if filename != "" { - err = LoadSQL(db, filename) + err = LoadSQL(gormDB, filename) if err != nil { return nil, nil, fmt.Errorf("failed to load SQL file: %v", err) } } - store, err := NewSqlStore(ctx, db, types.SqliteStoreEngine, nil, false) + conn, err := db.NewConn(ctx, gormDB, db.SqliteStoreEngine, nil) if err != nil { return nil, nil, fmt.Errorf("failed to create test store: %v", err) } + store, err := NewSqlStore(ctx, conn, nil, false) + if err != nil { + _ = conn.Close() + return nil, nil, fmt.Errorf("failed to create test store: %v", err) + } err = addAllGroupToAccount(ctx, store) if err != nil { @@ -790,9 +810,6 @@ func getSqlStoreEngine(ctx context.Context, sqliteStore *SqlStore, kind types.En closeConnection := func() { cleanup() store.Close(ctx) - if store.pool != nil { - store.pool.Close() - } if store != sqliteStore { // The sqlite store only seeded the engine under test; without this // every test leaks its connection and the opener goroutines. @@ -949,9 +966,6 @@ func postgresSchemaTemplate(ctx context.Context, baseDSN string, admin *gorm.DB) // TEMPLATE refuses a source that still has sessions, so release both handles // before the first clone. tplStore.Close(ctx) - if tplStore.pool != nil { - tplStore.pool.Close() - } schemaTemplates[key] = &schemaTemplate{dbName: name} return name, nil @@ -1036,11 +1050,11 @@ func mysqlTableNames(ctx context.Context, sqlDB *sql.DB, dbName string) ([]strin // cloneMysqlSchema replays the template's CREATE TABLE statements into the // database the DSN points at. func cloneMysqlSchema(ctx context.Context, dsn string, tableDDL []string) error { - db, err := gorm.Open(mysql.Open(dsn+"?charset=utf8&parseTime=True&loc=Local"), getGormConfig()) + gormDB, err := gorm.Open(mysql.Open(db.MysqlDSN(dsn)), db.GormConfig()) if err != nil { return fmt.Errorf("connect to test database: %w", err) } - sqlDB, err := db.DB() + sqlDB, err := gormDB.DB() if err != nil { return err } @@ -1080,19 +1094,19 @@ func closeGormDB(db *gorm.DB) { } func openDBWithRetry(dsn string, engine types.Engine, maxRetries int) (*gorm.DB, error) { - var db *gorm.DB + var gormDB *gorm.DB var err error for i := range maxRetries { switch engine { case types.PostgresStoreEngine: - db, err = gorm.Open(postgres.Open(dsn), &gorm.Config{}) + gormDB, err = gorm.Open(postgres.Open(dsn), &gorm.Config{}) case types.MysqlStoreEngine: - db, err = gorm.Open(mysql.Open(dsn+"?charset=utf8&parseTime=True&loc=Local"), &gorm.Config{}) + gormDB, err = gorm.Open(mysql.Open(db.MysqlDSN(dsn)), &gorm.Config{}) } if err == nil { - return db, nil + return gormDB, nil } if i < maxRetries-1 { @@ -1106,14 +1120,14 @@ func openDBWithRetry(dsn string, engine types.Engine, maxRetries int) (*gorm.DB, // createRandomDB creates a uniquely named database for one test. On postgres a // non-empty template is copied server-side with CREATE DATABASE ... TEMPLATE. -func createRandomDB(dsn string, db *gorm.DB, engine types.Engine, template string) (string, func(), error) { +func createRandomDB(dsn string, admin *gorm.DB, engine types.Engine, template string) (string, func(), error) { dbName := newTestDBName("test_db") createStmt := fmt.Sprintf("CREATE DATABASE %s", dbName) if template != "" && engine == types.PostgresStoreEngine { createStmt = fmt.Sprintf("CREATE DATABASE %s TEMPLATE %s", dbName, template) } - if err := execWithTemplateRetry(db, createStmt); err != nil { + if err := execWithTemplateRetry(admin, createStmt); err != nil { return "", nil, fmt.Errorf("failed to create database: %v", err) } @@ -1148,7 +1162,7 @@ func createRandomDB(dsn string, db *gorm.DB, engine types.Engine, template strin err = dropDB.Exec(fmt.Sprintf("DROP DATABASE IF EXISTS %s WITH (FORCE)", dbName)).Error case types.MysqlStoreEngine: - dropDB, err = gorm.Open(mysql.Open(originalDSN+"?charset=utf8&parseTime=True&loc=Local"), &gorm.Config{ + dropDB, err = gorm.Open(mysql.Open(db.MysqlDSN(originalDSN)), &gorm.Config{ SkipDefaultTransaction: true, PrepareStmt: false, }) @@ -1226,7 +1240,7 @@ func MigrateFileStoreToSqlite(ctx context.Context, dataDir string) error { return fmt.Errorf("%s doesn't exist, couldn't continue the operation", fileStorePath) } - sqlStorePath := path.Join(dataDir, storeSqliteFileName) + sqlStorePath := path.Join(dataDir, db.SqliteFileName) if _, err := os.Stat(sqlStorePath); err == nil { return fmt.Errorf("%s already exists, couldn't continue the operation", sqlStorePath) } diff --git a/management/server/store/store_mock.go b/management/server/store/store_mock.go index 2db74fb86..068ac6ff2 100644 --- a/management/server/store/store_mock.go +++ b/management/server/store/store_mock.go @@ -18,7 +18,6 @@ import ( dns "github.com/netbirdio/netbird/dns" types "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types" - accesslogs "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs" domain "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain" proxy "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy" service "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" @@ -247,20 +246,6 @@ func (mr *MockStoreMockRecorder) CountProxiesByAccountID(ctx, accountID any) *go return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CountProxiesByAccountID", reflect.TypeOf((*MockStore)(nil).CountProxiesByAccountID), ctx, accountID) } -// CreateAccessLog mocks base method. -func (m *MockStore) CreateAccessLog(ctx context.Context, log *accesslogs.AccessLogEntry) error { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "CreateAccessLog", ctx, log) - ret0, _ := ret[0].(error) - return ret0 -} - -// CreateAccessLog indicates an expected call of CreateAccessLog. -func (mr *MockStoreMockRecorder) CreateAccessLog(ctx, log any) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CreateAccessLog", reflect.TypeOf((*MockStore)(nil).CreateAccessLog), ctx, log) -} - // CreateAgentNetworkAccessLog mocks base method. func (m *MockStore) CreateAgentNetworkAccessLog(ctx context.Context, entry *types.AgentNetworkAccessLog, groups []types.AgentNetworkAccessLogGroup) error { m.ctrl.T.Helper() @@ -669,21 +654,6 @@ func (mr *MockStoreMockRecorder) DeleteNetworkRouter(ctx, accountID, routerID an return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteNetworkRouter", reflect.TypeOf((*MockStore)(nil).DeleteNetworkRouter), ctx, accountID, routerID) } -// DeleteOldAccessLogs mocks base method. -func (m *MockStore) DeleteOldAccessLogs(ctx context.Context, olderThan time.Time) (int64, error) { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "DeleteOldAccessLogs", ctx, olderThan) - ret0, _ := ret[0].(int64) - ret1, _ := ret[1].(error) - return ret0, ret1 -} - -// DeleteOldAccessLogs indicates an expected call of DeleteOldAccessLogs. -func (mr *MockStoreMockRecorder) DeleteOldAccessLogs(ctx, olderThan any) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteOldAccessLogs", reflect.TypeOf((*MockStore)(nil).DeleteOldAccessLogs), ctx, olderThan) -} - // DeleteOldAgentNetworkAccessLogs mocks base method. func (m *MockStore) DeleteOldAgentNetworkAccessLogs(ctx context.Context, accountID string, olderThan time.Time) (int64, error) { m.ctrl.T.Helper() @@ -968,22 +938,6 @@ func (mr *MockStoreMockRecorder) GetAccount(ctx, accountID any) *gomock.Call { return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccount", reflect.TypeOf((*MockStore)(nil).GetAccount), ctx, accountID) } -// GetAccountAccessLogs mocks base method. -func (m *MockStore) GetAccountAccessLogs(ctx context.Context, lockStrength LockingStrength, accountID string, filter accesslogs.AccessLogFilter) ([]*accesslogs.AccessLogEntry, int64, error) { - m.ctrl.T.Helper() - ret := m.ctrl.Call(m, "GetAccountAccessLogs", ctx, lockStrength, accountID, filter) - ret0, _ := ret[0].([]*accesslogs.AccessLogEntry) - ret1, _ := ret[1].(int64) - ret2, _ := ret[2].(error) - return ret0, ret1, ret2 -} - -// GetAccountAccessLogs indicates an expected call of GetAccountAccessLogs. -func (mr *MockStoreMockRecorder) GetAccountAccessLogs(ctx, lockStrength, accountID, filter any) *gomock.Call { - mr.mock.ctrl.T.Helper() - return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountAccessLogs", reflect.TypeOf((*MockStore)(nil).GetAccountAccessLogs), ctx, lockStrength, accountID, filter) -} - // GetAccountAgentNetworkBudgetRules mocks base method. func (m *MockStore) GetAccountAgentNetworkBudgetRules(ctx context.Context, lockStrength LockingStrength, accountID string) ([]*types.AccountBudgetRule, error) { m.ctrl.T.Helper() @@ -2177,7 +2131,7 @@ func (m *MockStore) GetNetworkResourceByIDOrPublicID(ctx context.Context, lockSt } // GetNetworkResourceByIDOrPublicID indicates an expected call of GetNetworkResourceByIDOrPublicID. -func (mr *MockStoreMockRecorder) GetNetworkResourceByIDOrPublicID(ctx, lockStrength, accountID, resourceID interface{}) *gomock.Call { +func (mr *MockStoreMockRecorder) GetNetworkResourceByIDOrPublicID(ctx, lockStrength, accountID, resourceID any) *gomock.Call { mr.mock.ctrl.T.Helper() return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetNetworkResourceByIDOrPublicID", reflect.TypeOf((*MockStore)(nil).GetNetworkResourceByIDOrPublicID), ctx, lockStrength, accountID, resourceID) } @@ -2522,7 +2476,7 @@ func (m *MockStore) GetPolicyByIDOrPublicID(ctx context.Context, lockStrength Lo } // GetPolicyByIDOrPublicID indicates an expected call of GetPolicyByIDOrPublicID. -func (mr *MockStoreMockRecorder) GetPolicyByIDOrPublicID(ctx, lockStrength, accountID, policyID interface{}) *gomock.Call { +func (mr *MockStoreMockRecorder) GetPolicyByIDOrPublicID(ctx, lockStrength, accountID, policyID any) *gomock.Call { mr.mock.ctrl.T.Helper() return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPolicyByIDOrPublicID", reflect.TypeOf((*MockStore)(nil).GetPolicyByIDOrPublicID), ctx, lockStrength, accountID, policyID) } @@ -2717,7 +2671,7 @@ func (m *MockStore) GetRouteByIDOrPublicID(ctx context.Context, lockStrength Loc } // GetRouteByIDOrPublicID indicates an expected call of GetRouteByIDOrPublicID. -func (mr *MockStoreMockRecorder) GetRouteByIDOrPublicID(ctx, lockStrength, accountID, routeID interface{}) *gomock.Call { +func (mr *MockStoreMockRecorder) GetRouteByIDOrPublicID(ctx, lockStrength, accountID, routeID any) *gomock.Call { mr.mock.ctrl.T.Helper() return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetRouteByIDOrPublicID", reflect.TypeOf((*MockStore)(nil).GetRouteByIDOrPublicID), ctx, lockStrength, accountID, routeID) } diff --git a/management/server/types/store.go b/management/server/types/store.go index 2ca4383b2..a13d52f91 100644 --- a/management/server/types/store.go +++ b/management/server/types/store.go @@ -1,10 +1,12 @@ package types -type Engine string +import "github.com/netbirdio/netbird/management/internals/shared/db" + +type Engine = db.Engine const ( - PostgresStoreEngine Engine = "postgres" + PostgresStoreEngine = db.PostgresStoreEngine FileStoreEngine Engine = "jsonfile" - SqliteStoreEngine Engine = "sqlite" - MysqlStoreEngine Engine = "mysql" + SqliteStoreEngine = db.SqliteStoreEngine + MysqlStoreEngine = db.MysqlStoreEngine )