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:
Dmitri Dolguikh
2026-08-07 19:23:52 +02:00
parent ac3b713dda
commit 935d1c5863
38 changed files with 325 additions and 375 deletions

View File

@@ -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{})
}

View File

@@ -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{})
}

View File

@@ -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{

View File

@@ -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},
})

View File

@@ -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)

View File

@@ -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) {

View File

@@ -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{

View File

@@ -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{

View File

@@ -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)

View File

@@ -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",

View File

@@ -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",

View File

@@ -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

View File

@@ -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}`)
}

View File

@@ -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{

View File

@@ -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",

View File

@@ -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",

View File

@@ -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": {},

View File

@@ -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{}{

View File

@@ -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

View File

@@ -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
}

View File

@@ -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
}

View File

@@ -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
}

View File

@@ -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])
}

View File

@@ -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
}

View File

@@ -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
}

View File

@@ -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
}

View File

@@ -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 {

View File

@@ -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
}

View File

@@ -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
}

View File

@@ -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
}

View File

@@ -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
}

View File

@@ -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 {

View File

@@ -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
}

View File

@@ -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
}

View File

@@ -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
}

View File

@@ -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
}

View File

@@ -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
}

View 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
}