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

@@ -1,245 +0,0 @@
package networkmap_pgsql
import (
"context"
"fmt"
"strings"
"github.com/jackc/pgx/v5"
"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"
)
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})
if err != nil {
return nil, err
}
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 := conn.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)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get groups: %w", err))
}
nsGroups, err := conn.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)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get network resources: %w", err))
}
routers, err := conn.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)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get network: %w", err))
}
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 := conn.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)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get posture checks: %w", err))
}
routes, err := conn.GetRoutes(ctx, accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get routes: %w", err))
}
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 := conn.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)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get dns settings: %w", err))
}
domains, err := conn.GetDomains(ctx, accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, err)
}
services, err := conn.GetPrivateServices(ctx, accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, err)
}
proxyTargetedDomainResourceIDs, err := conn.GetProxyTargetedDomainResourceIDs(ctx, accountId)
if err != nil {
return rollbackAndReturnError(ctx, tx, fmt.Errorf("failed to get proxy targeted domain resources: %w", err))
}
resourcePolicies := make(map[string][]*nmdata.Policy)
for _, resource := range networkResources {
if !resource.Enabled {
continue
}
networkResourceGroups := resourceToGroupIdx[resource.ID]
for _, policy := range policies {
if !policy.Enabled {
continue
}
if _, ok := policyToDestinationResourceIdx[policy.ID][resource.ID]; ok {
resourcePolicies[resource.ID] = append(resourcePolicies[resource.ID], &policy) // TODO (dmitri) maybe use public id?
continue
}
if groupIds, ok := policyToDestinationGroupIdx[policy.ID]; ok {
for networkResourceGroup := range networkResourceGroups {
if _, ok := groupIds[networkResourceGroup]; ok {
resourcePolicies[resource.ID] = append(resourcePolicies[resource.ID], &policy)
break
}
}
}
}
}
if err = tx.Commit(ctx); err != nil {
log.WithContext(ctx).Warnf("failed to commit network map read transaction: %v", err)
}
toret := networkmap.NetworkMapData{
AccountSettings: &acctSettings,
DNSSettings: &dnsSettings,
Network: &network,
Peers: toMap(peers, func(p nmdata.Peer) string { return p.ID }),
Groups: toMap(groups, func(g nmdata.Group) string { return g.ID }),
Policies: toSliceOfPtrs(policies),
ResourcePolicies: resourcePolicies,
Routes: toSliceOfPtrs(routes),
Routers: routers,
NameServerGroups: toSliceOfPtrs(nsGroups),
NetworkResources: toSliceOfPtrs(networkResources),
PostureChecks: toMap(postureChecks, func(pc nmdata.PostureChecks) string { return pc.ID }),
AllowedUserIDs: allowedUserIds,
GroupIDToUserIDs: groupsToUserIds,
NetworkXIDToPublicID: networkXIDToPublicID, // TODO (dmitri) maybe we can switch to public ids everywhere?
AppliedZoneCandidates: dnsZones,
PrivateServiceCandidates: buildPrivateServiceCandidates(services, domains, proxyPeers),
PostureCheckXIDToPublicID: postureCheckXIDToPublicID,
ProxyTargetedDomainResourceIDs: proxyTargetedDomainResourceIDs,
}
return &toret, nil
}
func rollbackAndReturnError(ctx context.Context, tx pgx.Tx, err error) (*networkmap.NetworkMapData, error) {
if errr := tx.Rollback(ctx); errr != nil {
log.WithContext(ctx).Warnf("failed to rollback network map read transaction: %v", errr)
}
return nil, err
}
func toMap[T any](all []T, id func(t T) string) map[string]*T {
toret := make(map[string]*T, len(all))
for _, t := range all {
toret[id(t)] = &t
}
return toret
}
func toSliceOfPtrs[T any](all []T) []*T {
toret := make([]*T, 0, len(all))
for _, t := range all {
toret = append(toret, &t)
}
return toret
}
func serviceDomainZone(svc networkmapdb.Service, ds []networkmapdb.Domain) string {
if domainFromSuffix(svc.Domain.String, svc.ProxyCluster.String) {
return svc.ProxyCluster.String
}
var zoneName string
for _, domain := range ds {
if domain.TargetCluster.String != svc.ProxyCluster.String {
continue
}
if domainFromSuffix(svc.Domain.String, domain.Domain.String) && len(domain.Domain.String) > len(zoneName) {
zoneName = domain.Domain.String
}
}
return zoneName
}
func domainFromSuffix(domain, suffix string) bool {
if suffix == "" {
return false
}
return domain == suffix || strings.HasSuffix(domain, "."+suffix)
}
func buildPrivateServiceCandidates(svcs []networkmapdb.Service, domains []networkmapdb.Domain, proxyPeersByCluster map[string][]*nmdata.Peer) []networkmap.PrivateServiceCandidate {
var out []networkmap.PrivateServiceCandidate
if len(proxyPeersByCluster) == 0 {
return out
}
for _, svc := range svcs {
if !svc.Enabled.Bool || !svc.Private.Bool {
continue
}
if len(svc.AccessGroups) == 0 {
continue
}
domainZone := serviceDomainZone(svc, domains)
if domainZone == "" {
continue
}
var records []nmdata.SimpleRecord
for _, proxyPeer := range proxyPeersByCluster[svc.ProxyCluster.String] {
if !proxyPeer.IP.IsValid() {
continue
}
records = append(records, nmdata.SimpleRecord{
Name: dns.Fqdn(svc.Domain.String),
Type: int(dns.TypeA),
Class: "IN",
TTL: 5,
RData: proxyPeer.IP.String(),
})
}
if len(records) == 0 {
continue
}
out = append(out, networkmap.PrivateServiceCandidate{
AccessGroups: svc.AccessGroups,
Zone: nmdata.CustomZone{
Domain: dns.Fqdn(domainZone),
Records: records,
NonAuthoritative: true,
SearchDomainDisabled: true,
},
})
}
return out
}

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 {