moved GetNetworkMapData implementation to NetworkMapDBStoreImpl

Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
This commit is contained in:
Dmitri Dolguikh
2026-08-11 15:42:11 +02:00
parent fe5e5c08eb
commit 9427fa87e1
4 changed files with 96 additions and 54 deletions

View File

@@ -3,8 +3,6 @@ package networkmapdb
import (
"context"
"golang.org/x/exp/maps"
"github.com/netbirdio/netbird/management/server/integrations/integrated_validator"
"github.com/netbirdio/netbird/management/server/settings"
"github.com/netbirdio/netbird/shared/management/networkmap"
@@ -12,7 +10,7 @@ import (
)
type NetworkMapDBStore interface { //nolint:revive // established name across the codebase
GetNetworkMapData(ctx context.Context, accountId string) (*networkmap.NetworkMapData, error)
BeginTx(ctx context.Context) (NetworkMapDBStoreConn, error)
}
type NetworkMapDBStoreConn interface { //nolint:revive // established name across the codebase
@@ -33,6 +31,9 @@ type NetworkMapDBStoreConn interface { //nolint:revive // established name acros
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)
CommitTx(ctx context.Context) error
RollbackTx(ctx context.Context) error
}
type NetworkMapDBStoreImpl struct { //nolint:revive // established name across the codebase
@@ -48,22 +49,3 @@ func NewNetworkMapDBStoreImpl(store NetworkMapDBStore, integratedPeerValidator i
extraSettingsManager: extraSettingsManager,
}
}
func (s *NetworkMapDBStoreImpl) GetNetworkMapData(ctx context.Context, accountId string) (*networkmap.NetworkMapData, error) {
nmdata, err := s.store.GetNetworkMapData(ctx, accountId)
if err != nil {
return nil, err
}
extraSettings, err := s.extraSettingsManager.GetExtraSettings(ctx, accountId)
if err != nil {
return nil, err
}
nmdata.ValidatedPeers, err = s.integratedPeerValidator.GetValidatedPeers(ctx, accountId, maps.Values(nmdata.Groups), maps.Values(nmdata.Peers), extraSettings)
if err != nil {
return nil, err
}
return nmdata, nil
}

View File

@@ -1,92 +1,89 @@
package networkmap_pgsql
package networkmapdb
import (
"context"
"fmt"
"strings"
"github.com/jackc/pgx/v5"
"github.com/miekg/dns"
log "github.com/sirupsen/logrus"
"golang.org/x/exp/maps"
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"
)
func (pg *PgStore) GetNetworkMapData(ctx context.Context, accountId string) (*networkmap.NetworkMapData, error) {
tx, err := pg.Pool.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.RepeatableRead, AccessMode: pgx.ReadOnly})
func (s *NetworkMapDBStoreImpl) GetNetworkMapData(ctx context.Context, accountId string) (*networkmap.NetworkMapData, error) {
tx, err := s.store.BeginTx(ctx)
if err != nil {
return nil, err
}
conn := pg.UsingConnection(tx.Conn())
acctSettings, err := conn.GetAccountSettings(ctx, accountId)
acctSettings, err := tx.GetAccountSettings(ctx, accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get account settings: %w", err))
}
dnsZones, err := conn.GetAppliedZoneCandidates(ctx, accountId)
dnsZones, err := tx.GetAppliedZoneCandidates(ctx, accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get applied zone candidates: %w", err))
}
groups, resourceToGroupIdx, err := conn.GetGroups(ctx, accountId)
groups, resourceToGroupIdx, err := tx.GetGroups(ctx, accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get groups: %w", err))
}
nsGroups, err := conn.GetNameServerGroups(ctx, accountId)
nsGroups, err := tx.GetNameServerGroups(ctx, accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get nameserver groups: %w", err))
}
networkResources, err := conn.GetNetworkResources(ctx, accountId)
networkResources, err := tx.GetNetworkResources(ctx, accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get network resources: %w", err))
}
routers, err := conn.GetNetworkRouters(ctx, accountId)
routers, err := tx.GetNetworkRouters(ctx, accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get network routers: %w", err))
}
network, err := conn.GetNetwork(ctx, accountId)
network, err := tx.GetNetwork(ctx, accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get network: %w", err))
}
peers, proxyPeers, err := conn.GetPeers(ctx, accountId)
peers, proxyPeers, err := tx.GetPeers(ctx, accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get peers: %w", err))
}
policies, policyToDestinationResourceIdx, policyToDestinationGroupIdx, err := conn.GetPolicies(ctx, accountId)
policies, policyToDestinationResourceIdx, policyToDestinationGroupIdx, err := tx.GetPolicies(ctx, accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get policies: %w", err))
}
postureChecks, postureCheckXIDToPublicID, err := conn.GetPostureChecks(ctx, accountId)
postureChecks, postureCheckXIDToPublicID, err := tx.GetPostureChecks(ctx, accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get posture checks: %w", err))
}
routes, err := conn.GetRoutes(ctx, accountId)
routes, err := tx.GetRoutes(ctx, accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get routes: %w", err))
}
networkXIDToPublicID, err := conn.GetNetworkXIDToPublicIdMap(ctx, accountId)
networkXIDToPublicID, err := tx.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 := conn.GetAllowedUsers(ctx, accountId)
allowedUserIds, groupsToUserIds, err := tx.GetAllowedUsers(ctx, accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get allowed users: %w", err))
}
dnsSettings, err := conn.GetDnsSettings(ctx, accountId)
dnsSettings, err := tx.GetDnsSettings(ctx, accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get dns settings: %w", err))
}
domains, err := conn.GetDomains(ctx, accountId)
domains, err := tx.GetDomains(ctx, accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, err)
}
services, err := conn.GetPrivateServices(ctx, accountId)
services, err := tx.GetPrivateServices(ctx, accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, err)
}
proxyTargetedDomainResourceIDs, err := conn.GetProxyTargetedDomainResourceIDs(ctx, accountId)
proxyTargetedDomainResourceIDs, err := tx.GetProxyTargetedDomainResourceIDs(ctx, accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get proxy targeted domain resources: %w", err))
}
@@ -116,7 +113,7 @@ func (pg *PgStore) GetNetworkMapData(ctx context.Context, accountId string) (*ne
}
}
if err = tx.Commit(ctx); err != nil {
if err = tx.CommitTx(ctx); err != nil {
log.WithContext(ctx).Warnf("failed to commit network map read transaction: %v", err)
}
@@ -142,11 +139,21 @@ func (pg *PgStore) GetNetworkMapData(ctx context.Context, accountId string) (*ne
ProxyTargetedDomainResourceIDs: proxyTargetedDomainResourceIDs,
}
extraSettings, err := s.extraSettingsManager.GetExtraSettings(ctx, accountId)
if err != nil {
return nil, err
}
toret.ValidatedPeers, err = s.integratedPeerValidator.GetValidatedPeers(ctx, accountId, maps.Values(toret.Groups), maps.Values(toret.Peers), extraSettings)
if err != nil {
return nil, err
}
return &toret, nil
}
func rollbackAndReturnError(ctx context.Context, tx pgx.Tx, err error) (*networkmap.NetworkMapData, error) {
if errr := tx.Rollback(ctx); errr != nil {
func rollbackAndReturnError(ctx context.Context, tx NetworkMapDBStoreConn, err error) (*networkmap.NetworkMapData, error) {
if errr := tx.RollbackTx(ctx); errr != nil {
log.WithContext(ctx).Warnf("failed to rollback network map read transaction: %v", errr)
}
return nil, err
@@ -168,7 +175,7 @@ func toSliceOfPtrs[T any](all []T) []*T {
return toret
}
func serviceDomainZone(svc networkmapdb.Service, ds []networkmapdb.Domain) string {
func serviceDomainZone(svc Service, ds []Domain) string {
if domainFromSuffix(svc.Domain.String, svc.ProxyCluster.String) {
return svc.ProxyCluster.String
}
@@ -193,7 +200,7 @@ func domainFromSuffix(domain, suffix string) bool {
return domain == suffix || strings.HasSuffix(domain, "."+suffix)
}
func buildPrivateServiceCandidates(svcs []networkmapdb.Service, domains []networkmapdb.Domain, proxyPeersByCluster map[string][]*nmdata.Peer) []networkmap.PrivateServiceCandidate {
func buildPrivateServiceCandidates(svcs []Service, domains []Domain, proxyPeersByCluster map[string][]*nmdata.Peer) []networkmap.PrivateServiceCandidate {
var out []networkmap.PrivateServiceCandidate
if len(proxyPeersByCluster) == 0 {

View File

@@ -3,9 +3,11 @@ package networkmap_pgsql
import (
"context"
"fmt"
"reflect"
"time"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgconn"
"github.com/jackc/pgx/v5/pgxpool"
networkmapdb "github.com/netbirdio/netbird/management/internals/network_map_db"
)
@@ -24,7 +26,12 @@ type PgStore struct {
}
type PgStoreConn struct {
Conn *pgx.Conn
Conn pgInterface
}
type pgInterface interface {
Query(ctx context.Context, sql string, args ...any) (pgx.Rows, error)
Exec(ctx context.Context, sql string, arguments ...any) (pgconn.CommandTag, error)
}
var _ networkmapdb.NetworkMapDBStoreConn = &PgStoreConn{}
@@ -42,6 +49,30 @@ func (p *PgStore) UsingConnection(c *pgx.Conn) networkmapdb.NetworkMapDBStoreCon
return &PgStoreConn{Conn: c}
}
func (p *PgStore) BeginTx(ctx context.Context) (networkmapdb.NetworkMapDBStoreConn, error) {
tx, err := p.Pool.BeginTx(ctx, pgx.TxOptions{IsoLevel: pgx.RepeatableRead, AccessMode: pgx.ReadOnly})
if err != nil {
return nil, err
}
return &PgStoreConn{Conn: tx}, nil
}
func (c *PgStoreConn) RollbackTx(ctx context.Context) error {
tx, ok := c.Conn.(pgx.Tx)
if !ok {
return fmt.Errorf("expected an pgx.Tx got %s", reflect.TypeOf(c.Conn).Kind())
}
return tx.Rollback(ctx)
}
func (c *PgStoreConn) CommitTx(ctx context.Context) error {
tx, ok := c.Conn.(pgx.Tx)
if !ok {
return fmt.Errorf("expected an sql.Tx got %s", reflect.TypeOf(c.Conn).Kind())
}
return tx.Commit(ctx)
}
func connectToPgDb(ctx context.Context, dsn string) (*pgxpool.Pool, error) {
config, err := pgxpool.ParseConfig(dsn)
if err != nil {

View File

@@ -3,8 +3,10 @@ package networkmap_sqlite
import (
"context"
"errors"
"fmt"
"net/url"
"path/filepath"
"reflect"
"runtime"
"strings"
@@ -73,8 +75,28 @@ func NewSqliteStore(ctx context.Context, storeFile, dataDir string) (*SqliteStor
return &SqliteStore{Db: db}, nil
}
func (s *SqliteStore) WithTx(tx *sql.Tx) *SqliteStoreConn {
return &SqliteStoreConn{Conn: tx}
func (s *SqliteStore) BeginTx(ctx context.Context) (*SqliteStoreConn, error) {
tx, err := s.Db.BeginTx(ctx, &sql.TxOptions{ReadOnly: true, Isolation: sql.LevelRepeatableRead})
if err != nil {
return nil, err
}
return &SqliteStoreConn{Conn: tx}, nil
}
func (sc *SqliteStoreConn) RollbackTx(ctx context.Context) error {
tx, ok := sc.Conn.(*sql.Tx)
if !ok {
return fmt.Errorf("expected an sql.Tx got %s", reflect.TypeOf(sc.Conn).Kind())
}
return tx.Rollback()
}
func (sc *SqliteStoreConn) CommitTx(ctx context.Context) error {
tx, ok := sc.Conn.(*sql.Tx)
if !ok {
return fmt.Errorf("expected an sql.Tx got %s", reflect.TypeOf(sc.Conn).Kind())
}
return tx.Commit()
}
func (s *SqliteStore) UsingConn() *SqliteStoreConn {