[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
This commit is contained in:
Pascal Fischer
2026-09-29 00:45:55 +02:00
committed by GitHub
parent 002755c412
commit 782c943410
46 changed files with 1247 additions and 562 deletions
@@ -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())
}
@@ -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
@@ -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()
@@ -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)
}
}
@@ -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"))
}
@@ -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)
}
@@ -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)
}
+14 -2
View File
@@ -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,
+121
View File
@@ -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()
}
+148
View File
@@ -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)
}
@@ -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
}
@@ -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())
}
+10
View File
@@ -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"
)
+12
View File
@@ -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"
)
+160
View File
@@ -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()
}
}
@@ -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"))
}
@@ -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())
}
}
@@ -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("")
+76 -321
View File
@@ -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))
})
}
+4 -4
View File
@@ -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,
@@ -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,
@@ -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 < ?)",
@@ -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 {
@@ -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")
+3 -3
View File
@@ -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
}
@@ -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
}
@@ -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)
}
@@ -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
}
+1 -1
View File
@@ -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
}
@@ -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
}
@@ -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
}
+2 -2
View File
@@ -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
}
@@ -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
}
+2 -2
View File
@@ -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)
}
@@ -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
}
@@ -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
}
+1 -1
View File
@@ -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
}
+2 -2
View File
@@ -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
@@ -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
}
@@ -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
}
+87 -1
View File
@@ -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)
})
}
+2 -2
View File
@@ -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
}
@@ -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) {
+49 -35
View File
@@ -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)
}
+3 -49
View File
@@ -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)
}
+6 -4
View File
@@ -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
)