Merge branch 'main' into reverse-proxy-allow-match-or

# Conflicts:
#	management/server/store/sql_store.go
#	management/server/store/sql_store_service_test.go
This commit is contained in:
Viktor Liu
2026-07-31 21:57:57 +02:00
216 changed files with 13000 additions and 4032 deletions
+217 -143
View File
@@ -6,6 +6,7 @@ import (
"encoding/json"
"errors"
"fmt"
"math"
"net"
"net/netip"
"net/url"
@@ -2258,125 +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, restrictions,
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 restrictions []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,
&restrictions,
&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(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 = 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
}
@@ -2387,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
}
@@ -2433,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)
@@ -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
}
@@ -5,6 +5,7 @@ import (
"os"
"runtime"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
@@ -45,6 +46,94 @@ func TestSqlStore_GetAccount_PrivateServiceRoundtrip(t *testing.T) {
})
}
// 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")
})
}
func TestSqlStore_GetAccount_ServiceRestrictionsRoundtrip(t *testing.T) {
if os.Getenv("CI") == "true" && (runtime.GOOS == "darwin" || runtime.GOOS == "windows") {
t.Skip("skip CI tests on darwin and windows")
@@ -78,8 +167,8 @@ func TestSqlStore_GetAccount_ServiceRestrictionsRoundtrip(t *testing.T) {
// including allow_match, survives the read path (Postgres pgx path
// included via runTestForAllEngines).
got := loaded.Services[0].Restrictions
assert.Equal(t, []string{"203.0.113.0/24"}, got.AllowedCIDRs)
assert.Equal(t, []string{"US"}, got.AllowedCountries)
assert.Equal(t, "any", got.AllowMatch)
assert.Equal(t, []string{"203.0.113.0/24"}, got.AllowedCIDRs, "restrictions allowed CIDRs")
assert.Equal(t, []string{"US"}, got.AllowedCountries, "restrictions allowed countries")
assert.Equal(t, "any", got.AllowMatch, "restrictions allow match")
})
}
+48
View File
@@ -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{})
@@ -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 {
+68
View File
@@ -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")
})
}
+4
View File
@@ -321,6 +321,10 @@ func (am *DefaultAccountManager) DeleteUser(ctx context.Context, accountID, init
return err
}
if targetUser.AccountID != accountID {
return status.NewUserNotFoundError(targetUserID)
}
if targetUser.Role == types.UserRoleOwner {
return status.NewOwnerDeletePermissionError()
}
+46
View File
@@ -802,6 +802,52 @@ func TestUser_DeleteUser_SelfDelete(t *testing.T) {
}
}
func TestUser_DeleteUser_OtherAccount(t *testing.T) {
testStore, cleanup, err := store.NewTestStoreFromSQL(context.Background(), "", t.TempDir())
if err != nil {
t.Fatalf("Error when creating store: %s", err)
}
t.Cleanup(cleanup)
account := newAccountWithId(context.Background(), mockAccountID, mockUserID, "", "", "", false)
if err = testStore.SaveAccount(context.Background(), account); err != nil {
t.Fatalf("Error when saving account: %s", err)
}
otherAccount := newAccountWithId(context.Background(), "otherAccount", "otherOwner", "", "", "", false)
otherAccount.Users["otherRegularUser"] = &types.User{
Id: "otherRegularUser",
AccountID: "otherAccount",
Role: types.UserRoleUser,
}
otherAccount.Users["otherServiceUser"] = &types.User{
Id: "otherServiceUser",
AccountID: "otherAccount",
Role: types.UserRoleUser,
IsServiceUser: true,
ServiceUserName: "otherServiceUser",
}
if err = testStore.SaveAccount(context.Background(), otherAccount); err != nil {
t.Fatalf("Error when saving other account: %s", err)
}
am := DefaultAccountManager{
Store: testStore,
eventStore: &activity.InMemoryEventStore{},
permissionsManager: permissions.NewManager(testStore),
}
for _, targetUserID := range []string{"otherRegularUser", "otherServiceUser"} {
t.Run(targetUserID, func(t *testing.T) {
err := am.DeleteUser(context.Background(), mockAccountID, mockUserID, targetUserID)
assert.Equal(t, status.NewUserNotFoundError(targetUserID), err)
_, err = testStore.GetUserByUserID(context.Background(), store.LockingStrengthNone, targetUserID)
assert.NoError(t, err, "user of another account must not be deleted")
})
}
}
func TestUser_DeleteUser_regularUser(t *testing.T) {
store, cleanup, err := store.NewTestStoreFromSQL(context.Background(), "", t.TempDir())
if err != nil {