diff --git a/.goreleaser_ui.yaml b/.goreleaser_ui.yaml index 197fcd440..1157e6379 100644 --- a/.goreleaser_ui.yaml +++ b/.goreleaser_ui.yaml @@ -24,6 +24,8 @@ builds: ldflags: - -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser mod_timestamp: "{{ .CommitTimestamp }}" + tags: + - production - id: netbird-ui-windows-amd64 dir: client/ui @@ -39,6 +41,8 @@ builds: - -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser - -H windowsgui mod_timestamp: "{{ .CommitTimestamp }}" + tags: + - production - id: netbird-ui-windows-arm64 dir: client/ui @@ -55,6 +59,8 @@ builds: - -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser - -H windowsgui mod_timestamp: "{{ .CommitTimestamp }}" + tags: + - production archives: - id: linux-arch diff --git a/.goreleaser_ui_darwin.yaml b/.goreleaser_ui_darwin.yaml index 96e15371a..47b991344 100644 --- a/.goreleaser_ui_darwin.yaml +++ b/.goreleaser_ui_darwin.yaml @@ -29,6 +29,8 @@ builds: ldflags: - -s -w -X github.com/netbirdio/netbird/version.version={{.Version}} -X main.commit={{.Commit}} -X main.date={{.CommitDate}} -X main.builtBy=goreleaser mod_timestamp: "{{ .CommitTimestamp }}" + tags: + - production universal_binaries: - id: netbird-ui-darwin diff --git a/client/android/profile_manager.go b/client/android/profile_manager.go index 87c001396..9a051137c 100644 --- a/client/android/profile_manager.go +++ b/client/android/profile_manager.go @@ -189,6 +189,19 @@ func (pm *ProfileManager) LogoutProfile(id string) error { return nil } +// RenameProfile changes a profile's display name. The profile ID, and therefore +// its on-disk filename, is left untouched: only the "name" field of the config +// is rewritten. This works for the default profile too, whose config lives in +// netbird.cfg rather than under profiles/. +func (pm *ProfileManager) RenameProfile(id string, newName string) error { + if err := pm.serviceMgr.RenameProfile(profilemanager.ID(id), androidUsername, newName); err != nil { + return fmt.Errorf("failed to rename profile: %w", err) + } + + log.Infof("renamed profile %s to: %s", id, newName) + return nil +} + // RemoveProfile deletes a profile func (pm *ProfileManager) RemoveProfile(id string) error { // Use ServiceManager (removes profile from profiles/ directory) diff --git a/client/internal/routemanager/sysctl/sysctl_linux.go b/client/internal/routemanager/sysctl/sysctl_linux.go index f96a57f37..46b7c9fb7 100644 --- a/client/internal/routemanager/sysctl/sysctl_linux.go +++ b/client/internal/routemanager/sysctl/sysctl_linux.go @@ -20,6 +20,8 @@ const ( rpFilterPath = "net.ipv4.conf.all.rp_filter" rpFilterInterfacePath = "net.ipv4.conf.%s.rp_filter" srcValidMarkPath = "net.ipv4.conf.all.src_valid_mark" + percentEscape = "%25" + dotEscape = "%2E" ) type iface interface { @@ -56,7 +58,11 @@ func Setup(wgIface iface) (map[string]int, error) { continue } - i := fmt.Sprintf(rpFilterInterfacePath, intf.Name) + // Escape '%' and '.' so they survive the dot-to-slash conversion in Set() + safeName := strings.ReplaceAll(intf.Name, "%", percentEscape) + safeName = strings.ReplaceAll(safeName, ".", dotEscape) + + i := fmt.Sprintf(rpFilterInterfacePath, safeName) oldVal, err := Set(i, 2, true) if err != nil { result = multierror.Append(result, err) @@ -70,7 +76,11 @@ func Setup(wgIface iface) (map[string]int, error) { // Set sets a sysctl configuration, if onlyIfOne is true it will only set the new value if it's set to 1 func Set(key string, desiredValue int, onlyIfOne bool) (int, error) { - path := fmt.Sprintf("/proc/sys/%s", strings.ReplaceAll(key, ".", "/")) + path := strings.ReplaceAll(key, ".", "/") + // Unescape interface dots and percent signs + path = strings.ReplaceAll(path, dotEscape, ".") + path = strings.ReplaceAll(path, percentEscape, "%") + path = fmt.Sprintf("/proc/sys/%s", path) currentValue, err := os.ReadFile(path) if err != nil { return -1, fmt.Errorf("read sysctl %s: %w", key, err) diff --git a/management/internals/shared/grpc/components_envelope_response.go b/management/internals/shared/grpc/components_envelope_response.go index cedd1b889..820708c98 100644 --- a/management/internals/shared/grpc/components_envelope_response.go +++ b/management/internals/shared/grpc/components_envelope_response.go @@ -50,7 +50,7 @@ func ToComponentSyncResponse( // TODO (dmitri) consider using invariants? // enableSSH := computeSSHEnabledForPeer(components, peer) - peerConfig := toPeerConfig(peer, components.Network, dnsName, settings, httpConfig, deviceFlowConfig, enableSSH) + peerConfig := toPeerConfig(peer, components.Network, dnsName, settings, httpConfig, deviceFlowConfig, enableSSH, components.ForceRoutingPeerDNSResolution) includeIPv6 := peer.SupportsIPv6() && peer.IPv6.IsValid() useSourcePrefixes := peer.SupportsSourcePrefixes() diff --git a/management/internals/shared/grpc/conversion.go b/management/internals/shared/grpc/conversion.go index 696d28f5c..74ceb3370 100644 --- a/management/internals/shared/grpc/conversion.go +++ b/management/internals/shared/grpc/conversion.go @@ -119,7 +119,7 @@ func toNetbirdConfig(config *nbconfig.Config, turnCredentials *Token, relayToken return nbConfig } -func toPeerConfig(peer *nbpeer.Peer, network *types.Network, dnsName string, settings *types.Settings, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, enableSSH bool) *proto.PeerConfig { +func toPeerConfig(peer *nbpeer.Peer, network *types.Network, dnsName string, settings *types.Settings, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, enableSSH bool, forceRoutingPeerDNS bool) *proto.PeerConfig { netmask, _ := network.Net.Mask.Size() fqdn := peer.FQDN(dnsName) @@ -135,7 +135,7 @@ func toPeerConfig(peer *nbpeer.Peer, network *types.Network, dnsName string, set Address: fmt.Sprintf("%s/%d", peer.IP.String(), netmask), SshConfig: sshConfig, Fqdn: fqdn, - RoutingPeerDnsResolutionEnabled: settings.RoutingPeerDNSResolutionEnabled, + RoutingPeerDnsResolutionEnabled: settings.RoutingPeerDNSResolutionEnabled || peer.ProxyMeta.Embedded || forceRoutingPeerDNS, LazyConnectionEnabled: settings.LazyConnectionEnabled, AutoUpdate: &proto.AutoUpdateSettings{ Version: settings.AutoUpdateVersion, @@ -162,12 +162,12 @@ func ToSyncResponse(ctx context.Context, config *nbconfig.Config, httpConfig *nb useSourcePrefixes := peer.SupportsSourcePrefixes() response := &proto.SyncResponse{ - PeerConfig: toPeerConfig(peer, networkMap.Network, dnsName, settings, httpConfig, deviceFlowConfig, networkMap.EnableSSH), + PeerConfig: toPeerConfig(peer, networkMap.Network, dnsName, settings, httpConfig, deviceFlowConfig, networkMap.EnableSSH, networkMap.ForceRoutingPeerDNSResolution), NetworkMap: &proto.NetworkMap{ Serial: networkMap.Network.CurrentSerial(), Routes: networkmap.ToProtocolRoutes(networkMap.Routes), DNSConfig: networkmap.ToProtocolDNSConfig(networkMap.DNSConfig, dnsCache, dnsFwdPort), - PeerConfig: toPeerConfig(peer, networkMap.Network, dnsName, settings, httpConfig, deviceFlowConfig, networkMap.EnableSSH), + PeerConfig: toPeerConfig(peer, networkMap.Network, dnsName, settings, httpConfig, deviceFlowConfig, networkMap.EnableSSH, networkMap.ForceRoutingPeerDNSResolution), }, Checks: toProtocolChecks(ctx, checks), } diff --git a/management/internals/shared/grpc/conversion_test.go b/management/internals/shared/grpc/conversion_test.go index 402b4fd07..38d370740 100644 --- a/management/internals/shared/grpc/conversion_test.go +++ b/management/internals/shared/grpc/conversion_test.go @@ -2,6 +2,7 @@ package grpc import ( "fmt" + "net" "net/netip" "reflect" "testing" @@ -14,6 +15,7 @@ import ( "github.com/netbirdio/netbird/management/internals/controllers/network_map" "github.com/netbirdio/netbird/management/internals/controllers/network_map/controller/cache" nbconfig "github.com/netbirdio/netbird/management/internals/server/config" + nbpeer "github.com/netbirdio/netbird/management/server/peer" "github.com/netbirdio/netbird/management/server/types" "github.com/netbirdio/netbird/shared/management/networkmap" ) @@ -301,3 +303,35 @@ func TestToNetbirdConfig_RelayInvariant(t *testing.T) { assert.True(t, nbCfg.Metrics.Enabled, "metrics flag should carry the settings value") }) } + +func TestToPeerConfig_RoutingPeerDNSResolution(t *testing.T) { + network := &types.Network{Net: net.IPNet{IP: net.IPv4(100, 0, 0, 0), Mask: net.CIDRMask(8, 32)}} + + newPeer := func(embedded bool) *nbpeer.Peer { + p := &nbpeer.Peer{IP: netip.MustParseAddr("100.0.0.1")} + p.ProxyMeta.Embedded = embedded + return p + } + + tests := []struct { + name string + globalFlag bool + embedded bool + forceParam bool + wantEnabled bool + }{ + {name: "global off, regular peer, no force", wantEnabled: false}, + {name: "global on wins", globalFlag: true, wantEnabled: true}, + {name: "embedded proxy peer forced", embedded: true, wantEnabled: true}, + {name: "routing peer forced via param", forceParam: true, wantEnabled: true}, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + settings := &types.Settings{RoutingPeerDNSResolutionEnabled: tt.globalFlag} + cfg := toPeerConfig(newPeer(tt.embedded), network, "netbird.selfhosted", settings, nil, nil, false, tt.forceParam) + assert.Equal(t, tt.wantEnabled, cfg.RoutingPeerDnsResolutionEnabled, + "RoutingPeerDnsResolutionEnabled should reflect global || embedded || forced") + }) + } +} diff --git a/management/internals/shared/grpc/server.go b/management/internals/shared/grpc/server.go index 3b7d62ac7..485f05a92 100644 --- a/management/internals/shared/grpc/server.go +++ b/management/internals/shared/grpc/server.go @@ -921,7 +921,7 @@ func (s *Server) prepareLoginResponse(ctx context.Context, peer *nbpeer.Peer, ne // if peer has reached this point then it has logged in loginResp := &proto.LoginResponse{ NetbirdConfig: toNetbirdConfig(s.config, nil, relayToken, nil, settings), - PeerConfig: toPeerConfig(peer, network, s.networkMapController.GetDNSDomain(settings), settings, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, enableSSH), + PeerConfig: toPeerConfig(peer, network, s.networkMapController.GetDNSDomain(settings), settings, s.config.HttpConfig, s.config.DeviceAuthorizationFlow, enableSSH, false), Checks: toProtocolChecks(ctx, postureChecks), } diff --git a/management/server/store/sql_store.go b/management/server/store/sql_store.go index 3ad870ad3..670d9f781 100644 --- a/management/server/store/sql_store.go +++ b/management/server/store/sql_store.go @@ -6,6 +6,7 @@ import ( "encoding/json" "errors" "fmt" + "math" "net" "net/netip" "net/url" @@ -2258,117 +2259,30 @@ func (s *SqlStore) getPostureChecks(ctx context.Context, accountID string) ([]*p return checks, nil } -func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpservice.Service, error) { - const serviceQuery = `SELECT id, account_id, name, domain, enabled, auth, - meta_created_at, meta_certificate_issued_at, meta_status, proxy_cluster, - pass_host_header, rewrite_redirects, session_private_key, session_public_key, - mode, listen_port, port_auto_assigned, source, source_peer, terminated, - private, access_groups - FROM services WHERE account_id = $1` +// serviceSelectColumns and targetSelectColumns are the column lists the Postgres +// pgx read path scans. They must stay in sync with the rpservice.Service and +// rpservice.Target gorm models; TestPgxServiceColumnsMatchGorm enforces this. +const serviceSelectColumns = `id, account_id, name, domain, enabled, auth, restrictions, + meta_created_at, meta_certificate_issued_at, meta_last_renewed_at, meta_status, proxy_cluster, + pass_host_header, rewrite_redirects, session_private_key, session_public_key, + mode, listen_port, port_auto_assigned, source, source_peer, terminated, + private, access_groups` - const targetsQuery = `SELECT id, account_id, service_id, path, host, port, protocol, - target_id, target_type, enabled - FROM targets WHERE service_id = ANY($1)` +const targetSelectColumns = `id, account_id, service_id, path, host, port, protocol, + target_id, target_type, enabled, proxy_protocol, + skip_tls_verify, request_timeout, session_idle_timeout, path_rewrite, custom_headers, + direct_upstream, middlewares, capture_max_request_bytes, capture_max_response_bytes, + capture_content_types, agent_network, disable_access_log` + +func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpservice.Service, error) { + const serviceQuery = `SELECT ` + serviceSelectColumns + ` FROM services WHERE account_id = $1` serviceRows, err := s.pool.Query(ctx, serviceQuery, accountID) if err != nil { return nil, err } - services, err := pgx.CollectRows(serviceRows, func(row pgx.CollectableRow) (*rpservice.Service, error) { - var s rpservice.Service - var auth []byte - var accessGroups []byte - var createdAt, certIssuedAt sql.NullTime - var status, proxyCluster, sessionPrivateKey, sessionPublicKey sql.NullString - var mode, source, sourcePeer sql.NullString - var terminated, portAutoAssigned, private sql.NullBool - var listenPort sql.NullInt64 - err := row.Scan( - &s.ID, - &s.AccountID, - &s.Name, - &s.Domain, - &s.Enabled, - &auth, - &createdAt, - &certIssuedAt, - &status, - &proxyCluster, - &s.PassHostHeader, - &s.RewriteRedirects, - &sessionPrivateKey, - &sessionPublicKey, - &mode, - &listenPort, - &portAutoAssigned, - &source, - &sourcePeer, - &terminated, - &private, - &accessGroups, - ) - if err != nil { - return nil, err - } - - if auth != nil { - if err := json.Unmarshal(auth, &s.Auth); err != nil { - return nil, err - } - } - - if len(accessGroups) > 0 { - if err := json.Unmarshal(accessGroups, &s.AccessGroups); err != nil { - return nil, fmt.Errorf("unmarshal access_groups: %w", err) - } - } - - if private.Valid { - s.Private = private.Bool - } - - s.Meta = rpservice.Meta{} - if createdAt.Valid { - s.Meta.CreatedAt = createdAt.Time - } - if certIssuedAt.Valid { - t := certIssuedAt.Time - s.Meta.CertificateIssuedAt = &t - } - if status.Valid { - s.Meta.Status = status.String - } - if proxyCluster.Valid { - s.ProxyCluster = proxyCluster.String - } - if sessionPrivateKey.Valid { - s.SessionPrivateKey = sessionPrivateKey.String - } - if sessionPublicKey.Valid { - s.SessionPublicKey = sessionPublicKey.String - } - if mode.Valid { - s.Mode = mode.String - } - if source.Valid { - s.Source = source.String - } - if sourcePeer.Valid { - s.SourcePeer = sourcePeer.String - } - if terminated.Valid { - s.Terminated = terminated.Bool - } - if portAutoAssigned.Valid { - s.PortAutoAssigned = portAutoAssigned.Bool - } - if listenPort.Valid { - s.ListenPort = uint16(listenPort.Int64) - } - s.Targets = []*rpservice.Target{} - return &s, nil - }) + services, err := pgx.CollectRows(serviceRows, scanService) if err != nil { return nil, err } @@ -2379,39 +2293,12 @@ func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpserv serviceIDs := make([]string, len(services)) serviceMap := make(map[string]*rpservice.Service) - for i, s := range services { - serviceIDs[i] = s.ID - serviceMap[s.ID] = s + for i, svc := range services { + serviceIDs[i] = svc.ID + serviceMap[svc.ID] = svc } - targetRows, err := s.pool.Query(ctx, targetsQuery, serviceIDs) - if err != nil { - return nil, err - } - - targets, err := pgx.CollectRows(targetRows, func(row pgx.CollectableRow) (*rpservice.Target, error) { - var t rpservice.Target - var path sql.NullString - err := row.Scan( - &t.ID, - &t.AccountID, - &t.ServiceID, - &path, - &t.Host, - &t.Port, - &t.Protocol, - &t.TargetId, - &t.TargetType, - &t.Enabled, - ) - if err != nil { - return nil, err - } - if path.Valid { - t.Path = &path.String - } - return &t, nil - }) + targets, err := s.getServiceTargets(ctx, serviceIDs) if err != nil { return nil, err } @@ -2425,6 +2312,201 @@ func (s *SqlStore) getServices(ctx context.Context, accountID string) ([]*rpserv return services, nil } +func scanService(row pgx.CollectableRow) (*rpservice.Service, error) { + var s rpservice.Service + var auth []byte + var restrictions []byte + var accessGroups []byte + var createdAt, certIssuedAt, lastRenewedAt sql.NullTime + var status, proxyCluster, sessionPrivateKey, sessionPublicKey sql.NullString + var mode, source, sourcePeer sql.NullString + var terminated, portAutoAssigned, private sql.NullBool + var listenPort sql.NullInt64 + err := row.Scan( + &s.ID, + &s.AccountID, + &s.Name, + &s.Domain, + &s.Enabled, + &auth, + &restrictions, + &createdAt, + &certIssuedAt, + &lastRenewedAt, + &status, + &proxyCluster, + &s.PassHostHeader, + &s.RewriteRedirects, + &sessionPrivateKey, + &sessionPublicKey, + &mode, + &listenPort, + &portAutoAssigned, + &source, + &sourcePeer, + &terminated, + &private, + &accessGroups, + ) + if err != nil { + return nil, err + } + + if auth != nil { + if err := json.Unmarshal(auth, &s.Auth); err != nil { + return nil, err + } + } + + if len(restrictions) > 0 { + if err := json.Unmarshal(restrictions, &s.Restrictions); err != nil { + return nil, fmt.Errorf("unmarshal restrictions: %w", err) + } + } + + if len(accessGroups) > 0 { + if err := json.Unmarshal(accessGroups, &s.AccessGroups); err != nil { + return nil, fmt.Errorf("unmarshal access_groups: %w", err) + } + } + + if private.Valid { + s.Private = private.Bool + } + + s.Meta = serviceMetaFromRow(createdAt, certIssuedAt, lastRenewedAt, status) + if proxyCluster.Valid { + s.ProxyCluster = proxyCluster.String + } + if sessionPrivateKey.Valid { + s.SessionPrivateKey = sessionPrivateKey.String + } + if sessionPublicKey.Valid { + s.SessionPublicKey = sessionPublicKey.String + } + if mode.Valid { + s.Mode = mode.String + } + if source.Valid { + s.Source = source.String + } + if sourcePeer.Valid { + s.SourcePeer = sourcePeer.String + } + if terminated.Valid { + s.Terminated = terminated.Bool + } + if portAutoAssigned.Valid { + s.PortAutoAssigned = portAutoAssigned.Bool + } + if listenPort.Valid { + if listenPort.Int64 < 0 || listenPort.Int64 > math.MaxUint16 { + return nil, fmt.Errorf("listen_port %d out of range", listenPort.Int64) + } + s.ListenPort = uint16(listenPort.Int64) + } + s.Targets = []*rpservice.Target{} + return &s, nil +} + +func serviceMetaFromRow(createdAt, certIssuedAt, lastRenewedAt sql.NullTime, status sql.NullString) rpservice.Meta { + meta := rpservice.Meta{} + if createdAt.Valid { + meta.CreatedAt = createdAt.Time + } + if certIssuedAt.Valid { + t := certIssuedAt.Time + meta.CertificateIssuedAt = &t + } + if lastRenewedAt.Valid { + t := lastRenewedAt.Time + meta.LastRenewedAt = &t + } + if status.Valid { + meta.Status = status.String + } + return meta +} + +func (s *SqlStore) getServiceTargets(ctx context.Context, serviceIDs []string) ([]*rpservice.Target, error) { + const targetsQuery = `SELECT ` + targetSelectColumns + ` FROM targets WHERE service_id = ANY($1)` + + rows, err := s.pool.Query(ctx, targetsQuery, serviceIDs) + if err != nil { + return nil, err + } + + return pgx.CollectRows(rows, scanTarget) +} + +func scanTarget(row pgx.CollectableRow) (*rpservice.Target, error) { + var t rpservice.Target + var path sql.NullString + var pathRewrite sql.NullString + var proxyProtocol, skipTLSVerify, directUpstream, agentNetwork, disableAccessLog sql.NullBool + var requestTimeout, sessionIdleTimeout, captureMaxRequestBytes, captureMaxResponseBytes sql.NullInt64 + var customHeaders, middlewares, captureContentTypes []byte + err := row.Scan( + &t.ID, + &t.AccountID, + &t.ServiceID, + &path, + &t.Host, + &t.Port, + &t.Protocol, + &t.TargetId, + &t.TargetType, + &t.Enabled, + &proxyProtocol, + &skipTLSVerify, + &requestTimeout, + &sessionIdleTimeout, + &pathRewrite, + &customHeaders, + &directUpstream, + &middlewares, + &captureMaxRequestBytes, + &captureMaxResponseBytes, + &captureContentTypes, + &agentNetwork, + &disableAccessLog, + ) + if err != nil { + return nil, err + } + if path.Valid { + t.Path = &path.String + } + + t.ProxyProtocol = proxyProtocol.Bool + t.Options.SkipTLSVerify = skipTLSVerify.Bool + t.Options.RequestTimeout = time.Duration(requestTimeout.Int64) + t.Options.SessionIdleTimeout = time.Duration(sessionIdleTimeout.Int64) + t.Options.PathRewrite = rpservice.PathRewriteMode(pathRewrite.String) + t.Options.DirectUpstream = directUpstream.Bool + t.Options.CaptureMaxRequestBytes = captureMaxRequestBytes.Int64 + t.Options.CaptureMaxResponseBytes = captureMaxResponseBytes.Int64 + t.Options.AgentNetwork = agentNetwork.Bool + t.Options.DisableAccessLog = disableAccessLog.Bool + + if len(customHeaders) > 0 { + if err := json.Unmarshal(customHeaders, &t.Options.CustomHeaders); err != nil { + return nil, fmt.Errorf("unmarshal custom_headers: %w", err) + } + } + if len(middlewares) > 0 { + if err := json.Unmarshal(middlewares, &t.Options.Middlewares); err != nil { + return nil, fmt.Errorf("unmarshal middlewares: %w", err) + } + } + if len(captureContentTypes) > 0 { + if err := json.Unmarshal(captureContentTypes, &t.Options.CaptureContentTypes); err != nil { + return nil, fmt.Errorf("unmarshal capture_content_types: %w", err) + } + } + return &t, nil +} + func (s *SqlStore) getNetworks(ctx context.Context, accountID string) ([]*networkTypes.Network, error) { const query = `SELECT id, account_id, public_id, name, description FROM networks WHERE account_id = $1` rows, err := s.pool.Query(ctx, query, accountID) diff --git a/management/server/store/sql_store_pgx_parity_test.go b/management/server/store/sql_store_pgx_parity_test.go new file mode 100644 index 000000000..1f17817d0 --- /dev/null +++ b/management/server/store/sql_store_pgx_parity_test.go @@ -0,0 +1,74 @@ +package store + +import ( + "strings" + "sync" + "testing" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "gorm.io/gorm/schema" + + rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service" +) + +// TestPgxServiceColumnsMatchGorm guards the Postgres pgx read path against +// drifting from the gorm model. The SQLite/MySQL gorm path loads rows by struct, +// so a new column on a model is picked up automatically, but the hand-written +// pgx SELECT in sql_store.go must be updated by hand. This test fails when a +// gorm column is missing from the pgx column list, which otherwise silently +// returns zero-valued on Postgres with no compile error. +func TestPgxServiceColumnsMatchGorm(t *testing.T) { + tests := []struct { + name string + model any + selectColumns string + // excluded lists gorm columns intentionally not loaded by the pgx path. + excluded map[string]struct{} + }{ + { + name: "service", + model: &rpservice.Service{}, + selectColumns: serviceSelectColumns, + }, + { + name: "target", + model: &rpservice.Target{}, + selectColumns: targetSelectColumns, + }, + } + + for _, tc := range tests { + t.Run(tc.name, func(t *testing.T) { + selected := parseColumnList(tc.selectColumns) + for _, col := range gormColumnNames(t, tc.model) { + if _, ok := tc.excluded[col]; ok { + continue + } + _, ok := selected[col] + assert.Truef(t, ok, + "gorm column %q is not read by the Postgres pgx SELECT; add it to %sSelectColumns in sql_store.go (or to the test's excluded set if it is intentionally not loaded)", + col, tc.name) + } + }) + } +} + +func parseColumnList(cols string) map[string]struct{} { + set := make(map[string]struct{}) + for _, c := range strings.Split(cols, ",") { + if c = strings.TrimSpace(c); c != "" { + set[c] = struct{}{} + } + } + return set +} + +// gormColumnNames returns the DB column names gorm would migrate for the model, +// using the same default naming strategy the store configures. +func gormColumnNames(t *testing.T, model any) []string { + t.Helper() + sch, err := schema.Parse(model, &sync.Map{}, schema.NamingStrategy{}) + require.NoError(t, err) + return sch.DBNames +} diff --git a/management/server/store/sql_store_service_test.go b/management/server/store/sql_store_service_test.go index 34999da4b..0e14fbdab 100644 --- a/management/server/store/sql_store_service_test.go +++ b/management/server/store/sql_store_service_test.go @@ -5,6 +5,7 @@ import ( "os" "runtime" "testing" + "time" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" @@ -44,3 +45,91 @@ func TestSqlStore_GetAccount_PrivateServiceRoundtrip(t *testing.T) { assert.Equal(t, []string{"grp-admins", "grp-ops"}, got.AccessGroups) }) } + +// TestSqlStore_GetAccount_ServiceTargetOptionsRoundtrip guards the Postgres pgx +// read path (getServices) against silently dropping columns present on the gorm +// model. Before the fix these fields loaded correctly on SQLite but came back +// zero-valued on Postgres because the hand-written SELECT and scan omitted them. +func TestSqlStore_GetAccount_ServiceTargetOptionsRoundtrip(t *testing.T) { + if os.Getenv("CI") == "true" && (runtime.GOOS == "darwin" || runtime.GOOS == "windows") { + t.Skip("skip CI tests on darwin and windows") + } + + runTestForAllEngines(t, "", func(t *testing.T, store Store) { + ctx := context.Background() + account := newAccountWithId(ctx, "account_svc_opts", "testuser", "") + require.NoError(t, store.SaveAccount(ctx, account)) + + renewedAt := time.Now().UTC().Truncate(time.Second) + targetPath := "/api" + svc := &rpservice.Service{ + ID: "svc-opts", + AccountID: account.Id, + Name: "opts-svc", + Domain: "opts.example", + Enabled: true, + Mode: rpservice.ModeHTTP, + Restrictions: rpservice.AccessRestrictions{ + AllowedCIDRs: []string{"10.0.0.0/8"}, + BlockedCountries: []string{"XX"}, + CrowdSecMode: "block", + }, + Meta: rpservice.Meta{ + LastRenewedAt: &renewedAt, + }, + Targets: []*rpservice.Target{ + { + AccountID: account.Id, + ServiceID: "svc-opts", + Path: &targetPath, + Host: "backend.internal", + Port: 8080, + Protocol: "http", + TargetId: "tgt-1", + Enabled: true, + ProxyProtocol: true, + Options: rpservice.TargetOptions{ + SkipTLSVerify: true, + RequestTimeout: 30 * time.Second, + SessionIdleTimeout: 5 * time.Minute, + PathRewrite: rpservice.PathRewritePreserve, + CustomHeaders: map[string]string{"X-Foo": "bar"}, + DirectUpstream: true, + CaptureMaxRequestBytes: 1024, + CaptureMaxResponseBytes: 2048, + CaptureContentTypes: []string{"application/json"}, + AgentNetwork: true, + DisableAccessLog: true, + }, + }, + }, + } + require.NoError(t, store.CreateService(ctx, svc)) + + loaded, err := store.GetAccount(ctx, account.Id) + require.NoError(t, err) + require.Len(t, loaded.Services, 1) + + got := loaded.Services[0] + assert.Equal(t, []string{"10.0.0.0/8"}, got.Restrictions.AllowedCIDRs, "restrictions allowed CIDRs") + assert.Equal(t, []string{"XX"}, got.Restrictions.BlockedCountries, "restrictions blocked countries") + assert.Equal(t, "block", got.Restrictions.CrowdSecMode, "restrictions crowdsec mode") + require.NotNil(t, got.Meta.LastRenewedAt, "meta last renewed at") + assert.WithinDuration(t, renewedAt, *got.Meta.LastRenewedAt, time.Second, "meta last renewed at") + + require.Len(t, got.Targets, 1) + tg := got.Targets[0] + assert.True(t, tg.ProxyProtocol, "target proxy protocol") + assert.True(t, tg.Options.SkipTLSVerify, "options skip TLS verify") + assert.Equal(t, 30*time.Second, tg.Options.RequestTimeout, "options request timeout") + assert.Equal(t, 5*time.Minute, tg.Options.SessionIdleTimeout, "options session idle timeout") + assert.Equal(t, rpservice.PathRewritePreserve, tg.Options.PathRewrite, "options path rewrite") + assert.Equal(t, map[string]string{"X-Foo": "bar"}, tg.Options.CustomHeaders, "options custom headers") + assert.True(t, tg.Options.DirectUpstream, "options direct upstream") + assert.Equal(t, int64(1024), tg.Options.CaptureMaxRequestBytes, "options capture max request bytes") + assert.Equal(t, int64(2048), tg.Options.CaptureMaxResponseBytes, "options capture max response bytes") + assert.Equal(t, []string{"application/json"}, tg.Options.CaptureContentTypes, "options capture content types") + assert.True(t, tg.Options.AgentNetwork, "options agent network") + assert.True(t, tg.Options.DisableAccessLog, "options disable access log") + }) +} diff --git a/management/server/types/account.go b/management/server/types/account.go index 588e63a09..1a3a30544 100644 --- a/management/server/types/account.go +++ b/management/server/types/account.go @@ -1517,6 +1517,54 @@ func (a *Account) GetResourceRoutersMap() map[string]map[string]*routerTypes.Net return routers } +// forcesRoutingPeerDNSResolution reports whether the given peer must run +// routing-peer DNS resolution regardless of the account-global +// RoutingPeerDNSResolutionEnabled setting. It returns true when the peer is a +// router for a domain network resource that is targeted by an enabled +// reverse-proxy service, so the peer's DNS forwarder starts and can resolve +// the target for the embedded proxy peers. Embedded proxy peers themselves are +// handled at PeerConfig build time. +func (a *Account) forcesRoutingPeerDNSResolution(peerID string, routers map[string]map[string]*routerTypes.NetworkRouter) bool { + targeted := a.proxyTargetedDomainResourceIDs() + if len(targeted) == 0 { + return false + } + + for _, resource := range a.NetworkResources { + if resource == nil || !resource.Enabled || resource.Type != resourceTypes.Domain { + continue + } + if _, ok := targeted[resource.ID]; !ok { + continue + } + if _, isRouter := routers[resource.NetworkID][peerID]; isRouter { + return true + } + } + + return false +} + +// proxyTargetedDomainResourceIDs returns the set of domain network resource IDs +// targeted by an enabled, non-terminated reverse-proxy service. +func (a *Account) proxyTargetedDomainResourceIDs() map[string]struct{} { + ids := make(map[string]struct{}) + for _, svc := range a.Services { + if svc == nil || !svc.Enabled || svc.Terminated { + continue + } + for _, target := range svc.Targets { + if target == nil || !target.Enabled { + continue + } + if target.TargetType == service.TargetTypeDomain { + ids[target.TargetId] = struct{}{} + } + } + } + return ids +} + // getPoliciesSourcePeers collects all unique peers from the source groups defined in the given policies. func getPoliciesSourcePeers(policies []*Policy, groups map[string]*Group) map[string]struct{} { sourcePeers := make(map[string]struct{}) diff --git a/management/server/types/account_components.go b/management/server/types/account_components.go index af27788d8..6fc904c0b 100644 --- a/management/server/types/account_components.go +++ b/management/server/types/account_components.go @@ -140,6 +140,8 @@ func (a *Account) GetPeerNetworkMapComponents( RouterPeers: make(map[string]*ComponentPeer), NetworkXIDToPublicID: make(map[string]string, len(a.Networks)), PostureCheckXIDToPublicID: make(map[string]string, len(a.PostureChecks)), + + ForceRoutingPeerDNSResolution: a.forcesRoutingPeerDNSResolution(peerID, routers), } for _, n := range a.Networks { if n != nil { diff --git a/management/server/types/account_test.go b/management/server/types/account_test.go index 67d9e1c6f..80f2a950a 100644 --- a/management/server/types/account_test.go +++ b/management/server/types/account_test.go @@ -1751,3 +1751,71 @@ func hasPrivateAccessPolicy(account *Account, serviceID string) bool { } return false } + +func TestForcesRoutingPeerDNSResolution(t *testing.T) { + buildAccountRes := func(serviceEnabled, targetEnabled, resourceEnabled bool, targetType service.TargetType, resType resourceTypes.NetworkResourceType) *Account { + return &Account{ + Id: "accountID", + Groups: map[string]*Group{ + "router-group": {ID: "router-group", Peers: []string{"router-peer-grp"}}, + }, + NetworkRouters: []*routerTypes.NetworkRouter{ + {ID: "r1", NetworkID: "net-1", AccountID: "accountID", Peer: "router-peer", Enabled: true}, + {ID: "r2", NetworkID: "net-1", AccountID: "accountID", PeerGroups: []string{"router-group"}, Enabled: true}, + }, + NetworkResources: []*resourceTypes.NetworkResource{ + {ID: "res-domain", AccountID: "accountID", NetworkID: "net-1", Type: resType, Domain: "example.org", Enabled: resourceEnabled}, + }, + Services: []*service.Service{ + { + ID: "svc-1", AccountID: "accountID", Enabled: serviceEnabled, + Targets: []*service.Target{ + {TargetId: "res-domain", TargetType: targetType, Enabled: targetEnabled}, + }, + }, + }, + } + } + + buildAccount := func(serviceEnabled, targetEnabled, resourceEnabled bool, targetType service.TargetType) *Account { + return buildAccountRes(serviceEnabled, targetEnabled, resourceEnabled, targetType, resourceTypes.Domain) + } + + t.Run("router peer for RP-targeted domain resource is forced", func(t *testing.T) { + account := buildAccount(true, true, true, service.TargetTypeDomain) + routers := account.GetResourceRoutersMap() + assert.True(t, account.forcesRoutingPeerDNSResolution("router-peer", routers), "direct router peer should be forced") + assert.True(t, account.forcesRoutingPeerDNSResolution("router-peer-grp", routers), "group-member router peer should be forced") + }) + + t.Run("non-router peer is not forced", func(t *testing.T) { + account := buildAccount(true, true, true, service.TargetTypeDomain) + assert.False(t, account.forcesRoutingPeerDNSResolution("other-peer", account.GetResourceRoutersMap())) + }) + + t.Run("not forced when service disabled", func(t *testing.T) { + account := buildAccount(false, true, true, service.TargetTypeDomain) + assert.False(t, account.forcesRoutingPeerDNSResolution("router-peer", account.GetResourceRoutersMap())) + }) + + t.Run("not forced when target disabled", func(t *testing.T) { + account := buildAccount(true, false, true, service.TargetTypeDomain) + assert.False(t, account.forcesRoutingPeerDNSResolution("router-peer", account.GetResourceRoutersMap())) + }) + + t.Run("not forced when resource disabled", func(t *testing.T) { + account := buildAccount(true, true, false, service.TargetTypeDomain) + assert.False(t, account.forcesRoutingPeerDNSResolution("router-peer", account.GetResourceRoutersMap())) + }) + + t.Run("not forced for non-domain target type", func(t *testing.T) { + account := buildAccount(true, true, true, service.TargetTypePeer) + assert.False(t, account.forcesRoutingPeerDNSResolution("router-peer", account.GetResourceRoutersMap())) + }) + + t.Run("not forced when targeted resource is not a domain", func(t *testing.T) { + account := buildAccountRes(true, true, true, service.TargetTypeDomain, resourceTypes.Host) + assert.False(t, account.forcesRoutingPeerDNSResolution("router-peer", account.GetResourceRoutersMap()), + "a domain target pointing at a non-domain resource must not force resolution") + }) +} diff --git a/shared/management/types/network.go b/shared/management/types/network.go index 72a5cc5b3..34ce60436 100644 --- a/shared/management/types/network.go +++ b/shared/management/types/network.go @@ -47,6 +47,10 @@ type NetworkMap struct { ForwardingRules []*ForwardingRule AuthorizedUsers map[string]map[string]struct{} EnableSSH bool + // ForceRoutingPeerDNSResolution forces the peer to run/use routing-peer DNS + // resolution regardless of the account-global setting, for reverse-proxy + // domain targets. + ForceRoutingPeerDNSResolution bool } func (nm *NetworkMap) Merge(other *NetworkMap) { @@ -56,6 +60,7 @@ func (nm *NetworkMap) Merge(other *NetworkMap) { nm.FirewallRules = mergeUnique(nm.FirewallRules, other.FirewallRules) nm.RoutesFirewallRules = mergeUnique(nm.RoutesFirewallRules, other.RoutesFirewallRules) nm.ForwardingRules = mergeUnique(nm.ForwardingRules, other.ForwardingRules) + nm.ForceRoutingPeerDNSResolution = nm.ForceRoutingPeerDNSResolution || other.ForceRoutingPeerDNSResolution } type comparableObject[T any] interface { diff --git a/shared/management/types/networkmap_components.go b/shared/management/types/networkmap_components.go index a708e99e1..c4c437e4b 100644 --- a/shared/management/types/networkmap_components.go +++ b/shared/management/types/networkmap_components.go @@ -56,6 +56,11 @@ type NetworkMapComponents struct { // true when returning an empty-like map (returned instead of nil) empty bool + + // ForceRoutingPeerDNSResolution forces the peer to run/use routing-peer DNS + // resolution regardless of the account-global setting, for reverse-proxy + // domain targets. + ForceRoutingPeerDNSResolution bool } type routeIndexEntry struct { @@ -190,6 +195,8 @@ func (c *NetworkMapComponents) Calculate(ctx context.Context) *NetworkMap { RoutesFirewallRules: append(networkResourcesFirewallRules, routesFirewallRules...), AuthorizedUsers: authorizedUsers, EnableSSH: sshEnabled, + + ForceRoutingPeerDNSResolution: c.ForceRoutingPeerDNSResolution, } }