diff --git a/integration_tests/management/network_map_db/pgsql/account_settings_test.go b/integration_tests/management/network_map_db/pgsql/account_settings_test.go index e6de07a5a..d7927aaf1 100644 --- a/integration_tests/management/network_map_db/pgsql/account_settings_test.go +++ b/integration_tests/management/network_map_db/pgsql/account_settings_test.go @@ -7,7 +7,6 @@ import ( "testing" "time" - networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql" "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" "github.com/stretchr/testify/assert" ) @@ -21,7 +20,7 @@ func TestGetAccountSettings(t *testing.T) { settings_lazy_connection_enabled, settings_auto_update_version, settings_auto_update_always, settings_metrics_push_enabled) values('account-3',null,null,null,null,null,null,null,null,null,null,null)`) - accountSettings, err := networkmap_pgsql.GetAccountSettingsViaPgxConnection(ctx, conn(t, ctx), "account-1") + accountSettings, err := conn(t, ctx).GetAccountSettings(ctx, "account-1") assert.NoError(t, err) assert.Equal(t, accountSettings, nmdata.AccountSettingsInfo{ PeerLoginExpirationEnabled: true, @@ -37,7 +36,7 @@ func TestGetAccountSettings(t *testing.T) { MetricsPushEnabled: false, }) - accountSettings, err = networkmap_pgsql.GetAccountSettingsViaPgxConnection(ctx, conn(t, ctx), "account-2") + accountSettings, err = conn(t, ctx).GetAccountSettings(ctx, "account-2") assert.NoError(t, err) assert.Equal(t, accountSettings, nmdata.AccountSettingsInfo{ PeerLoginExpirationEnabled: true, @@ -53,7 +52,7 @@ func TestGetAccountSettings(t *testing.T) { MetricsPushEnabled: false, }) - accountSettings, err = networkmap_pgsql.GetAccountSettingsViaPgxConnection(ctx, conn(t, ctx), "account-3") + accountSettings, err = conn(t, ctx).GetAccountSettings(ctx, "account-3") assert.NoError(t, err) assert.Equal(t, accountSettings, nmdata.AccountSettingsInfo{}) } diff --git a/integration_tests/management/network_map_db/pgsql/dns_settings_test.go b/integration_tests/management/network_map_db/pgsql/dns_settings_test.go index f857dded2..95ac84aed 100644 --- a/integration_tests/management/network_map_db/pgsql/dns_settings_test.go +++ b/integration_tests/management/network_map_db/pgsql/dns_settings_test.go @@ -13,13 +13,13 @@ import ( func TestGetDnsSettings(t *testing.T) { ctx := context.TODO() - settings, err := pgstore.GetDnsSettings(ctx, "account-1") + settings, err := conn(t, ctx).GetDnsSettings(ctx, "account-1") assert.NoError(t, err) assert.Equal(t, settings, nmdata.DNSSettings{ DisabledManagementGroups: []string{"disabled-group-1", "disabled-group-2"}, }) - settings, err = pgstore.GetDnsSettings(ctx, "account-2") + settings, err = conn(t, ctx).GetDnsSettings(ctx, "account-2") assert.NoError(t, err) assert.Equal(t, settings, nmdata.DNSSettings{}) } diff --git a/integration_tests/management/network_map_db/pgsql/dns_test.go b/integration_tests/management/network_map_db/pgsql/dns_test.go index 5a3c37b1e..63a609193 100644 --- a/integration_tests/management/network_map_db/pgsql/dns_test.go +++ b/integration_tests/management/network_map_db/pgsql/dns_test.go @@ -7,7 +7,6 @@ import ( "testing" "github.com/miekg/dns" - networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql" "github.com/netbirdio/netbird/shared/management/networkmap" "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" "github.com/stretchr/testify/assert" @@ -35,7 +34,7 @@ func TestGetAppliedZoneCandidatesViaPgxConnection(t *testing.T) { `insert into records (id, account_id, zone_id, name, type, ttl, content) VALUES('record-4','account-1','zone-2','test2.test-2.com','CNAME',1800,'test3.test-2.com')`) - zoneCandidates, err := networkmap_pgsql.GetAppliedZoneCandidatesViaPgxConnection(ctx, conn(t, ctx), "account-1") + zoneCandidates, err := conn(t, ctx).GetAppliedZoneCandidates(ctx, "account-1") assert.NoError(t, err) assert.Contains(t, zoneCandidates, networkmap.AppliedZoneCandidate{ diff --git a/integration_tests/management/network_map_db/pgsql/domain_test.go b/integration_tests/management/network_map_db/pgsql/domain_test.go index 2ebb16b5d..8434a76c3 100644 --- a/integration_tests/management/network_map_db/pgsql/domain_test.go +++ b/integration_tests/management/network_map_db/pgsql/domain_test.go @@ -7,7 +7,7 @@ import ( "database/sql" "testing" - networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql" + networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db" "github.com/stretchr/testify/assert" ) @@ -24,15 +24,15 @@ func TestGetDomains(t *testing.T) { `insert into domains (id, account_id, domain, target_cluster) VALUES('domain-3','account-1',null,null)`) - domains, err := pgstore.GetDomains(ctx, "account-1") + domains, err := conn(t, ctx).GetDomains(ctx, "account-1") assert.NoError(t, err) assert.Len(t, domains, 2) - assert.Contains(t, domains, networkmap_pgsql.Domain{ + assert.Contains(t, domains, networkmapdb.Domain{ Domain: sql.NullString{String: "test-1.com", Valid: true}, TargetCluster: sql.NullString{String: "target-1.cluster.local", Valid: true}, }) - assert.Contains(t, domains, networkmap_pgsql.Domain{ + assert.Contains(t, domains, networkmapdb.Domain{ Domain: sql.NullString{String: "test-2.com", Valid: true}, TargetCluster: sql.NullString{String: "target-2.cluster.local", Valid: true}, }) diff --git a/integration_tests/management/network_map_db/pgsql/group_test.go b/integration_tests/management/network_map_db/pgsql/group_test.go index 8acf4475d..3ccf96eb0 100644 --- a/integration_tests/management/network_map_db/pgsql/group_test.go +++ b/integration_tests/management/network_map_db/pgsql/group_test.go @@ -6,7 +6,6 @@ import ( "context" "testing" - networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql" "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -15,10 +14,7 @@ import ( func TestGetGroups(t *testing.T) { ctx := context.TODO() - s, err := networkmap_pgsql.NewPostgresqlStore(ctx, dsn) - assert.NoError(t, err) - - groups, resourceToGroupIdx, err := s.GetGroups(ctx, "account-1") + groups, resourceToGroupIdx, err := conn(t, ctx).GetGroups(ctx, "account-1") assert.NoError(t, err) assert.Contains(t, groups, @@ -45,17 +41,13 @@ func TestGetGroups(t *testing.T) { func TestGetGroupsWithoutExpectedFields(t *testing.T) { ctx := context.TODO() - s, err := networkmap_pgsql.NewPostgresqlStore(ctx, dsn) - assert.NoError(t, err) - execQuery(t, ctx, "insert into accounts (id) VALUES('random-id')") execQuery(t, ctx, "insert into groups (id, account_id) VALUES('g2-test-group-id-1','random-id')") - assert.NoError(t, err) - groups, _, err := s.GetGroups(ctx, "random-id") + groups, _, err := conn(t, ctx).GetGroups(ctx, "random-id") assert.NoError(t, err) require.Len(t, groups, 1) assert.NotEmpty(t, groups[0].PublicID) diff --git a/integration_tests/management/network_map_db/pgsql/main_test.go b/integration_tests/management/network_map_db/pgsql/main_test.go index 90760c57c..a8d482654 100644 --- a/integration_tests/management/network_map_db/pgsql/main_test.go +++ b/integration_tests/management/network_map_db/pgsql/main_test.go @@ -5,149 +5,51 @@ package networkmap_pgsql import ( "context" _ "embed" - "fmt" "os" - "regexp" - "strings" "testing" "time" - "github.com/google/uuid" - "github.com/jackc/pgx/v5" + log "github.com/sirupsen/logrus" "github.com/stretchr/testify/assert" - log "github.com/sirupsen/logrus" - "gorm.io/driver/postgres" - "gorm.io/gorm" - + networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db" networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql" - gormstore "github.com/netbirdio/netbird/management/server/store" - "github.com/netbirdio/netbird/management/server/testutil" + "github.com/netbirdio/netbird/management/server/types" ) //go:embed base_data.sql var baseData string var ( - dsn string pgstore *networkmap_pgsql.PgStore + engine string ) func TestMain(m *testing.M) { - _, tmpdsn, err := testutil.CreatePostgresTestContainer() - if err != nil { - log.Fatalf("error starting postres container %v", err) + kind, _ := os.LookupEnv("NETBIRD_STORE_ENGINE") + switch kind { + case "": + engine = string(types.PostgresStoreEngine) + case string(types.PostgresStoreEngine), string(types.SqliteStoreEngine): + engine = kind + default: + log.Fatalf("unsupported db '%s' in NETBIRD_STORE_ENGINE env var", kind) } - var db *gorm.DB - for i := range 5 { - db, err = gorm.Open(postgres.Open(tmpdsn), &gorm.Config{}) - - if err == nil { - break - } - - if i < 5 { - waitTime := time.Duration(100*(i+1)) * time.Millisecond - time.Sleep(waitTime) - continue - } - - log.Fatalf("error connecting to postres db %v", err) - } - - var cleanup func() - dsn, cleanup, err = createRandomDB(tmpdsn, db) - sqlDB, _ := db.DB() - if sqlDB != nil { - sqlDB.Close() - } - if err != nil { - log.Fatalf("error creating postres db %v", err) - } - - _, err = gormstore.NewPostgresqlStoreForTests(context.TODO(), dsn, nil, false) - if err != nil { - log.Fatalf("error running migrations %v", err) - } - - ctx := context.TODO() - s, err := networkmap_pgsql.NewPostgresqlStore(ctx, dsn) - if err != nil { - log.Fatal("error creating postgres store %w", err) - } - - for _, query := range strings.Split(baseData, ";") { - if _, err := s.Pool.Exec(ctx, query); err != nil { - log.Fatalf("error initializing db: %s", err.Error()) - } - } - - pgstore, err = networkmap_pgsql.NewPostgresqlStore(ctx, dsn) - if err != nil { - log.Fatalf("error creating pg store %v", err.Error()) - } + store, cleanup := createPGTestStore(baseData) + pgstore = store code := m.Run() cleanup() - os.Exit(code) } -func createRandomDB(dsn string, db *gorm.DB) (string, func(), error) { - dbName := fmt.Sprintf("test_db_%s", strings.ReplaceAll(uuid.New().String(), "-", "_")) - - if err := db.Exec(fmt.Sprintf("CREATE DATABASE %s", dbName)).Error; err != nil { - return "", nil, fmt.Errorf("failed to create database: %v", err) - } - - originalDSN := dsn - - cleanup := func() { - var dropDB *gorm.DB - var err error - - dropDB, err = gorm.Open(postgres.Open(originalDSN), &gorm.Config{ - SkipDefaultTransaction: true, - PrepareStmt: false, - }) - if err != nil { - log.Errorf("failed to connect for dropping database %s: %v", dbName, err) - return - } - defer func() { - if sqlDB, _ := dropDB.DB(); sqlDB != nil { - sqlDB.Close() - } - }() - - if sqlDB, _ := dropDB.DB(); sqlDB != nil { - sqlDB.SetMaxOpenConns(1) - sqlDB.SetMaxIdleConns(0) - sqlDB.SetConnMaxLifetime(time.Second) - } - - err = dropDB.Exec(fmt.Sprintf("DROP DATABASE IF EXISTS %s WITH (FORCE)", dbName)).Error - - if err != nil { - log.Errorf("failed to drop database %s: %v", dbName, err) - } - } - - return replaceDBName(dsn, dbName), cleanup, nil -} - -func replaceDBName(dsn, newDBName string) string { - re := regexp.MustCompile(`(?P
[:/@])(?P[^/?]+)(?P \?|$)`) - return re.ReplaceAllString(dsn, `${pre}`+newDBName+`${post}`) -} - -func conn(t *testing.T, ctx context.Context) *pgx.Conn { +func conn(t *testing.T, ctx context.Context) networkmapdb.NetworkMapDBStoreConn { t.Helper() c, err := pgstore.Pool.Acquire(ctx) assert.NoError(t, err) - return c.Conn() + return pgstore.UsingConnection(c.Conn()) } func execQuery(t *testing.T, ctx context.Context, q string) { diff --git a/integration_tests/management/network_map_db/pgsql/nameserver_test.go b/integration_tests/management/network_map_db/pgsql/nameserver_test.go index b352cfe96..d6243a6e3 100644 --- a/integration_tests/management/network_map_db/pgsql/nameserver_test.go +++ b/integration_tests/management/network_map_db/pgsql/nameserver_test.go @@ -7,7 +7,6 @@ import ( "net/netip" "testing" - networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql" "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" "github.com/stretchr/testify/assert" ) @@ -25,7 +24,7 @@ func TestGetNameServerGroups(t *testing.T) { `insert into name_server_groups (id, public_id, name, description, name_servers, groups, domains, enabled, search_domains_enabled,"primary",account_id) VALUES('nsgroup-3','nsgroup-3-public',null,null,null,null,null,TRUE,FALSE,FALSE,'account-1')`) - nsgroups, err := networkmap_pgsql.GetNameServerGroupsViaPgxConnection(ctx, conn(t, ctx), "account-1") + nsgroups, err := conn(t, ctx).GetNameServerGroups(ctx, "account-1") assert.NoError(t, err) assert.Contains(t, nsgroups, nmdata.NameServerGroup{ diff --git a/integration_tests/management/network_map_db/pgsql/network_resource_test.go b/integration_tests/management/network_map_db/pgsql/network_resource_test.go index cf300badb..4325ed3ba 100644 --- a/integration_tests/management/network_map_db/pgsql/network_resource_test.go +++ b/integration_tests/management/network_map_db/pgsql/network_resource_test.go @@ -7,7 +7,6 @@ import ( "net/netip" "testing" - networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql" "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" "github.com/stretchr/testify/assert" ) @@ -15,9 +14,6 @@ import ( func TestGetNetworkResources(t *testing.T) { ctx := context.TODO() - s, err := networkmap_pgsql.NewPostgresqlStore(ctx, dsn) - assert.NoError(t, err) - execQuery(t, ctx, `insert into network_resources (id, account_id, network_id, public_id, name, description, type, domain, prefix, enabled) VALUES('net-resource-1','account-1','network-1','net-resource-public-1','network-resource-1','network-resource-1','subnet','','"10.0.0.0/16"',TRUE)`) @@ -28,7 +24,7 @@ func TestGetNetworkResources(t *testing.T) { `insert into network_resources (id, account_id, network_id, public_id, name, description, type, domain, prefix, enabled) VALUES('net-resource-3','account-1','network-3','net-resource-public-3','network-resource-3','network-resource-3','host','','"10.0.0.1/32"',TRUE)`) - resources, err := s.GetNetworkResources(ctx, "account-1") + resources, err := conn(t, ctx).GetNetworkResources(ctx, "account-1") assert.NoError(t, err) assert.Contains(t, resources, nmdata.NetworkResource{ diff --git a/integration_tests/management/network_map_db/pgsql/network_router_test.go b/integration_tests/management/network_map_db/pgsql/network_router_test.go index 4eab2d60d..fa7ea2a04 100644 --- a/integration_tests/management/network_map_db/pgsql/network_router_test.go +++ b/integration_tests/management/network_map_db/pgsql/network_router_test.go @@ -6,7 +6,6 @@ import ( "context" "testing" - networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql" "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" "github.com/stretchr/testify/assert" ) @@ -14,9 +13,6 @@ import ( func TestGetNetworkRouters(t *testing.T) { ctx := context.TODO() - s, err := networkmap_pgsql.NewPostgresqlStore(ctx, dsn) - assert.NoError(t, err) - execQuery(t, ctx, `insert into network_routers (id, account_id, public_id, peer, network_id, masquerade, metric, enabled, peer_groups) VALUES('test-nr-id-1','account-1','public-id-1','peer-id-1','network-id-1',TRUE,999,TRUE,'["group-one-resource-id"]')`) @@ -24,7 +20,7 @@ func TestGetNetworkRouters(t *testing.T) { `insert into network_routers (id, account_id, public_id, peer, network_id, masquerade, metric, enabled, peer_groups) VALUES('test-nr-id-2','account-1','public-id-2','','network-id-2',TRUE,333,TRUE,'["group-two-resources-id","group-no-resources-id"]')`) - routers, err := s.GetNetworkRouters(ctx, "account-1") + routers, err := conn(t, ctx).GetNetworkRouters(ctx, "account-1") assert.NoError(t, err) assert.NotEmpty(t, routers) diff --git a/integration_tests/management/network_map_db/pgsql/network_test.go b/integration_tests/management/network_map_db/pgsql/network_test.go index a43449f71..fbccee504 100644 --- a/integration_tests/management/network_map_db/pgsql/network_test.go +++ b/integration_tests/management/network_map_db/pgsql/network_test.go @@ -8,7 +8,6 @@ import ( "net" "testing" - networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql" "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" "github.com/stretchr/testify/assert" ) @@ -16,10 +15,7 @@ import ( func TestGetNetwork(t *testing.T) { ctx := context.TODO() - s, err := networkmap_pgsql.NewPostgresqlStore(ctx, dsn) - assert.NoError(t, err) - - network, err := s.GetNetwork(ctx, "account-1") + network, err := conn(t, ctx).GetNetwork(ctx, "account-1") assert.NoError(t, err) assert.Equal(t, network, nmdata.Network{ Identifier: "network-1", @@ -28,7 +24,7 @@ func TestGetNetwork(t *testing.T) { Serial: 1, }) - network, err = s.GetNetwork(ctx, "account-2") + network, err = conn(t, ctx).GetNetwork(ctx, "account-2") assert.NoError(t, err) assert.Equal(t, network, nmdata.Network{ Identifier: "network-2", diff --git a/integration_tests/management/network_map_db/pgsql/networks_test.go b/integration_tests/management/network_map_db/pgsql/networks_test.go index e4468c9d3..5af771522 100644 --- a/integration_tests/management/network_map_db/pgsql/networks_test.go +++ b/integration_tests/management/network_map_db/pgsql/networks_test.go @@ -6,7 +6,6 @@ import ( "context" "testing" - networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql" "github.com/stretchr/testify/assert" ) @@ -18,7 +17,7 @@ func TestGetNetworks(t *testing.T) { execQuery(t, ctx, `insert into networks (id, account_id, public_id) VALUES('network-2','account-1','network-2-public')`) - networksIdx, err := networkmap_pgsql.GetNetworkXIDToPublicIdMapViaPgxConnection(ctx, conn(t, ctx), "account-1") + networksIdx, err := conn(t, ctx).GetNetworkXIDToPublicIdMap(ctx, "account-1") assert.NoError(t, err) assert.Equal(t, networksIdx, map[string]string{ "network-1": "network-1-public", diff --git a/integration_tests/management/network_map_db/pgsql/peer_test.go b/integration_tests/management/network_map_db/pgsql/peer_test.go index 0d965aba5..af1dffdbd 100644 --- a/integration_tests/management/network_map_db/pgsql/peer_test.go +++ b/integration_tests/management/network_map_db/pgsql/peer_test.go @@ -8,7 +8,6 @@ import ( "net/netip" "testing" - networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql" "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" "github.com/stretchr/testify/assert" ) @@ -16,7 +15,7 @@ import ( func TestGetPeers(t *testing.T) { ctx := context.TODO() - peers, clusterToPeersIdx, err := networkmap_pgsql.GetPeersViaPgxConnection(ctx, conn(t, ctx), "account-1") + peers, clusterToPeersIdx, err := conn(t, ctx).GetPeers(ctx, "account-1") assert.NoError(t, err) // shouldn't be returned in the index, as it's not connected diff --git a/integration_tests/management/network_map_db/pgsql/pg_test_store.go b/integration_tests/management/network_map_db/pgsql/pg_test_store.go new file mode 100644 index 000000000..2c7afacc5 --- /dev/null +++ b/integration_tests/management/network_map_db/pgsql/pg_test_store.go @@ -0,0 +1,126 @@ +//go:build integration + +package networkmap_pgsql + +import ( + "context" + "fmt" + "regexp" + "strings" + "time" + + log "github.com/sirupsen/logrus" + + "github.com/google/uuid" + networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql" + gormstore "github.com/netbirdio/netbird/management/server/store" + "github.com/netbirdio/netbird/management/server/testutil" + "gorm.io/driver/postgres" + "gorm.io/gorm" +) + +func createPGTestStore(baseData string) (*networkmap_pgsql.PgStore, func()) { + _, tmpdsn, err := testutil.CreatePostgresTestContainer() + if err != nil { + log.Fatalf("error starting postres container %v", err) + } + + var db *gorm.DB + for i := range 5 { + db, err = gorm.Open(postgres.Open(tmpdsn), &gorm.Config{}) + + if err == nil { + break + } + + if i < 5 { + waitTime := time.Duration(100*(i+1)) * time.Millisecond + time.Sleep(waitTime) + continue + } + + log.Fatalf("error connecting to postres db %v", err) + } + + var cleanup func() + dsn, cleanup, err := createRandomDB(tmpdsn, db) + sqlDB, _ := db.DB() + if sqlDB != nil { + sqlDB.Close() + } + if err != nil { + log.Fatalf("error creating postres db %v", err) + } + + _, err = gormstore.NewPostgresqlStoreForTests(context.TODO(), dsn, nil, false) + if err != nil { + log.Fatalf("error running migrations %v", err) + } + + ctx := context.TODO() + s, err := networkmap_pgsql.NewPostgresqlStore(ctx, dsn) + if err != nil { + log.Fatal("error creating postgres store %w", err) + } + + for _, query := range strings.Split(baseData, ";") { + if _, err := s.Pool.Exec(ctx, query); err != nil { + log.Fatalf("error initializing db: %s", err.Error()) + } + } + + pgstore, err := networkmap_pgsql.NewPostgresqlStore(ctx, dsn) + if err != nil { + log.Fatalf("error creating pg store %v", err.Error()) + } + + return pgstore, cleanup +} + +func createRandomDB(dsn string, db *gorm.DB) (string, func(), error) { + dbName := fmt.Sprintf("test_db_%s", strings.ReplaceAll(uuid.New().String(), "-", "_")) + + if err := db.Exec(fmt.Sprintf("CREATE DATABASE %s", dbName)).Error; err != nil { + return "", nil, fmt.Errorf("failed to create database: %v", err) + } + + originalDSN := dsn + + cleanup := func() { + var dropDB *gorm.DB + var err error + + dropDB, err = gorm.Open(postgres.Open(originalDSN), &gorm.Config{ + SkipDefaultTransaction: true, + PrepareStmt: false, + }) + if err != nil { + log.Errorf("failed to connect for dropping database %s: %v", dbName, err) + return + } + defer func() { + if sqlDB, _ := dropDB.DB(); sqlDB != nil { + sqlDB.Close() + } + }() + + if sqlDB, _ := dropDB.DB(); sqlDB != nil { + sqlDB.SetMaxOpenConns(1) + sqlDB.SetMaxIdleConns(0) + sqlDB.SetConnMaxLifetime(time.Second) + } + + err = dropDB.Exec(fmt.Sprintf("DROP DATABASE IF EXISTS %s WITH (FORCE)", dbName)).Error + + if err != nil { + log.Errorf("failed to drop database %s: %v", dbName, err) + } + } + + return replaceDBName(dsn, dbName), cleanup, nil +} + +func replaceDBName(dsn, newDBName string) string { + re := regexp.MustCompile(`(?P [:/@])(?P[^/?]+)(?P \?|$)`) + return re.ReplaceAllString(dsn, `${pre}`+newDBName+`${post}`) +} diff --git a/integration_tests/management/network_map_db/pgsql/policy_test.go b/integration_tests/management/network_map_db/pgsql/policy_test.go index 47b5836e2..a3f3db794 100644 --- a/integration_tests/management/network_map_db/pgsql/policy_test.go +++ b/integration_tests/management/network_map_db/pgsql/policy_test.go @@ -6,7 +6,6 @@ import ( "context" "testing" - networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql" "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" "github.com/stretchr/testify/assert" ) @@ -54,7 +53,7 @@ func TestGetPolicies(t *testing.T) { values('policy-4-rule-1','policy-4',false,null,null,null,null,'["group-two-resources-id"]', null,'{"ID":"domain-3","Type":"domain"}',null,null,null,null)`) - policies, policyToDestinationResourceIdx, policyToDestinationGroupIdx, err := networkmap_pgsql.GetPoliciesViaPgxConnection(ctx, conn(t, ctx), "account-1") + policies, policyToDestinationResourceIdx, policyToDestinationGroupIdx, err := conn(t, ctx).GetPolicies(ctx, "account-1") assert.NoError(t, err) assert.Contains(t, policies, nmdata.Policy{ diff --git a/integration_tests/management/network_map_db/pgsql/posture_test.go b/integration_tests/management/network_map_db/pgsql/posture_test.go index 37812f35d..2b4bb3f3d 100644 --- a/integration_tests/management/network_map_db/pgsql/posture_test.go +++ b/integration_tests/management/network_map_db/pgsql/posture_test.go @@ -7,7 +7,6 @@ import ( "net/netip" "testing" - networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql" "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" "github.com/stretchr/testify/assert" ) @@ -34,7 +33,7 @@ func TestGetPostureChecks(t *testing.T) { `insert into posture_checks (id, account_id, public_id, checks) VALUES('posturecheck-3','account-1','posturecheck-3-public', null)`) - postureChecks, idToPublicIDIdx, err := networkmap_pgsql.GetPostureChecksViaPgxConnection(ctx, conn(t, ctx), "account-1") + postureChecks, idToPublicIDIdx, err := conn(t, ctx).GetPostureChecks(ctx, "account-1") assert.NoError(t, err) assert.Equal(t, idToPublicIDIdx, map[string]string{ "posturecheck-1": "posturecheck-1-public", diff --git a/integration_tests/management/network_map_db/pgsql/route_test.go b/integration_tests/management/network_map_db/pgsql/route_test.go index 3b9a7822f..12e9302f9 100644 --- a/integration_tests/management/network_map_db/pgsql/route_test.go +++ b/integration_tests/management/network_map_db/pgsql/route_test.go @@ -7,7 +7,6 @@ import ( "net/netip" "testing" - networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql" "github.com/netbirdio/netbird/shared/management/domain" "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" "github.com/stretchr/testify/assert" @@ -37,7 +36,7 @@ func TestGetRoutes(t *testing.T) { VALUES('route-3','account-1','route-3-public',null,null,null,null,'route-3', null,null,null,null,null,null,null,null,null)`) - routes, err := networkmap_pgsql.GetRoutesViaPgxConnection(ctx, conn(t, ctx), "account-1") + routes, err := conn(t, ctx).GetRoutes(ctx, "account-1") assert.NoError(t, err) assert.Contains(t, routes, nmdata.Route{ ID: "route-1", diff --git a/integration_tests/management/network_map_db/pgsql/service_test.go b/integration_tests/management/network_map_db/pgsql/service_test.go index de8028b24..8d8facaa0 100644 --- a/integration_tests/management/network_map_db/pgsql/service_test.go +++ b/integration_tests/management/network_map_db/pgsql/service_test.go @@ -9,7 +9,7 @@ import ( "github.com/stretchr/testify/assert" - networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql" + networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db" ) func TestGetPrivateServicesViaPgxConnection(t *testing.T) { @@ -28,23 +28,23 @@ func TestGetPrivateServicesViaPgxConnection(t *testing.T) { values('service-3','account-1',null,null,null,null,null)`) assert.NoError(t, err) - services, err := networkmap_pgsql.GetPrivateServicesViaPgxConnection(ctx, conn(t, ctx), "account-1") + services, err := conn(t, ctx).GetPrivateServices(ctx, "account-1") assert.NoError(t, err) - assert.Contains(t, services, networkmap_pgsql.Service{ + assert.Contains(t, services, networkmapdb.Service{ Enabled: sql.NullBool{Bool: true, Valid: true}, Private: sql.NullBool{Bool: true, Valid: true}, AccessGroups: []string{"group-one-resource-id"}, ProxyCluster: sql.NullString{String: "test-1.com", Valid: true}, Domain: sql.NullString{String: "test-2.com", Valid: true}, }) - assert.Contains(t, services, networkmap_pgsql.Service{ + assert.Contains(t, services, networkmapdb.Service{ Enabled: sql.NullBool{Bool: true, Valid: true}, Private: sql.NullBool{Bool: true, Valid: true}, AccessGroups: []string{"group-one-resource-id", "group-two-resources-id"}, ProxyCluster: sql.NullString{String: "test-3.com", Valid: true}, Domain: sql.NullString{String: "test-4.com", Valid: true}, }) - assert.Contains(t, services, networkmap_pgsql.Service{ + assert.Contains(t, services, networkmapdb.Service{ Enabled: sql.NullBool{Bool: false, Valid: false}, Private: sql.NullBool{Bool: false, Valid: false}, AccessGroups: []string{}, @@ -102,7 +102,7 @@ func TestGetProxyTargetedDomainResourceIDsViaPgxConnection(t *testing.T) { `insert into targets (target_id, account_id, service_id, enabled, target_type) values(null,'account-1','service-4',true,'cluster')`) - servtargetedDomains, err := networkmap_pgsql.GetProxyTargetedDomainResourceIDsViaPgxConnection(ctx, conn(t, ctx), "account-1") + servtargetedDomains, err := conn(t, ctx).GetProxyTargetedDomainResourceIDs(ctx, "account-1") assert.NoError(t, err) assert.Equal(t, servtargetedDomains, map[string]struct{}{ "target-1": {}, diff --git a/integration_tests/management/network_map_db/pgsql/user_test.go b/integration_tests/management/network_map_db/pgsql/user_test.go index f70f7a65e..132f749e2 100644 --- a/integration_tests/management/network_map_db/pgsql/user_test.go +++ b/integration_tests/management/network_map_db/pgsql/user_test.go @@ -6,7 +6,6 @@ import ( "context" "testing" - networkmap_pgsql "github.com/netbirdio/netbird/management/internals/network_map_db/pgsql" "github.com/stretchr/testify/assert" ) @@ -40,7 +39,7 @@ func TestGetAllowedUsers(t *testing.T) { `insert into groups (id, name, account_id) VALUES('all-group-3','All','account-1')`) - userIdx, groupIdToUserIds, err := networkmap_pgsql.GetAllowedUsersViaPgxConnection(ctx, conn(t, ctx), "account-1") + userIdx, groupIdToUserIds, err := conn(t, ctx).GetAllowedUsers(ctx, "account-1") assert.NoError(t, err) assert.Equal(t, userIdx, map[string]struct{}{ diff --git a/management/internals/network_map_db/db_store.go b/management/internals/network_map_db/db_store.go index 634b5725f..fe31c5a5f 100644 --- a/management/internals/network_map_db/db_store.go +++ b/management/internals/network_map_db/db_store.go @@ -24,7 +24,12 @@ const ( ) type NetworkMapDBStore interface { //nolint:revive // established name across the codebase + GetNetworkMapData(ctx context.Context, accountId string) (*networkmap.NetworkMapData, error) +} + +type NetworkMapDBStoreConn interface { //nolint:revive // established name across the codebase GetGroups(ctx context.Context, accountId string) ([]nmdata.Group, map[string]map[string]any, error) + GetDomains(ctx context.Context, accountId string) ([]Domain, error) GetPeers(ctx context.Context, accountId string) ([]nmdata.Peer, map[string][]*nmdata.Peer, error) GetPolicies(ctx context.Context, accountId string) ([]nmdata.Policy, map[string]map[string]any, map[string]map[string]any, error) GetRoutes(ctx context.Context, accountId string) ([]nmdata.Route, error) @@ -35,9 +40,24 @@ type NetworkMapDBStore interface { //nolint:revive // established name across th GetAppliedZoneCandidates(ctx context.Context, accountId string) ([]networkmap.AppliedZoneCandidate, error) GetAccountSettings(ctx context.Context, accountId string) (nmdata.AccountSettingsInfo, error) GetPostureChecks(ctx context.Context, accountId string) ([]nmdata.PostureChecks, map[string]string, error) - GetNetworkMapData(ctx context.Context, accountId string) (*networkmap.NetworkMapData, error) GetAllowedUsers(ctx context.Context, accountId string) (map[string]struct{}, map[string][]string, error) GetDnsSettings(ctx context.Context, accountId string) (nmdata.DNSSettings, error) + GetNetworkXIDToPublicIdMap(ctx context.Context, accountId string) (map[string]string, error) + GetPrivateServices(ctx context.Context, accountId string) ([]Service, error) + GetProxyTargetedDomainResourceIDs(ctx context.Context, accountId string) (map[string]struct{}, error) +} + +type Domain struct { + Domain sql.NullString + TargetCluster sql.NullString +} + +type Service struct { + Enabled sql.NullBool + Private sql.NullBool + AccessGroups []string + ProxyCluster sql.NullString + Domain sql.NullString } type NetworkMapDBStoreImpl struct { //nolint:revive // established name across the codebase diff --git a/management/internals/network_map_db/pgsql/account_settings.go b/management/internals/network_map_db/pgsql/account_settings.go index 92364dc1e..6fbbc8ef8 100644 --- a/management/internals/network_map_db/pgsql/account_settings.go +++ b/management/internals/network_map_db/pgsql/account_settings.go @@ -28,17 +28,8 @@ const ( ` ) -func (pg *PgStore) GetAccountSettings(ctx context.Context, accountId string) (nmdata.AccountSettingsInfo, error) { - c, err := pg.Pool.Acquire(ctx) - if err != nil { - return nmdata.AccountSettingsInfo{}, err - } - return GetAccountSettingsViaPgxConnection(ctx, c.Conn(), accountId) - -} - -func GetAccountSettingsViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) (nmdata.AccountSettingsInfo, error) { - rows, err := con.Query(ctx, GetAccountSettingsQuery, accountId) +func (pgc *PgStoreConn) GetAccountSettings(ctx context.Context, accountId string) (nmdata.AccountSettingsInfo, error) { + rows, err := pgc.Conn.Query(ctx, GetAccountSettingsQuery, accountId) if err != nil { return nmdata.AccountSettingsInfo{}, err } diff --git a/management/internals/network_map_db/pgsql/dns.go b/management/internals/network_map_db/pgsql/dns.go index 4f349e1f0..b45a2acd9 100644 --- a/management/internals/network_map_db/pgsql/dns.go +++ b/management/internals/network_map_db/pgsql/dns.go @@ -27,16 +27,8 @@ const ( ` ) -func (pg *PgStore) GetAppliedZoneCandidates(ctx context.Context, accountId string) ([]networkmap.AppliedZoneCandidate, error) { - c, err := pg.Pool.Acquire(ctx) - if err != nil { - return nil, err - } - return GetAppliedZoneCandidatesViaPgxConnection(ctx, c.Conn(), accountId) -} - -func GetAppliedZoneCandidatesViaPgxConnection(ctx context.Context, conn *pgx.Conn, accountId string) ([]networkmap.AppliedZoneCandidate, error) { - rows, err := conn.Query(ctx, GetAccountZonesQuery, accountId) +func (pgc *PgStoreConn) GetAppliedZoneCandidates(ctx context.Context, accountId string) ([]networkmap.AppliedZoneCandidate, error) { + rows, err := pgc.Conn.Query(ctx, GetAccountZonesQuery, accountId) if err != nil { return nil, err } diff --git a/management/internals/network_map_db/pgsql/dns_settings.go b/management/internals/network_map_db/pgsql/dns_settings.go index 7dc921827..aec44e0f2 100644 --- a/management/internals/network_map_db/pgsql/dns_settings.go +++ b/management/internals/network_map_db/pgsql/dns_settings.go @@ -16,16 +16,8 @@ const ( ` ) -func (pg *PgStore) GetDnsSettings(ctx context.Context, accountId string) (nmdata.DNSSettings, error) { - c, err := pg.Pool.Acquire(ctx) - if err != nil { - return nmdata.DNSSettings{}, err - } - return GetDnsSettingsViaPgxConnection(ctx, c.Conn(), accountId) -} - -func GetDnsSettingsViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) (nmdata.DNSSettings, error) { - rows, err := con.Query(ctx, GetDnsSettingsQuery, accountId) +func (pgc *PgStoreConn) GetDnsSettings(ctx context.Context, accountId string) (nmdata.DNSSettings, error) { + rows, err := pgc.Conn.Query(ctx, GetDnsSettingsQuery, accountId) if err != nil { return nmdata.DNSSettings{}, err } diff --git a/management/internals/network_map_db/pgsql/domain.go b/management/internals/network_map_db/pgsql/domain.go index aca2bb70f..8730007c5 100644 --- a/management/internals/network_map_db/pgsql/domain.go +++ b/management/internals/network_map_db/pgsql/domain.go @@ -2,9 +2,9 @@ package networkmap_pgsql import ( "context" - "database/sql" "github.com/jackc/pgx/v5" + networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db" ) const ( @@ -15,24 +15,11 @@ const ( ` ) -func (pg *PgStore) GetDomains(ctx context.Context, accountId string) ([]Domain, error) { - c, err := pg.Pool.Acquire(ctx) - if err != nil { - return nil, err - } - return GetDomainsViaPgxConnection(ctx, c.Conn(), accountId) -} - -func GetDomainsViaPgxConnection(ctx context.Context, conn *pgx.Conn, accountId string) ([]Domain, error) { - rows, err := conn.Query(ctx, GetDomainsQuery, accountId) +func (pgc *PgStoreConn) GetDomains(ctx context.Context, accountId string) ([]networkmapdb.Domain, error) { + rows, err := pgc.Conn.Query(ctx, GetDomainsQuery, accountId) if err != nil { return nil, err } - return pgx.CollectRows(rows, pgx.RowToStructByName[Domain]) -} - -type Domain struct { - Domain sql.NullString - TargetCluster sql.NullString + return pgx.CollectRows(rows, pgx.RowToStructByName[networkmapdb.Domain]) } diff --git a/management/internals/network_map_db/pgsql/group.go b/management/internals/network_map_db/pgsql/group.go index 957874f85..874e743a5 100644 --- a/management/internals/network_map_db/pgsql/group.go +++ b/management/internals/network_map_db/pgsql/group.go @@ -26,16 +26,8 @@ const ( // we also return a resource-to-group index. // an alternative is to add json indexes, query this directly. Not sure how expensive // json indexes are. TODO (dmitri) verify and maybe change the implementation here. -func (pg *PgStore) GetGroups(ctx context.Context, accountId string) ([]nmdata.Group, map[string]map[string]any, error) { - c, err := pg.Pool.Acquire(ctx) - if err != nil { - return nil, nil, err - } - return GetGroupsViaPgxConnection(ctx, c.Conn(), accountId) -} - -func GetGroupsViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.Group, map[string]map[string]any, error) { - rows, err := con.Query(ctx, GetGroupsQuery, accountId) +func (pgc *PgStoreConn) GetGroups(ctx context.Context, accountId string) ([]nmdata.Group, map[string]map[string]any, error) { + rows, err := pgc.Conn.Query(ctx, GetGroupsQuery, accountId) if err != nil { return nil, nil, err } diff --git a/management/internals/network_map_db/pgsql/nameserver.go b/management/internals/network_map_db/pgsql/nameserver.go index 718285153..4fed409b0 100644 --- a/management/internals/network_map_db/pgsql/nameserver.go +++ b/management/internals/network_map_db/pgsql/nameserver.go @@ -19,16 +19,8 @@ const ( ` ) -func (pg *PgStore) GetNameServerGroups(ctx context.Context, accountId string) ([]nmdata.NameServerGroup, error) { - c, err := pg.Pool.Acquire(ctx) - if err != nil { - return nil, err - } - return GetNameServerGroupsViaPgxConnection(ctx, c.Conn(), accountId) -} - -func GetNameServerGroupsViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.NameServerGroup, error) { - rows, err := con.Query(ctx, GetNameserversQuery, accountId) +func (pgc *PgStoreConn) GetNameServerGroups(ctx context.Context, accountId string) ([]nmdata.NameServerGroup, error) { + rows, err := pgc.Conn.Query(ctx, GetNameserversQuery, accountId) if err != nil { return nil, err } diff --git a/management/internals/network_map_db/pgsql/network.go b/management/internals/network_map_db/pgsql/network.go index beaa32a64..0b252b4e1 100644 --- a/management/internals/network_map_db/pgsql/network.go +++ b/management/internals/network_map_db/pgsql/network.go @@ -19,16 +19,8 @@ const ( ` ) -func (pg *PgStore) GetNetwork(ctx context.Context, accountId string) (nmdata.Network, error) { - c, err := pg.Pool.Acquire(ctx) - if err != nil { - return nmdata.Network{}, err - } - return GetNetworkViaPgxConnection(ctx, c.Conn(), accountId) -} - -func GetNetworkViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) (nmdata.Network, error) { - rows, err := con.Query(ctx, GetNetworkQuery, accountId) +func (pgc *PgStoreConn) GetNetwork(ctx context.Context, accountId string) (nmdata.Network, error) { + rows, err := pgc.Conn.Query(ctx, GetNetworkQuery, accountId) if err != nil { return nmdata.Network{}, err } diff --git a/management/internals/network_map_db/pgsql/network_map_data.go b/management/internals/network_map_db/pgsql/network_map_data.go index a949c01c4..d44a79a53 100644 --- a/management/internals/network_map_db/pgsql/network_map_data.go +++ b/management/internals/network_map_db/pgsql/network_map_data.go @@ -9,6 +9,7 @@ import ( "github.com/miekg/dns" log "github.com/sirupsen/logrus" + networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db" "github.com/netbirdio/netbird/shared/management/networkmap" "github.com/netbirdio/netbird/shared/management/networkmap/nmdata" ) @@ -19,71 +20,73 @@ func (pg *PgStore) GetNetworkMapData(ctx context.Context, accountId string) (*ne return nil, err } - acctSettings, err := GetAccountSettingsViaPgxConnection(ctx, tx.Conn(), accountId) + conn := pg.UsingConnection(tx.Conn()) + + acctSettings, err := conn.GetAccountSettings(ctx, accountId) if err != nil { return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get account settings: %w", err)) } - dnsZones, err := GetAppliedZoneCandidatesViaPgxConnection(ctx, tx.Conn(), accountId) + dnsZones, err := conn.GetAppliedZoneCandidates(ctx, accountId) if err != nil { return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get applied zone candidates: %w", err)) } - groups, resourceToGroupIdx, err := GetGroupsViaPgxConnection(ctx, tx.Conn(), accountId) + groups, resourceToGroupIdx, err := conn.GetGroups(ctx, accountId) if err != nil { return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get groups: %w", err)) } - nsGroups, err := GetNameServerGroupsViaPgxConnection(ctx, tx.Conn(), accountId) + nsGroups, err := conn.GetNameServerGroups(ctx, accountId) if err != nil { return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get nameserver groups: %w", err)) } - networkResources, err := GetNetworkResourcesViaPgxConnection(ctx, tx.Conn(), accountId) + networkResources, err := conn.GetNetworkResources(ctx, accountId) if err != nil { return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get network resources: %w", err)) } - routers, err := GetNetworkRoutersViaPgxConnection(ctx, tx.Conn(), accountId) + routers, err := conn.GetNetworkRouters(ctx, accountId) if err != nil { return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get network routers: %w", err)) } - network, err := GetNetworkViaPgxConnection(ctx, tx.Conn(), accountId) + network, err := conn.GetNetwork(ctx, accountId) if err != nil { return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get network: %w", err)) } - peers, proxyPeers, err := GetPeersViaPgxConnection(ctx, tx.Conn(), accountId) + peers, proxyPeers, err := conn.GetPeers(ctx, accountId) if err != nil { return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get peers: %w", err)) } - policies, policyToDestinationResourceIdx, policyToDestinationGroupIdx, err := GetPoliciesViaPgxConnection(ctx, tx.Conn(), accountId) + policies, policyToDestinationResourceIdx, policyToDestinationGroupIdx, err := conn.GetPolicies(ctx, accountId) if err != nil { return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get policies: %w", err)) } - postureChecks, postureCheckXIDToPublicID, err := GetPostureChecksViaPgxConnection(ctx, tx.Conn(), accountId) + postureChecks, postureCheckXIDToPublicID, err := conn.GetPostureChecks(ctx, accountId) if err != nil { return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get posture checks: %w", err)) } - routes, err := GetRoutesViaPgxConnection(ctx, tx.Conn(), accountId) + routes, err := conn.GetRoutes(ctx, accountId) if err != nil { return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get routes: %w", err)) } - networkXIDToPublicID, err := GetNetworkXIDToPublicIdMapViaPgxConnection(ctx, tx.Conn(), accountId) + networkXIDToPublicID, err := conn.GetNetworkXIDToPublicIdMap(ctx, accountId) if err != nil { return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get network xid to public id map: %w", err)) } - allowedUserIds, groupsToUserIds, err := GetAllowedUsersViaPgxConnection(ctx, tx.Conn(), accountId) + allowedUserIds, groupsToUserIds, err := conn.GetAllowedUsers(ctx, accountId) if err != nil { return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get allowed users: %w", err)) } - dnsSettings, err := GetDnsSettingsViaPgxConnection(ctx, tx.Conn(), accountId) + dnsSettings, err := conn.GetDnsSettings(ctx, accountId) if err != nil { return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get dns settings: %w", err)) } - domains, err := GetDomainsViaPgxConnection(ctx, tx.Conn(), accountId) + domains, err := conn.GetDomains(ctx, accountId) if err != nil { return rollbackAndReturnError(ctx, tx, err) } - services, err := GetPrivateServicesViaPgxConnection(ctx, tx.Conn(), accountId) + services, err := conn.GetPrivateServices(ctx, accountId) if err != nil { return rollbackAndReturnError(ctx, tx, err) } - proxyTargetedDomainResourceIDs, err := GetProxyTargetedDomainResourceIDsViaPgxConnection(ctx, tx.Conn(), accountId) + proxyTargetedDomainResourceIDs, err := conn.GetProxyTargetedDomainResourceIDs(ctx, accountId) if err != nil { return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get proxy targeted domain resources: %w", err)) } @@ -165,7 +168,7 @@ func toSliceOfPtrs[T any](all []T) []*T { return toret } -func serviceDomainZone(svc Service, ds []Domain) string { +func serviceDomainZone(svc networkmapdb.Service, ds []networkmapdb.Domain) string { if domainFromSuffix(svc.Domain.String, svc.ProxyCluster.String) { return svc.ProxyCluster.String } @@ -190,7 +193,7 @@ func domainFromSuffix(domain, suffix string) bool { return domain == suffix || strings.HasSuffix(domain, "."+suffix) } -func buildPrivateServiceCandidates(svcs []Service, domains []Domain, proxyPeersByCluster map[string][]*nmdata.Peer) []networkmap.PrivateServiceCandidate { +func buildPrivateServiceCandidates(svcs []networkmapdb.Service, domains []networkmapdb.Domain, proxyPeersByCluster map[string][]*nmdata.Peer) []networkmap.PrivateServiceCandidate { var out []networkmap.PrivateServiceCandidate if len(proxyPeersByCluster) == 0 { diff --git a/management/internals/network_map_db/pgsql/network_resource.go b/management/internals/network_map_db/pgsql/network_resource.go index cf08d35a7..27f4e75c6 100644 --- a/management/internals/network_map_db/pgsql/network_resource.go +++ b/management/internals/network_map_db/pgsql/network_resource.go @@ -19,16 +19,8 @@ const ( ` ) -func (pg *PgStore) GetNetworkResources(ctx context.Context, accountId string) ([]nmdata.NetworkResource, error) { - c, err := pg.Pool.Acquire(ctx) - if err != nil { - return nil, err - } - return GetNetworkResourcesViaPgxConnection(ctx, c.Conn(), accountId) -} - -func GetNetworkResourcesViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.NetworkResource, error) { - rows, err := con.Query(ctx, GetNetworkResourcesQuery, accountId) +func (pgc *PgStoreConn) GetNetworkResources(ctx context.Context, accountId string) ([]nmdata.NetworkResource, error) { + rows, err := pgc.Conn.Query(ctx, GetNetworkResourcesQuery, accountId) if err != nil { return nil, err } diff --git a/management/internals/network_map_db/pgsql/network_router.go b/management/internals/network_map_db/pgsql/network_router.go index f4a76d69f..42d5e3b28 100644 --- a/management/internals/network_map_db/pgsql/network_router.go +++ b/management/internals/network_map_db/pgsql/network_router.go @@ -25,16 +25,8 @@ const ( ` ) -func (pg *PgStore) GetNetworkRouters(ctx context.Context, accountId string) (map[string]map[string]*nmdata.NetworkRouter, error) { - c, err := pg.Pool.Acquire(ctx) - if err != nil { - return nil, err - } - return GetNetworkRoutersViaPgxConnection(ctx, c.Conn(), accountId) -} - -func GetNetworkRoutersViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) (map[string]map[string]*nmdata.NetworkRouter, error) { - rows, err := con.Query(ctx, GetNetworkRouterQuery, accountId) +func (pgc *PgStoreConn) GetNetworkRouters(ctx context.Context, accountId string) (map[string]map[string]*nmdata.NetworkRouter, error) { + rows, err := pgc.Conn.Query(ctx, GetNetworkRouterQuery, accountId) if err != nil { return nil, err } diff --git a/management/internals/network_map_db/pgsql/networks.go b/management/internals/network_map_db/pgsql/networks.go index 2862025f5..501b74d81 100644 --- a/management/internals/network_map_db/pgsql/networks.go +++ b/management/internals/network_map_db/pgsql/networks.go @@ -14,24 +14,16 @@ const ( ` ) -func (pg *PgStore) GetNetworks(ctx context.Context, accountId string) ([]network, error) { - c, err := pg.Pool.Acquire(ctx) - if err != nil { - return nil, err - } - return GetNetworksViaPgxConnection(ctx, c.Conn(), accountId) -} - -func GetNetworksViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]network, error) { - rows, err := con.Query(ctx, GetNetworksQuery, accountId) +func (pgc *PgStoreConn) GetNetworks(ctx context.Context, accountId string) ([]network, error) { + rows, err := pgc.Conn.Query(ctx, GetNetworksQuery, accountId) if err != nil { return nil, err } return pgx.CollectRows(rows, pgx.RowToStructByName[network]) } -func GetNetworkXIDToPublicIdMapViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) (map[string]string, error) { - networks, err := GetNetworksViaPgxConnection(ctx, con, accountId) +func (pgc *PgStoreConn) GetNetworkXIDToPublicIdMap(ctx context.Context, accountId string) (map[string]string, error) { + networks, err := pgc.GetNetworks(ctx, accountId) if err != nil { return nil, err } diff --git a/management/internals/network_map_db/pgsql/peer.go b/management/internals/network_map_db/pgsql/peer.go index 7a4360e74..3c92ae0b5 100644 --- a/management/internals/network_map_db/pgsql/peer.go +++ b/management/internals/network_map_db/pgsql/peer.go @@ -22,16 +22,8 @@ const ( ` ) -func (pg *PgStore) GetPeers(ctx context.Context, accountId string) ([]nmdata.Peer, map[string][]*nmdata.Peer, error) { - c, err := pg.Pool.Acquire(ctx) - if err != nil { - return nil, nil, err - } - return GetPeersViaPgxConnection(ctx, c.Conn(), accountId) -} - -func GetPeersViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.Peer, map[string][]*nmdata.Peer, error) { - rows, err := con.Query(ctx, GetPeersQuery, accountId) +func (pgc *PgStoreConn) GetPeers(ctx context.Context, accountId string) ([]nmdata.Peer, map[string][]*nmdata.Peer, error) { + rows, err := pgc.Conn.Query(ctx, GetPeersQuery, accountId) if err != nil { return nil, nil, err } diff --git a/management/internals/network_map_db/pgsql/pg_store.go b/management/internals/network_map_db/pgsql/pg_store.go index e18a0c830..4843dd5d1 100644 --- a/management/internals/network_map_db/pgsql/pg_store.go +++ b/management/internals/network_map_db/pgsql/pg_store.go @@ -5,6 +5,7 @@ import ( "fmt" "time" + "github.com/jackc/pgx/v5" "github.com/jackc/pgx/v5/pgxpool" networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db" ) @@ -22,6 +23,12 @@ type PgStore struct { Pool *pgxpool.Pool } +type PgStoreConn struct { + Conn *pgx.Conn +} + +var _ networkmapdb.NetworkMapDBStoreConn = &PgStoreConn{} + func NewPostgresqlStore(ctx context.Context, dsn string) (*PgStore, error) { pool, err := connectToPgDb(ctx, dsn) if err != nil { @@ -31,6 +38,10 @@ func NewPostgresqlStore(ctx context.Context, dsn string) (*PgStore, error) { return &PgStore{Pool: pool}, nil } +func (p *PgStore) UsingConnection(c *pgx.Conn) networkmapdb.NetworkMapDBStoreConn { + return &PgStoreConn{Conn: c} +} + func connectToPgDb(ctx context.Context, dsn string) (*pgxpool.Pool, error) { config, err := pgxpool.ParseConfig(dsn) if err != nil { diff --git a/management/internals/network_map_db/pgsql/policy.go b/management/internals/network_map_db/pgsql/policy.go index da8e1e545..535d95acf 100644 --- a/management/internals/network_map_db/pgsql/policy.go +++ b/management/internals/network_map_db/pgsql/policy.go @@ -22,16 +22,8 @@ const ( ` ) -func (pg *PgStore) GetPolicies(ctx context.Context, accountId string) ([]nmdata.Policy, map[string]map[string]any, map[string]map[string]any, error) { - c, err := pg.Pool.Acquire(ctx) - if err != nil { - return nil, nil, nil, err - } - return GetPoliciesViaPgxConnection(ctx, c.Conn(), accountId) -} - -func GetPoliciesViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.Policy, map[string]map[string]any, map[string]map[string]any, error) { - rows, err := con.Query(ctx, GetPoliciesQuery, accountId) +func (pgc *PgStoreConn) GetPolicies(ctx context.Context, accountId string) ([]nmdata.Policy, map[string]map[string]any, map[string]map[string]any, error) { + rows, err := pgc.Conn.Query(ctx, GetPoliciesQuery, accountId) if err != nil { return nil, nil, nil, err } diff --git a/management/internals/network_map_db/pgsql/posture.go b/management/internals/network_map_db/pgsql/posture.go index f995811ca..e4ef8420c 100644 --- a/management/internals/network_map_db/pgsql/posture.go +++ b/management/internals/network_map_db/pgsql/posture.go @@ -19,16 +19,8 @@ const ( ` ) -func (pg *PgStore) GetPostureChecks(ctx context.Context, accountId string) ([]nmdata.PostureChecks, map[string]string, error) { - c, err := pg.Pool.Acquire(ctx) - if err != nil { - return nil, nil, err - } - return GetPostureChecksViaPgxConnection(ctx, c.Conn(), accountId) -} - -func GetPostureChecksViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.PostureChecks, map[string]string, error) { - rows, err := con.Query(ctx, GetPostureChecksQuery, accountId) +func (pgc *PgStoreConn) GetPostureChecks(ctx context.Context, accountId string) ([]nmdata.PostureChecks, map[string]string, error) { + rows, err := pgc.Conn.Query(ctx, GetPostureChecksQuery, accountId) if err != nil { return nil, nil, err } diff --git a/management/internals/network_map_db/pgsql/route.go b/management/internals/network_map_db/pgsql/route.go index 60752b50f..7020dd026 100644 --- a/management/internals/network_map_db/pgsql/route.go +++ b/management/internals/network_map_db/pgsql/route.go @@ -21,16 +21,8 @@ const ( ` ) -func (pg *PgStore) GetRoutes(ctx context.Context, accountId string) ([]nmdata.Route, error) { - c, err := pg.Pool.Acquire(ctx) - if err != nil { - return nil, err - } - return GetRoutesViaPgxConnection(ctx, c.Conn(), accountId) -} - -func GetRoutesViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) ([]nmdata.Route, error) { - rows, err := con.Query(ctx, GetRoutesQuery, accountId) +func (pgc *PgStoreConn) GetRoutes(ctx context.Context, accountId string) ([]nmdata.Route, error) { + rows, err := pgc.Conn.Query(ctx, GetRoutesQuery, accountId) if err != nil { return nil, err } diff --git a/management/internals/network_map_db/pgsql/service.go b/management/internals/network_map_db/pgsql/service.go index fce62a207..5d82046be 100644 --- a/management/internals/network_map_db/pgsql/service.go +++ b/management/internals/network_map_db/pgsql/service.go @@ -2,9 +2,9 @@ package networkmap_pgsql import ( "context" - "database/sql" "github.com/jackc/pgx/v5" + networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db" ) const ( @@ -23,25 +23,17 @@ const ( ` ) -func (pg *PgStore) GetPrivateServices(ctx context.Context, accountId string) ([]Service, error) { - c, err := pg.Pool.Acquire(ctx) - if err != nil { - return nil, err - } - return GetPrivateServicesViaPgxConnection(ctx, c.Conn(), accountId) -} - -func GetPrivateServicesViaPgxConnection(ctx context.Context, conn *pgx.Conn, accountId string) ([]Service, error) { - rows, err := conn.Query(ctx, GetServicesQuery, accountId) +func (pgc *PgStoreConn) GetPrivateServices(ctx context.Context, accountId string) ([]networkmapdb.Service, error) { + rows, err := pgc.Conn.Query(ctx, GetServicesQuery, accountId) if err != nil { return nil, err } - return pgx.CollectRows(rows, pgx.RowToStructByName[Service]) + return pgx.CollectRows(rows, pgx.RowToStructByName[networkmapdb.Service]) } -func GetProxyTargetedDomainResourceIDsViaPgxConnection(ctx context.Context, conn *pgx.Conn, accountId string) (map[string]struct{}, error) { - rows, err := conn.Query(ctx, GetProxyTargetedDomainResourcesQuery, accountId) +func (pgc *PgStoreConn) GetProxyTargetedDomainResourceIDs(ctx context.Context, accountId string) (map[string]struct{}, error) { + rows, err := pgc.Conn.Query(ctx, GetProxyTargetedDomainResourcesQuery, accountId) if err != nil { return nil, err } @@ -57,11 +49,3 @@ func GetProxyTargetedDomainResourceIDsViaPgxConnection(ctx context.Context, conn } return toret, nil } - -type Service struct { - Enabled sql.NullBool - Private sql.NullBool - AccessGroups []string - ProxyCluster sql.NullString - Domain sql.NullString -} diff --git a/management/internals/network_map_db/pgsql/user.go b/management/internals/network_map_db/pgsql/user.go index 04fa93835..9e22c3575 100644 --- a/management/internals/network_map_db/pgsql/user.go +++ b/management/internals/network_map_db/pgsql/user.go @@ -19,16 +19,8 @@ const ( ` ) -func (pg *PgStore) GetAllowedUsers(ctx context.Context, accountId string) (map[string]struct{}, map[string][]string, error) { - c, err := pg.Pool.Acquire(ctx) - if err != nil { - return nil, nil, err - } - return GetAllowedUsersViaPgxConnection(ctx, c.Conn(), accountId) -} - -func GetAllowedUsersViaPgxConnection(ctx context.Context, con *pgx.Conn, accountId string) (map[string]struct{}, map[string][]string, error) { - rows, err := con.Query(ctx, GetAllowedUserIdsQuery, accountId) +func (pgc *PgStoreConn) GetAllowedUsers(ctx context.Context, accountId string) (map[string]struct{}, map[string][]string, error) { + rows, err := pgc.Conn.Query(ctx, GetAllowedUserIdsQuery, accountId) if err != nil { return nil, nil, err } @@ -38,7 +30,7 @@ func GetAllowedUsersViaPgxConnection(ctx context.Context, con *pgx.Conn, account return nil, nil, err } - rows, err = con.Query(ctx, GetAllGroupIdQuery, accountId) + rows, err = pgc.Conn.Query(ctx, GetAllGroupIdQuery, accountId) if err != nil { return nil, nil, err } diff --git a/management/internals/network_map_db/sqlite/sqlite_store.go b/management/internals/network_map_db/sqlite/sqlite_store.go new file mode 100644 index 000000000..8d5226d6e --- /dev/null +++ b/management/internals/network_map_db/sqlite/sqlite_store.go @@ -0,0 +1,59 @@ +package networkmap_sqlite + +import ( + "context" + "net/url" + "path/filepath" + "runtime" + "strings" + + "database/sql" +) + +type SqliteStore struct { + Db *sql.DB +} + +func NewSqliteStore(ctx context.Context, storeFile, dataDir, dsn string) (*SqliteStore, 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 != ":memory:" && !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 := sql.Open("sqlite3", "") + if err != nil { + return nil, err + } + return &SqliteStore{Db: db}, nil +}