mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-04 20:49:06 +02:00
Merge branch 'main' into embedded-vnc
This commit is contained in:
@@ -7,11 +7,10 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
cachestore "github.com/eko/gocache/lib/v4/store"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.opentelemetry.io/otel/metric/noop"
|
||||
"go.uber.org/mock/gomock"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
|
||||
proxymanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy/manager"
|
||||
@@ -31,7 +30,7 @@ import (
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
func testCacheStore(t *testing.T) cachestore.StoreInterface {
|
||||
func testCacheStore(t *testing.T) nbcache.Store {
|
||||
t.Helper()
|
||||
s, err := nbcache.NewStore(context.Background(), 30*time.Minute, 10*time.Minute, 100)
|
||||
require.NoError(t, err)
|
||||
@@ -295,6 +294,7 @@ func TestPersistNewService(t *testing.T) {
|
||||
assert.Equal(t, status.AlreadyExists, sErr.Type())
|
||||
})
|
||||
}
|
||||
|
||||
func TestPreserveExistingAuthSecrets(t *testing.T) {
|
||||
mgr := &Manager{}
|
||||
|
||||
|
||||
@@ -55,6 +55,8 @@ const (
|
||||
SourceEphemeral = "ephemeral"
|
||||
)
|
||||
|
||||
var ErrUnsupportedIPAddressUpstreamHost = errors.New("unsupported ip address for a direct upstream host")
|
||||
|
||||
type TargetOptions struct {
|
||||
SkipTLSVerify bool `json:"skip_tls_verify"`
|
||||
RequestTimeout time.Duration `json:"request_timeout,omitempty"`
|
||||
@@ -388,6 +390,7 @@ func (s *Service) ToProtoMapping(operation Operation, authToken string, oidcConf
|
||||
|
||||
if s.Auth.BearerAuth != nil && s.Auth.BearerAuth.Enabled {
|
||||
auth.Oidc = true
|
||||
auth.AllowedGroupIds = append([]string(nil), s.Auth.BearerAuth.DistributionGroups...)
|
||||
}
|
||||
|
||||
for _, h := range s.Auth.HeaderAuths {
|
||||
@@ -961,8 +964,8 @@ func (s *Service) validateHTTPTargets() error {
|
||||
return err
|
||||
}
|
||||
case TargetTypeSubnet:
|
||||
if target.Host == "" {
|
||||
return fmt.Errorf("target %d has empty host but target_type is %q", i, target.TargetType)
|
||||
if err := validateSubnetTarget(i, target); err != nil {
|
||||
return err
|
||||
}
|
||||
case TargetTypeCluster:
|
||||
if err := validateClusterTarget(i, target); err != nil {
|
||||
@@ -985,6 +988,34 @@ func (s *Service) validateHTTPTargets() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func validateSubnetTarget(idx int, target *Target) error {
|
||||
host := strings.TrimSpace(target.Host)
|
||||
if host == "" {
|
||||
return fmt.Errorf("target %d has empty host but target_type is %q", idx, target.TargetType)
|
||||
}
|
||||
if strings.ContainsAny(host, " \t/") {
|
||||
return fmt.Errorf("target %d: host %q contains invalid characters", idx, host)
|
||||
}
|
||||
if _, _, err := net.SplitHostPort(host); err == nil {
|
||||
return fmt.Errorf("target %d: host %q must not include a port (set target.port instead)", idx, host)
|
||||
}
|
||||
noBrackets := strings.TrimSuffix(strings.TrimPrefix(host, "["), "]")
|
||||
maybeip, err := netip.ParseAddr(noBrackets)
|
||||
if err != nil { // not an ip
|
||||
return nil //nolint:nilerr
|
||||
}
|
||||
if maybeip.Zone() != "" {
|
||||
return fmt.Errorf("invalid direct upstream host ip %s %w", maybeip.String(), ErrUnsupportedIPAddressUpstreamHost)
|
||||
}
|
||||
if !target.Options.DirectUpstream {
|
||||
return nil
|
||||
}
|
||||
if maybeip.IsLoopback() || maybeip.IsMulticast() || maybeip.IsLinkLocalUnicast() {
|
||||
return fmt.Errorf("invalid direct upstream host ip %s %w", maybeip.String(), ErrUnsupportedIPAddressUpstreamHost)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// validateClusterTarget cluster targets should not have empty hosts and should have direct upstream enabled.
|
||||
func validateClusterTarget(idx int, target *Target) error {
|
||||
host := strings.TrimSpace(target.Host)
|
||||
@@ -1019,6 +1050,15 @@ func validateDirectUpstreamHost(idx int, target *Target) error {
|
||||
if _, _, err := net.SplitHostPort(host); err == nil {
|
||||
return fmt.Errorf("target %d: host %q must not include a port (set target.port instead)", idx, host)
|
||||
}
|
||||
noBrackets := strings.TrimSuffix(strings.TrimPrefix(host, "["), "]")
|
||||
maybeip, err := netip.ParseAddr(noBrackets)
|
||||
if err != nil { // not an ip
|
||||
return nil //nolint:nilerr
|
||||
}
|
||||
if maybeip.Zone() != "" || maybeip.IsLoopback() || maybeip.IsMulticast() || maybeip.IsLinkLocalUnicast() {
|
||||
return fmt.Errorf("invalid direct upstream host ip %s %w", maybeip.String(), ErrUnsupportedIPAddressUpstreamHost)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
@@ -216,6 +216,64 @@ func TestValidateTargetOptions_CustomHeaders(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestValidate_DirectUpstreamHost(t *testing.T) {
|
||||
target := Target{TargetId: "id-1", TargetType: TargetTypePeer, Host: "10.0.0.1", Port: 80, Protocol: "http", Enabled: true, Options: TargetOptions{DirectUpstream: true}}
|
||||
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "127.0.0.2")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.NotNil(t, validateDirectUpstreamHost(0, targetWithHost(&target, "127.0.0.2:80")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "::1")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "::1%lo0")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "[::1]")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.NotNil(t, validateDirectUpstreamHost(0, targetWithHost(&target, "[::1]:80")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "[::1%lo0]")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.NotNil(t, validateDirectUpstreamHost(0, targetWithHost(&target, "[::1%lo0]:80")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "169.254.100.100")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "fe80::1")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "[fe80::1]")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "224.100.100.100")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "ff00::ffff")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.ErrorIs(t, validateDirectUpstreamHost(0, targetWithHost(&target, "[ff00::ffff]")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
|
||||
// empty host
|
||||
assert.Nil(t, validateDirectUpstreamHost(0, &Target{Options: TargetOptions{DirectUpstream: true}, Host: " "}))
|
||||
// host with a space
|
||||
assert.NotNil(t, validateDirectUpstreamHost(0, &Target{Options: TargetOptions{DirectUpstream: true}, Host: "with space"}))
|
||||
// host with a tab
|
||||
assert.NotNil(t, validateDirectUpstreamHost(0, &Target{Options: TargetOptions{DirectUpstream: true}, Host: "with\ttab"}))
|
||||
// host with a slash
|
||||
assert.NotNil(t, validateDirectUpstreamHost(0, &Target{Options: TargetOptions{DirectUpstream: true}, Host: "with/slash"}))
|
||||
}
|
||||
|
||||
func TestValidate_ValidateSubnetTarget(t *testing.T) {
|
||||
target := Target{TargetId: "id-1", TargetType: TargetTypeSubnet, Host: "10.0.0.1", Port: 80, Protocol: "http", Enabled: true, Options: TargetOptions{DirectUpstream: true}}
|
||||
assert.ErrorIs(t, validateSubnetTarget(0, targetWithHost(&target, "127.0.0.2")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.NotNil(t, validateSubnetTarget(0, targetWithHost(&target, "127.0.0.2:80")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.ErrorIs(t, validateSubnetTarget(0, targetWithHost(&target, "::1")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.NotNil(t, validateSubnetTarget(0, targetWithHost(&target, "[::1]:80")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.ErrorIs(t, validateSubnetTarget(0, targetWithHost(&target, "::1%lo0")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.ErrorIs(t, validateSubnetTarget(0, targetWithHost(&target, "[::1%lo0]")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.NotNil(t, validateSubnetTarget(0, targetWithHost(&target, "[::1%lo0]:80")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.ErrorIs(t, validateSubnetTarget(0, targetWithHost(&target, "169.254.100.100")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.ErrorIs(t, validateSubnetTarget(0, targetWithHost(&target, "fe80::1")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.ErrorIs(t, validateSubnetTarget(0, targetWithHost(&target, "[fe80::1]")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.ErrorIs(t, validateSubnetTarget(0, targetWithHost(&target, "224.100.100.100")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.ErrorIs(t, validateSubnetTarget(0, targetWithHost(&target, "ff00::ffff")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
assert.ErrorIs(t, validateSubnetTarget(0, targetWithHost(&target, "[ff00::ffff]")), ErrUnsupportedIPAddressUpstreamHost)
|
||||
|
||||
// empty host
|
||||
assert.NotNil(t, validateSubnetTarget(0, &Target{Options: TargetOptions{DirectUpstream: true}, Host: " "}))
|
||||
// host with a space
|
||||
assert.NotNil(t, validateSubnetTarget(0, &Target{Options: TargetOptions{DirectUpstream: true}, Host: "with space"}))
|
||||
// host with a tab
|
||||
assert.NotNil(t, validateSubnetTarget(0, &Target{Options: TargetOptions{DirectUpstream: true}, Host: "with\ttab"}))
|
||||
// host with a slash
|
||||
assert.NotNil(t, validateSubnetTarget(0, &Target{Options: TargetOptions{DirectUpstream: true}, Host: "with/slash"}))
|
||||
}
|
||||
|
||||
func targetWithHost(t *Target, host string) *Target {
|
||||
t.Host = host
|
||||
return t
|
||||
}
|
||||
|
||||
func TestToProtoMapping_TargetOptions(t *testing.T) {
|
||||
rp := &Service{
|
||||
ID: "svc-1",
|
||||
@@ -250,6 +308,44 @@ func TestToProtoMapping_TargetOptions(t *testing.T) {
|
||||
assert.Equal(t, int64(30), opts.RequestTimeout.Seconds)
|
||||
}
|
||||
|
||||
// TestToProtoMapping_AllowedGroupIds covers the list the proxy gates session
|
||||
// cookies on: without it the proxy can only check a cookie's signature, which
|
||||
// makes a token minted for a user outside the groups a bearer credential.
|
||||
func TestToProtoMapping_AllowedGroupIds(t *testing.T) {
|
||||
t.Run("distribution groups reach the proxy", func(t *testing.T) {
|
||||
rp := &Service{
|
||||
ID: "svc-1",
|
||||
AccountID: "acc-1",
|
||||
Domain: "example.com",
|
||||
Auth: AuthConfig{
|
||||
BearerAuth: &BearerAuthConfig{
|
||||
Enabled: true,
|
||||
DistributionGroups: []string{"grp-1", "grp-2"},
|
||||
},
|
||||
},
|
||||
}
|
||||
pm := rp.ToProtoMapping(Create, "token", proxy.OIDCValidationConfig{})
|
||||
|
||||
assert.True(t, pm.GetAuth().GetOidc())
|
||||
assert.Equal(t, []string{"grp-1", "grp-2"}, pm.GetAuth().GetAllowedGroupIds())
|
||||
})
|
||||
|
||||
t.Run("a service open to the account carries no groups", func(t *testing.T) {
|
||||
rp := &Service{
|
||||
ID: "svc-1",
|
||||
AccountID: "acc-1",
|
||||
Domain: "example.com",
|
||||
Auth: AuthConfig{
|
||||
BearerAuth: &BearerAuthConfig{Enabled: true},
|
||||
},
|
||||
}
|
||||
pm := rp.ToProtoMapping(Create, "token", proxy.OIDCValidationConfig{})
|
||||
|
||||
assert.True(t, pm.GetAuth().GetOidc())
|
||||
assert.Empty(t, pm.GetAuth().GetAllowedGroupIds(), "an empty list must not restrict access")
|
||||
})
|
||||
}
|
||||
|
||||
func TestToProtoMapping_NoOptionsWhenDefault(t *testing.T) {
|
||||
rp := &Service{
|
||||
ID: "svc-1",
|
||||
|
||||
@@ -15,6 +15,7 @@ const (
|
||||
from zones
|
||||
left join records as r on r.zone_id = zones.id
|
||||
where zones.account_id=$1 and zones.enabled
|
||||
order by zones.id
|
||||
`
|
||||
)
|
||||
|
||||
|
||||
@@ -11,9 +11,11 @@ import (
|
||||
)
|
||||
|
||||
const (
|
||||
// Outer join: a groupless router must survive.
|
||||
GetNetworkRouterQuery = `
|
||||
select public_id, peer, network_id, masquerade, metric, enabled, peer_groups, group_peers.peer_id
|
||||
from network_routers, json_each(peer_groups)
|
||||
from network_routers
|
||||
left join json_each(network_routers.peer_groups) on true
|
||||
left join group_peers on group_peers.account_id=? and group_peers.group_id=json_each.value
|
||||
where network_routers.account_id=?
|
||||
`
|
||||
|
||||
@@ -42,17 +42,20 @@ func (sc *SqliteStoreConn) GetAllowedUsers(ctx context.Context, accountId string
|
||||
userIdIdx := make(map[string]struct{})
|
||||
groupIdToUserIds := make(map[string][]string)
|
||||
for _, user := range users {
|
||||
for _, allgid := range allGroupIds {
|
||||
groupIdToUserIds[allgid] = append(groupIdToUserIds[allgid], user.ID)
|
||||
}
|
||||
userIdIdx[user.ID] = struct{}{}
|
||||
autogroups := make([]string, 0)
|
||||
if user.AutoGroups == nil {
|
||||
continue
|
||||
}
|
||||
if err := json.Unmarshal(user.AutoGroups, &autogroups); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
userIdIdx[user.ID] = struct{}{}
|
||||
for _, groupId := range autogroups {
|
||||
groupIdToUserIds[groupId] = append(groupIdToUserIds[groupId], user.ID)
|
||||
}
|
||||
for _, allgid := range allGroupIds {
|
||||
groupIdToUserIds[allgid] = append(groupIdToUserIds[allgid], user.ID)
|
||||
}
|
||||
}
|
||||
|
||||
return userIdIdx, groupIdToUserIds, nil
|
||||
|
||||
@@ -21,8 +21,6 @@ import (
|
||||
"google.golang.org/grpc/credentials"
|
||||
"google.golang.org/grpc/keepalive"
|
||||
|
||||
cachestore "github.com/eko/gocache/lib/v4/store"
|
||||
|
||||
"github.com/netbirdio/netbird/encryption"
|
||||
"github.com/netbirdio/netbird/formatter/hook"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
@@ -75,8 +73,8 @@ func (s *BaseServer) Metrics() telemetry.AppMetrics {
|
||||
|
||||
// CacheStore returns a shared cache store backed by Redis or in-memory depending on the environment.
|
||||
// All consumers should reuse this store to avoid creating multiple Redis connections.
|
||||
func (s *BaseServer) CacheStore() cachestore.StoreInterface {
|
||||
return Create(s, func() cachestore.StoreInterface {
|
||||
func (s *BaseServer) CacheStore() nbcache.Store {
|
||||
return Create(s, func() nbcache.Store {
|
||||
cs, err := nbcache.NewStore(context.Background(), nbcache.DefaultStoreMaxTimeout, nbcache.DefaultStoreCleanupInterval, nbcache.DefaultStoreMaxConn)
|
||||
if err != nil {
|
||||
log.Fatalf("failed to create shared cache store: %v", err)
|
||||
|
||||
@@ -5,22 +5,23 @@ import (
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/eko/gocache/lib/v4/cache"
|
||||
"github.com/eko/gocache/lib/v4/store"
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
nbcache "github.com/netbirdio/netbird/management/server/cache"
|
||||
)
|
||||
|
||||
// PKCEVerifierStore manages PKCE verifiers for OAuth flows.
|
||||
// Supports both in-memory and Redis storage via NB_IDP_CACHE_REDIS_ADDRESS env var.
|
||||
type PKCEVerifierStore struct {
|
||||
cache *cache.Cache[string]
|
||||
cache nbcache.Store
|
||||
ctx context.Context
|
||||
}
|
||||
|
||||
// NewPKCEVerifierStore creates a PKCE verifier store using the provided shared cache store.
|
||||
func NewPKCEVerifierStore(ctx context.Context, cacheStore store.StoreInterface) *PKCEVerifierStore {
|
||||
func NewPKCEVerifierStore(ctx context.Context, cacheStore nbcache.Store) *PKCEVerifierStore {
|
||||
return &PKCEVerifierStore{
|
||||
cache: cache.New[string](cacheStore),
|
||||
cache: cacheStore,
|
||||
ctx: ctx,
|
||||
}
|
||||
}
|
||||
@@ -40,14 +41,14 @@ func (s *PKCEVerifierStore) Store(state, verifier string, ttl time.Duration) err
|
||||
// Returns the verifier and true if found, or empty string and false if not found.
|
||||
// This enforces single-use semantics for PKCE verifiers.
|
||||
func (s *PKCEVerifierStore) LoadAndDelete(state string) (string, bool) {
|
||||
verifier, err := s.cache.Get(s.ctx, state)
|
||||
verifier, found, err := s.cache.GetDel(s.ctx, state)
|
||||
if err != nil {
|
||||
log.Debugf("PKCE verifier not found for state")
|
||||
log.Warnf("Failed to consume PKCE verifier: %v", err)
|
||||
return "", false
|
||||
}
|
||||
|
||||
if err := s.cache.Delete(s.ctx, state); err != nil {
|
||||
log.Warnf("Failed to delete PKCE verifier for state: %v", err)
|
||||
if !found {
|
||||
log.Debug("PKCE verifier not found for state")
|
||||
return "", false
|
||||
}
|
||||
|
||||
return verifier, true
|
||||
|
||||
@@ -0,0 +1,85 @@
|
||||
package grpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestPKCEVerifierStoreLoadAndDelete(t *testing.T) {
|
||||
const (
|
||||
state = "state"
|
||||
verifier = "verifier"
|
||||
attempts = 64
|
||||
)
|
||||
|
||||
t.Run("exactly one concurrent caller consumes the verifier", func(t *testing.T) {
|
||||
store := NewPKCEVerifierStore(context.Background(), testCacheStore(t))
|
||||
if err := store.Store(state, verifier, time.Minute); err != nil {
|
||||
t.Fatalf("couldn't store PKCE verifier: %s", err)
|
||||
}
|
||||
|
||||
start := make(chan struct{})
|
||||
type result struct {
|
||||
verifier string
|
||||
found bool
|
||||
}
|
||||
results := make(chan result, attempts)
|
||||
for range attempts {
|
||||
go func() {
|
||||
<-start
|
||||
verifier, found := store.LoadAndDelete(state)
|
||||
results <- result{verifier: verifier, found: found}
|
||||
}()
|
||||
}
|
||||
close(start)
|
||||
|
||||
winners := 0
|
||||
for range attempts {
|
||||
result := <-results
|
||||
if result.found {
|
||||
winners++
|
||||
if result.verifier != verifier {
|
||||
t.Fatalf("unexpected verifier: got %q, expected %q", result.verifier, verifier)
|
||||
}
|
||||
}
|
||||
}
|
||||
if winners != 1 {
|
||||
t.Fatalf("expected exactly one PKCE verifier consumer, got %d", winners)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("replayed state is rejected", func(t *testing.T) {
|
||||
store := NewPKCEVerifierStore(context.Background(), testCacheStore(t))
|
||||
if err := store.Store(state, verifier, time.Minute); err != nil {
|
||||
t.Fatalf("couldn't store PKCE verifier: %s", err)
|
||||
}
|
||||
|
||||
if got, found := store.LoadAndDelete(state); !found || got != verifier {
|
||||
t.Fatalf("first load should return the verifier, got %q, found %t", got, found)
|
||||
}
|
||||
if got, found := store.LoadAndDelete(state); found {
|
||||
t.Fatalf("replayed state should not resolve, got %q", got)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("unknown state is rejected", func(t *testing.T) {
|
||||
store := NewPKCEVerifierStore(context.Background(), testCacheStore(t))
|
||||
|
||||
if got, found := store.LoadAndDelete("never-stored"); found {
|
||||
t.Fatalf("unknown state should not resolve, got %q", got)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("expired verifier is rejected", func(t *testing.T) {
|
||||
store := NewPKCEVerifierStore(context.Background(), testCacheStore(t))
|
||||
if err := store.Store(state, verifier, 50*time.Millisecond); err != nil {
|
||||
t.Fatalf("couldn't store PKCE verifier: %s", err)
|
||||
}
|
||||
|
||||
time.Sleep(100 * time.Millisecond)
|
||||
if got, found := store.LoadAndDelete(state); found {
|
||||
t.Fatalf("expired verifier should not resolve, got %q", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -1651,6 +1651,10 @@ var (
|
||||
// ErrUserBlocked reports a blocked user, who may not hold a proxy session.
|
||||
ErrUserBlocked = errors.New("user blocked")
|
||||
|
||||
// ErrUserNotInGroup reports a user outside the service's distribution
|
||||
// groups, who may not hold a proxy session for it.
|
||||
ErrUserNotInGroup = errors.New("user not in allowed groups")
|
||||
|
||||
errUserUnresolved = errors.New("user could not be resolved")
|
||||
)
|
||||
|
||||
@@ -1689,8 +1693,10 @@ func sameAccount(userAccountID, serviceAccountID string) bool {
|
||||
// GenerateSessionToken creates a signed session JWT for the given domain and
|
||||
// user. The user's group memberships are embedded in the token so policy-aware
|
||||
// middlewares on the proxy can authorise without an extra management round-trip.
|
||||
// A user the store cannot resolve, or whose account is pending approval or
|
||||
// blocked, gets no token at all, so the browser never receives a session cookie.
|
||||
// A user the store cannot resolve, whose account is pending approval or blocked,
|
||||
// or who is outside the service's distribution groups, gets no token at all: the
|
||||
// token is a bearer credential for the service, so authorisation has to run
|
||||
// before it is signed rather than only when the proxy presents it back.
|
||||
func (s *ProxyServiceServer) GenerateSessionToken(ctx context.Context, domain, userID string, method proxyauth.Method) (string, error) {
|
||||
service, err := s.getServiceByDomain(ctx, domain)
|
||||
if err != nil {
|
||||
@@ -1726,6 +1732,14 @@ func (s *ProxyServiceServer) GenerateSessionToken(ctx context.Context, domain, u
|
||||
return "", fmt.Errorf("session token for user %s: %w", userID, err)
|
||||
}
|
||||
|
||||
if err := s.checkGroupAccess(service, user); err != nil {
|
||||
log.WithContext(ctx).WithFields(log.Fields{
|
||||
"domain": domain,
|
||||
"user_id": userID,
|
||||
}).Debug("GenerateSessionToken: user not in the service's distribution groups")
|
||||
return "", fmt.Errorf("session token for user %s: %w", userID, ErrUserNotInGroup)
|
||||
}
|
||||
|
||||
groupIDs, groupNames := pairGroupIDsAndNames(userGroups)
|
||||
|
||||
token, err := sessionkey.SignToken(
|
||||
|
||||
@@ -9,7 +9,6 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
cachestore "github.com/eko/gocache/lib/v4/store"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/grpc/codes"
|
||||
@@ -21,7 +20,7 @@ import (
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
func testCacheStore(t *testing.T) cachestore.StoreInterface {
|
||||
func testCacheStore(t *testing.T) nbcache.Store {
|
||||
t.Helper()
|
||||
s, err := nbcache.NewStore(context.Background(), 30*time.Minute, 10*time.Minute, 100)
|
||||
require.NoError(t, err)
|
||||
|
||||
@@ -247,17 +247,8 @@ func (s *Server) Sync(req *proto.EncryptedMessage, srv proto.ManagementService_S
|
||||
sRealIP := realIP.String()
|
||||
peerMeta := extractPeerMeta(ctx, syncReq.GetMeta())
|
||||
|
||||
userID, err := s.accountManager.GetUserIDByPeerKey(ctx, peerKey.String())
|
||||
if err != nil {
|
||||
s.syncSem.Add(-1)
|
||||
if errStatus, ok := internalStatus.FromError(err); ok && errStatus.Type() == internalStatus.NotFound {
|
||||
return status.Errorf(codes.PermissionDenied, "peer is not registered")
|
||||
}
|
||||
return mapError(ctx, err)
|
||||
}
|
||||
|
||||
metahashed := metaHash(peerMeta)
|
||||
if userID == "" && !s.loginFilter.allowLogin(peerKey.String(), metahashed) {
|
||||
if !s.loginFilter.allowLogin(peerKey.String(), metahashed) {
|
||||
if s.appMetrics != nil {
|
||||
s.appMetrics.GRPCMetrics().CountSyncRequestBlocked()
|
||||
}
|
||||
|
||||
@@ -431,6 +431,57 @@ func TestValidateSession_MissingToken(t *testing.T) {
|
||||
assert.Contains(t, resp.DeniedReason, "missing")
|
||||
}
|
||||
|
||||
// TestGenerateSessionToken_UserNotInAllowedGroupGetsNoToken is the regression
|
||||
// guard for the group-authorisation bypass: the callback used to hand a signed
|
||||
// token to a user the service denies, and the proxy honoured that token as soon
|
||||
// as the user moved it into the nb_session cookie themselves. Authorisation has
|
||||
// to run before the token is signed.
|
||||
func TestGenerateSessionToken_UserNotInAllowedGroupGetsNoToken(t *testing.T) {
|
||||
setup := setupValidateSessionTest(t)
|
||||
defer setup.cleanup()
|
||||
|
||||
token, err := setup.proxyService.GenerateSessionToken(context.Background(), "restricted-proxy.example.com", "nonGroupUserId", auth.MethodOIDC)
|
||||
|
||||
require.Error(t, err, "a user outside the distribution groups must not receive a token")
|
||||
assert.ErrorIs(t, err, ErrUserNotInGroup, "the callback maps this sentinel onto the access denied page")
|
||||
assert.Empty(t, token, "no token may reach the browser")
|
||||
}
|
||||
|
||||
func TestGenerateSessionToken_UserInAllowedGroupGetsTokenWithGroups(t *testing.T) {
|
||||
setup := setupValidateSessionTest(t)
|
||||
defer setup.cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
svc, err := setup.store.GetServiceByID(ctx, store.LockingStrengthNone, "testAccountId", "restrictedProxyId")
|
||||
require.NoError(t, err)
|
||||
|
||||
token, err := setup.proxyService.GenerateSessionToken(ctx, "restricted-proxy.example.com", "allowedUserId", auth.MethodOIDC)
|
||||
require.NoError(t, err)
|
||||
require.NotEmpty(t, token)
|
||||
|
||||
pubKey, err := base64.StdEncoding.DecodeString(svc.SessionPublicKey)
|
||||
require.NoError(t, err)
|
||||
|
||||
userID, _, method, groups, _, err := auth.ValidateSessionJWT(token, "restricted-proxy.example.com", pubKey)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "allowedUserId", userID)
|
||||
assert.Equal(t, auth.MethodOIDC.String(), method)
|
||||
assert.Equal(t, []string{"allowedGroupId"}, groups, "the proxy gates the cookie on this claim, so it must carry the matched group")
|
||||
}
|
||||
|
||||
// TestGenerateSessionToken_UnrestrictedServiceAllowsAnyAccountUser keeps the new
|
||||
// gate scoped: a service without distribution groups is open to every user of
|
||||
// its account, as before.
|
||||
func TestGenerateSessionToken_UnrestrictedServiceAllowsAnyAccountUser(t *testing.T) {
|
||||
setup := setupValidateSessionTest(t)
|
||||
defer setup.cleanup()
|
||||
|
||||
token, err := setup.proxyService.GenerateSessionToken(context.Background(), "test-proxy.example.com", "nonGroupUserId", auth.MethodOIDC)
|
||||
|
||||
require.NoError(t, err, "an unrestricted service must keep working for any user of the account")
|
||||
assert.NotEmpty(t, token)
|
||||
}
|
||||
|
||||
type testValidateSessionServiceManager struct {
|
||||
store store.Store
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user