mirror of
https://github.com/netbirdio/netbird.git
synced 2026-08-30 03:21:29 +02:00
move read-only queries to an interface to reuse in tests
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
This commit is contained in:
@@ -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{})
|
||||
}
|
||||
|
||||
@@ -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{})
|
||||
}
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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},
|
||||
})
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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<pre>[:/@])(?P<dbname>[^/?]+)(?P<post>\?|$)`)
|
||||
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) {
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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{
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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<pre>[:/@])(?P<dbname>[^/?]+)(?P<post>\?|$)`)
|
||||
return re.ReplaceAllString(dsn, `${pre}`+newDBName+`${post}`)
|
||||
}
|
||||
@@ -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{
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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",
|
||||
|
||||
@@ -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": {},
|
||||
|
||||
@@ -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{}{
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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])
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
59
management/internals/network_map_db/sqlite/sqlite_store.go
Normal file
59
management/internals/network_map_db/sqlite/sqlite_store.go
Normal file
@@ -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
|
||||
}
|
||||
Reference in New Issue
Block a user