mirror of
https://github.com/netbirdio/netbird.git
synced 2026-08-25 00:51:28 +02:00
moved GetNetworkMapData implementation to NetworkMapDBStoreImpl
Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
This commit is contained in:
@@ -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
|
||||
}
|
||||
|
||||
@@ -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 {
|
||||
@@ -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 {
|
||||
|
||||
@@ -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 {
|
||||
|
||||
Reference in New Issue
Block a user