Merge branch 'main' into fix/pkce-flow-session-extend

# Conflicts:
#	client/ios/NetBirdSDK/login.go
#	client/server/server.go
#	shared/management/proto/management.pb.go
This commit is contained in:
Zoltán Papp
2026-10-05 15:37:54 +02:00
847 changed files with 67577 additions and 25637 deletions
+121
View File
@@ -0,0 +1,121 @@
package db
import (
"context"
"fmt"
"os"
"runtime"
"strconv"
"time"
"github.com/jackc/pgx/v5/pgxpool"
log "github.com/sirupsen/logrus"
"gorm.io/gorm"
)
const (
defaultTransactionTimeout = 5 * time.Minute
connMaxLifetime = time.Hour
connMaxIdleTime = 3 * time.Minute
)
// TxMetrics receives the duration of every committed top-level transaction.
type TxMetrics interface {
CountTransactionDuration(duration time.Duration)
}
// Conn is the database connection shared by all repositories: one gorm handle,
// the pgx pool of a Postgres deployment and the engine they talk to.
type Conn struct {
db *gorm.DB
pool *pgxpool.Pool
engine Engine
txTimeout time.Duration
metrics TxMetrics
}
// NewConn takes ownership of an open gorm handle and pool once it returns
// without error, applying the connection limits and transaction timeout
// configured through the environment.
func NewConn(ctx context.Context, gormDB *gorm.DB, engine Engine, pool *pgxpool.Pool) (*Conn, error) {
sqlDB, err := gormDB.DB()
if err != nil {
return nil, err
}
txTimeout := defaultTransactionTimeout
if v := os.Getenv("NB_STORE_TRANSACTION_TIMEOUT"); v != "" {
if parsed, err := time.ParseDuration(v); err == nil {
txTimeout = parsed
}
}
log.WithContext(ctx).Infof("Setting transaction timeout to %v", txTimeout)
conns := runtime.NumCPU()
configuredConns, err := strconv.Atoi(os.Getenv("NB_SQL_MAX_OPEN_CONNS"))
connsConfigured := err == nil
if connsConfigured {
conns = configuredConns
}
if engine == SqliteStoreEngine {
if connsConfigured {
log.WithContext(ctx).Warnf("setting NB_SQL_MAX_OPEN_CONNS is not supported for sqlite, using default value 1")
}
conns = 1
}
sqlDB.SetMaxOpenConns(conns)
sqlDB.SetMaxIdleConns(conns)
sqlDB.SetConnMaxLifetime(connMaxLifetime)
sqlDB.SetConnMaxIdleTime(connMaxIdleTime)
log.WithContext(ctx).Infof("Set max open db connections to %d, max idle to %d, max lifetime to %v, max idle time to %v",
conns, conns, connMaxLifetime, connMaxIdleTime)
return &Conn{db: gormDB, pool: pool, engine: engine, txTimeout: txTimeout}, nil
}
// DB returns the handle a query must run on: the transaction when tx is set,
// otherwise the shared connection.
func (c *Conn) DB(tx *Tx) *gorm.DB {
if tx != nil {
return tx.db
}
return c.db
}
// Pool returns the pgx pool for read paths that bypass gorm. It is nil on
// engines other than Postgres and inside a transaction, where the pool would
// not see the uncommitted writes.
func (c *Conn) Pool(tx *Tx) *pgxpool.Pool {
if tx != nil {
return nil
}
return c.pool
}
func (c *Conn) Engine() Engine {
return c.engine
}
// SetTxMetrics registers the sink that receives transaction durations.
func (c *Conn) SetTxMetrics(metrics TxMetrics) {
c.metrics = metrics
}
// AutoMigrate creates or updates the tables of the given models.
func (c *Conn) AutoMigrate(models ...any) error {
return c.db.AutoMigrate(models...)
}
// Close releases the gorm connection and the pgx pool.
func (c *Conn) Close() error {
if c.pool != nil {
c.pool.Close()
}
sqlDB, err := c.db.DB()
if err != nil {
return fmt.Errorf("get db: %w", err)
}
return sqlDB.Close()
}
+148
View File
@@ -0,0 +1,148 @@
package db
import (
"context"
"errors"
"path/filepath"
"testing"
"time"
"github.com/jackc/pgx/v5/pgxpool"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
)
type testRow struct {
ID uint `gorm:"primaryKey"`
Name string
}
func openTestConn(t *testing.T) *Conn {
t.Helper()
conn, err := OpenSqliteFile(context.Background(), t.TempDir(), SqliteFileName)
require.NoError(t, err)
t.Cleanup(func() { require.NoError(t, conn.Close()) })
require.NoError(t, conn.AutoMigrate(&testRow{}))
return conn
}
func countRows(t *testing.T, conn *Conn) int64 {
t.Helper()
var count int64
require.NoError(t, conn.DB(nil).Model(&testRow{}).Count(&count).Error)
return count
}
func TestNewConn_ReadsTransactionTimeoutFromEnv(t *testing.T) {
t.Setenv("NB_STORE_TRANSACTION_TIMEOUT", "1s")
conn := openTestConn(t)
assert.Equal(t, time.Second, conn.txTimeout)
assert.Equal(t, SqliteStoreEngine, conn.Engine())
}
func TestRunInTx_CommitsOnSuccess(t *testing.T) {
conn := openTestConn(t)
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
return conn.DB(tx).Create(&testRow{Name: "a"}).Error
})
require.NoError(t, err)
assert.EqualValues(t, 1, countRows(t, conn))
}
func TestRunInTx_RollsBackOnError(t *testing.T) {
conn := openTestConn(t)
failure := errors.New("boom")
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
require.NoError(t, conn.DB(tx).Create(&testRow{Name: "a"}).Error)
return failure
})
require.ErrorIs(t, err, failure)
assert.EqualValues(t, 0, countRows(t, conn))
}
func TestRunInTx_RollsBackOnPanic(t *testing.T) {
conn := openTestConn(t)
require.Panics(t, func() {
_ = conn.RunInTx(context.Background(), func(tx *Tx) error {
require.NoError(t, conn.DB(tx).Create(&testRow{Name: "a"}).Error)
panic("boom")
})
})
assert.EqualValues(t, 0, countRows(t, conn))
}
func TestRunInTx_FailsWhenTimeoutExceeded(t *testing.T) {
t.Setenv("NB_STORE_TRANSACTION_TIMEOUT", "50ms")
conn := openTestConn(t)
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
time.Sleep(100 * time.Millisecond)
return conn.DB(tx).Create(&testRow{Name: "a"}).Error
})
require.ErrorIs(t, err, context.DeadlineExceeded)
assert.EqualValues(t, 0, countRows(t, conn))
}
func TestRunInTx_ReportsDurationToMetrics(t *testing.T) {
conn := openTestConn(t)
metrics := &recordingMetrics{}
conn.SetTxMetrics(metrics)
require.NoError(t, conn.RunInTx(context.Background(), func(*Tx) error { return nil }))
assert.Equal(t, 1, metrics.calls)
}
func TestConn_DBSelectsTransactionHandle(t *testing.T) {
conn := openTestConn(t)
assert.Same(t, conn.db, conn.DB(nil))
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
assert.Same(t, tx.db, conn.DB(tx))
assert.NotSame(t, conn.db, conn.DB(tx))
return nil
})
require.NoError(t, err)
}
func TestConn_PoolIsUnavailableInsideTransaction(t *testing.T) {
conn := openTestConn(t)
conn.pool = &pgxpool.Pool{}
defer func() { conn.pool = nil }()
assert.Same(t, conn.pool, conn.Pool(nil))
err := conn.RunInTx(context.Background(), func(tx *Tx) error {
assert.Nil(t, conn.Pool(tx))
return nil
})
require.NoError(t, err)
}
type recordingMetrics struct {
calls int
}
func (m *recordingMetrics) CountTransactionDuration(time.Duration) {
m.calls++
}
func TestNewConn_MaxOpenConnsFromEnv(t *testing.T) {
t.Setenv("NB_SQL_MAX_OPEN_CONNS", "7")
gormDB, err := gorm.Open(sqlite.Open(filepath.Join(t.TempDir(), "store.db")), GormConfig())
require.NoError(t, err)
conn, err := NewConn(context.Background(), gormDB, PostgresStoreEngine, nil)
require.NoError(t, err)
t.Cleanup(func() { _ = conn.Close() })
sqlDB, err := conn.DB(nil).DB()
require.NoError(t, err)
assert.Equal(t, 7, sqlDB.Stats().MaxOpenConnections)
sqliteDB, err := openTestConn(t).DB(nil).DB()
require.NoError(t, err)
assert.Equal(t, 1, sqliteDB.Stats().MaxOpenConnections)
}
@@ -0,0 +1,23 @@
package dbtest
import (
"context"
"testing"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/shared/db"
)
// NewConn opens a fresh SQLite database in a temporary directory, migrates the
// given models and closes the connection when the test ends. It ignores
// NB_STORE_ENGINE_SQLITE_FILE, so a developer's configured database is never
// touched, and is safe to call from parallel tests.
func NewConn(t testing.TB, models ...any) *db.Conn {
t.Helper()
conn, err := db.OpenSqliteFile(context.Background(), t.TempDir(), db.SqliteFileName)
require.NoError(t, err)
t.Cleanup(func() { _ = conn.Close() })
require.NoError(t, conn.AutoMigrate(models...))
return conn
}
@@ -0,0 +1,31 @@
package dbtest
import (
"os"
"path/filepath"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"github.com/netbirdio/netbird/management/internals/shared/db"
)
func TestNewConn_IgnoresSqliteFileOverride(t *testing.T) {
override := filepath.Join(t.TempDir(), "configured.db")
t.Setenv("NB_STORE_ENGINE_SQLITE_FILE", override)
conn := NewConn(t)
assert.Equal(t, db.SqliteStoreEngine, conn.Engine())
_, err := os.Stat(override)
require.ErrorIs(t, err, os.ErrNotExist)
}
func TestNewConn_Parallel(t *testing.T) {
t.Parallel()
conn := NewConn(t)
assert.Equal(t, db.SqliteStoreEngine, conn.Engine())
}
+10
View File
@@ -0,0 +1,10 @@
package db
// Engine identifies the SQL engine behind a Conn.
type Engine string
const (
SqliteStoreEngine Engine = "sqlite"
PostgresStoreEngine Engine = "postgres"
MysqlStoreEngine Engine = "mysql"
)
+12
View File
@@ -0,0 +1,12 @@
package db
// LockingStrength is the row lock a query holds until its transaction ends.
type LockingStrength string
const (
LockingStrengthUpdate LockingStrength = "UPDATE"
LockingStrengthShare LockingStrength = "SHARE"
LockingStrengthNoKeyUpdate LockingStrength = "NO KEY UPDATE"
LockingStrengthKeyShare LockingStrength = "KEY SHARE"
LockingStrengthNone LockingStrength = "NONE"
)
+176
View File
@@ -0,0 +1,176 @@
package db
import (
"context"
"fmt"
"net/url"
"os"
"path/filepath"
"runtime"
"strings"
"time"
"github.com/jackc/pgx/v5/pgxpool"
"gorm.io/driver/mysql"
"gorm.io/driver/postgres"
"gorm.io/driver/sqlite"
"gorm.io/gorm"
"gorm.io/gorm/logger"
)
// SqliteFileName is the default SQLite database file inside the data directory.
const SqliteFileName = "store.db"
// PoolConfig sizes the pgx pool a Postgres deployment uses for the read paths
// that bypass gorm.
type PoolConfig struct {
MaxConns int32
MinConns int32
MaxConnLifetime time.Duration
HealthCheckPeriod time.Duration
}
var DefaultPoolConfig = PoolConfig{
MaxConns: 30,
MinConns: 1,
MaxConnLifetime: 60 * time.Minute,
HealthCheckPeriod: time.Minute,
}
// GormConfig is the configuration every engine is opened with.
func GormConfig() *gorm.Config {
return &gorm.Config{
Logger: logger.Default.LogMode(logger.Silent),
CreateBatchSize: 400,
}
}
// OpenSqlite opens the SQLite database in dataDir, or the file named by
// NB_STORE_ENGINE_SQLITE_FILE.
func OpenSqlite(ctx context.Context, dataDir string) (*Conn, error) {
storeFile := SqliteFileName
if envFile, ok := os.LookupEnv("NB_STORE_ENGINE_SQLITE_FILE"); ok && envFile != "" {
storeFile = envFile
}
return OpenSqliteFile(ctx, dataDir, storeFile)
}
// OpenSqliteFile opens the SQLite database storeFile, resolved against dataDir
// when relative. storeFile may carry SQLite URI query parameters.
func OpenSqliteFile(ctx context.Context, dataDir, storeFile string) (*Conn, error) {
// Separate file path from any SQLite URI query parameters (e.g., "store.db?mode=rwc")
filePath, query, hasQuery := strings.Cut(storeFile, "?")
connStr := filePath
if !filepath.IsAbs(filePath) {
connStr = filepath.Join(dataDir, filePath)
}
// Compose query parameters. User-provided ?_busy_timeout (or its mattn alias
// ?_timeout) overrides our default; otherwise inject 30s so SQLite waits at
// most that long on a lock instead of blocking the only Go-side connection.
// mattn/go-sqlite3 applies PRAGMA from the DSN on every fresh connection, so
// the value survives ConnMaxIdleTime/ConnMaxLifetime recycling. cache=shared
// stays the default on non-Windows for the same reason as before.
parsed, _ := url.ParseQuery(query)
var defaults []string
if parsed.Get("_busy_timeout") == "" && parsed.Get("_timeout") == "" {
defaults = append(defaults, "_busy_timeout=30000")
}
if !hasQuery && runtime.GOOS != "windows" {
// To avoid `The process cannot access the file because it is being used by another process` on Windows
defaults = append(defaults, "cache=shared")
}
parts := defaults
if hasQuery {
parts = append(parts, query)
}
if len(parts) > 0 {
connStr += "?" + strings.Join(parts, "&")
}
gormDB, err := gorm.Open(sqlite.Open(connStr), GormConfig())
if err != nil {
return nil, err
}
conn, err := NewConn(ctx, gormDB, SqliteStoreEngine, nil)
if err != nil {
closeGorm(gormDB)
return nil, err
}
return conn, nil
}
// OpenPostgres opens a Postgres database through gorm and a pgx pool sized by pool.
func OpenPostgres(ctx context.Context, dsn string, pool PoolConfig) (*Conn, error) {
gormDB, err := gorm.Open(postgres.Open(dsn), GormConfig())
if err != nil {
return nil, err
}
pgxPool, err := newPgxPool(ctx, dsn, pool)
if err != nil {
closeGorm(gormDB)
return nil, err
}
conn, err := NewConn(ctx, gormDB, PostgresStoreEngine, pgxPool)
if err != nil {
pgxPool.Close()
closeGorm(gormDB)
return nil, err
}
return conn, nil
}
// MysqlDSN adds the connection parameters every MySQL handle needs, keeping
// the options already present in dsn.
func MysqlDSN(dsn string) string {
separator := "?"
if strings.Contains(dsn, "?") {
separator = "&"
}
return dsn + separator + "charset=utf8&parseTime=True&loc=Local"
}
// OpenMysql opens a MySQL database through gorm.
func OpenMysql(ctx context.Context, dsn string) (*Conn, error) {
gormDB, err := gorm.Open(mysql.Open(MysqlDSN(dsn)), GormConfig())
if err != nil {
return nil, err
}
conn, err := NewConn(ctx, gormDB, MysqlStoreEngine, nil)
if err != nil {
closeGorm(gormDB)
return nil, err
}
return conn, nil
}
func newPgxPool(ctx context.Context, dsn string, cfg PoolConfig) (*pgxpool.Pool, error) {
config, err := pgxpool.ParseConfig(dsn)
if err != nil {
return nil, fmt.Errorf("unable to parse database config: %w", err)
}
config.MaxConns = cfg.MaxConns
config.MinConns = cfg.MinConns
config.MaxConnLifetime = cfg.MaxConnLifetime
config.HealthCheckPeriod = cfg.HealthCheckPeriod
pool, err := pgxpool.NewWithConfig(ctx, config)
if err != nil {
return nil, fmt.Errorf("unable to create connection pool: %w", err)
}
if err := pool.Ping(ctx); err != nil {
pool.Close()
return nil, fmt.Errorf("unable to ping database: %w", err)
}
return pool, nil
}
func closeGorm(gormDB *gorm.DB) {
if sqlDB, err := gormDB.DB(); err == nil {
_ = sqlDB.Close()
}
}
@@ -0,0 +1,12 @@
package db
import (
"testing"
"github.com/stretchr/testify/assert"
)
func TestMysqlDSN(t *testing.T) {
assert.Equal(t, "user:pw@tcp(host:3306)/db?charset=utf8&parseTime=True&loc=Local", MysqlDSN("user:pw@tcp(host:3306)/db"))
assert.Equal(t, "user:pw@tcp(host:3306)/db?tls=true&charset=utf8&parseTime=True&loc=Local", MysqlDSN("user:pw@tcp(host:3306)/db?tls=true"))
}
@@ -0,0 +1,105 @@
package db
import (
"context"
"errors"
"fmt"
"runtime/debug"
"time"
log "github.com/sirupsen/logrus"
"gorm.io/gorm"
)
// Tx is an open transaction handed to repository calls; nil means autocommit.
type Tx struct {
db *gorm.DB
}
// RunInTx runs fn in one transaction that commits when fn returns nil and rolls
// back otherwise, bounded by the configured transaction timeout.
func (c *Conn) RunInTx(ctx context.Context, fn func(tx *Tx) error) error {
timeoutCtx, cancel := context.WithTimeout(ctx, c.txTimeout)
defer cancel()
startTime := time.Now()
tx := c.db.WithContext(timeoutCtx).Begin()
if tx.Error != nil {
return tx.Error
}
defer func() {
if r := recover(); r != nil {
tx.Rollback()
panic(r)
}
}()
if err := c.applyStatementTimeouts(tx); err != nil {
tx.Rollback()
return err
}
err := c.withForeignKeyChecksDisabled(tx, func() error {
return fn(&Tx{db: tx})
})
if err != nil {
tx.Rollback()
c.logIfTimedOut(ctx, timeoutCtx, err, "transaction", startTime)
return err
}
if err := tx.Commit().Error; err != nil {
c.logIfTimedOut(ctx, timeoutCtx, err, "transaction commit", startTime)
return err
}
log.WithContext(ctx).Tracef("transaction took %v", time.Since(startTime))
if c.metrics != nil {
c.metrics.CountTransactionDuration(time.Since(startTime))
}
return nil
}
func (c *Conn) applyStatementTimeouts(tx *gorm.DB) error {
if c.engine != PostgresStoreEngine {
return nil
}
if err := tx.Exec("SET LOCAL statement_timeout = '1min'").Error; err != nil {
return fmt.Errorf("failed to set statement timeout: %w", err)
}
if err := tx.Exec("SET LOCAL lock_timeout = '1min'").Error; err != nil {
return fmt.Errorf("failed to set lock timeout: %w", err)
}
return nil
}
// withForeignKeyChecksDisabled runs fn with MySQL's FK checks off, which avoids
// deadlocks on MySQL and Aurora without needing SUPER privilege. The setting is
// session-scoped and survives a rollback, so it is turned back on whenever fn
// returns or panics; otherwise the pooled connection would keep it disabled.
func (c *Conn) withForeignKeyChecksDisabled(tx *gorm.DB, fn func() error) (err error) {
if c.engine != MysqlStoreEngine {
return fn()
}
if err := tx.Exec("SET FOREIGN_KEY_CHECKS = 0").Error; err != nil {
return fmt.Errorf("failed to disable FK checks: %w", err)
}
defer func() {
restoreErr := tx.Exec("SET FOREIGN_KEY_CHECKS = 1").Error
if restoreErr == nil {
return
}
if err == nil {
err = fmt.Errorf("failed to re-enable FK checks: %w", restoreErr)
return
}
log.WithContext(tx.Statement.Context).Warnf("failed to re-enable FK checks after failed transaction: %v", restoreErr)
}()
return fn()
}
func (c *Conn) logIfTimedOut(ctx, timeoutCtx context.Context, err error, phase string, startTime time.Time) {
if errors.Is(err, context.DeadlineExceeded) || errors.Is(timeoutCtx.Err(), context.DeadlineExceeded) {
log.WithContext(ctx).Warnf("%s exceeded %s timeout after %v, stack: %s", phase, c.txTimeout, time.Since(startTime), debug.Stack())
}
}
@@ -7,7 +7,6 @@ import (
"github.com/netbirdio/netbird/client/ssh/auth"
nbconfig "github.com/netbirdio/netbird/management/internals/server/config"
"github.com/netbirdio/netbird/management/server/posture"
"github.com/netbirdio/netbird/management/server/types"
sharedgrpc "github.com/netbirdio/netbird/shared/management/grpc"
"github.com/netbirdio/netbird/shared/management/networkmap"
@@ -37,7 +36,7 @@ func ToComponentSyncResponse(
components *types.NetworkMapComponents,
proxyPatch *types.NetworkMap,
dnsName string,
checks []*posture.Checks,
checks []*nmdata.PostureChecks,
settings *nmdata.AccountSettingsInfo,
extraSettings *types.ExtraSettings,
peerGroups []string,
@@ -18,7 +18,6 @@ import (
"github.com/netbirdio/netbird/management/internals/controllers/network_map/controller/cache"
nbconfig "github.com/netbirdio/netbird/management/internals/server/config"
"github.com/netbirdio/netbird/management/server/posture"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/management/networkmap"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
@@ -154,7 +153,7 @@ func toPeerConfig(peer *nmdata.Peer, network *nmdata.Network, dnsName string, se
return peerConfig
}
func ToSyncResponse(ctx context.Context, config *nbconfig.Config, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, peer *nmdata.Peer, turnCredentials *Token, relayCredentials *Token, networkMap *types.NetworkMap, dnsName string, checks []*posture.Checks, dnsCache *cache.DNSConfigCache, settings *nmdata.AccountSettingsInfo, extraSettings *types.ExtraSettings, peerGroups []string, dnsFwdPort int64) *proto.SyncResponse {
func ToSyncResponse(ctx context.Context, config *nbconfig.Config, httpConfig *nbconfig.HttpServerConfig, deviceFlowConfig *nbconfig.DeviceAuthorizationFlow, peer *nmdata.Peer, turnCredentials *Token, relayCredentials *Token, networkMap *types.NetworkMap, dnsName string, checks []*nmdata.PostureChecks, dnsCache *cache.DNSConfigCache, settings *nmdata.AccountSettingsInfo, extraSettings *types.ExtraSettings, peerGroups []string, dnsFwdPort int64) *proto.SyncResponse {
// IPv6 data in AllowedIPs and SourcePrefixes wildcard expansion depends on
// whether the target peer supports IPv6. Routes and firewall rules are already
// filtered at the source (network map builder).
@@ -14,6 +14,7 @@ const (
baseBlockDuration = 10 * time.Minute // Duration for which a peer is banned after exceeding the reconnection limit
reconnLimitForBan = 30 // Number of reconnections within the reconnTreshold that triggers a ban
metaChangeLimit = 5 // Number of reconnections with different metadata that triggers a ban of one peer
maxBanLevel = 6 // Highest ban level; the ban duration doubles per level up to this one
)
type lfConfig struct {
@@ -21,6 +22,7 @@ type lfConfig struct {
baseBlockDuration time.Duration
reconnLimitForBan int
metaChangeLimit int
maxBanLevel int
}
func initCfg() *lfConfig {
@@ -29,6 +31,7 @@ func initCfg() *lfConfig {
baseBlockDuration: baseBlockDuration,
reconnLimitForBan: reconnLimitForBan,
metaChangeLimit: metaChangeLimit,
maxBanLevel: maxBanLevel,
}
}
@@ -102,11 +105,18 @@ func (l *loginFilter) addLogin(wgPubKey string, metaHash uint64) {
return
}
if state.isBanned && now.After(state.banExpiresAt) {
if state.isBanned {
if now.Before(state.banExpiresAt) {
return
}
state.isBanned = false
}
if state.banLevel > 0 && now.Sub(state.lastSeen) > (2*l.cfg.baseBlockDuration) {
quietSince := state.lastSeen
if state.banExpiresAt.After(quietSince) {
quietSince = state.banExpiresAt
}
if state.banLevel > 0 && now.Sub(quietSince) > (2*l.cfg.baseBlockDuration) {
state.banLevel = 0
}
@@ -124,10 +134,17 @@ func (l *loginFilter) addLogin(wgPubKey string, metaHash uint64) {
return
}
if now.Sub(state.sessionStart) >= l.cfg.reconnThreshold {
state.sessionStart = now
state.sessionCounter = 0
}
state.sessionCounter++
if state.sessionCounter > l.cfg.reconnLimitForBan && now.Sub(state.sessionStart) < l.cfg.reconnThreshold {
if state.sessionCounter > l.cfg.reconnLimitForBan {
state.isBanned = true
state.banLevel++
if state.banLevel < l.cfg.maxBanLevel {
state.banLevel++
}
backoffFactor := math.Pow(2, float64(state.banLevel-1))
duration := time.Duration(float64(l.cfg.baseBlockDuration) * backoffFactor)
@@ -20,6 +20,7 @@ func testAdvancedCfg() *lfConfig {
baseBlockDuration: 100 * time.Millisecond,
reconnLimitForBan: 3,
metaChangeLimit: 2,
maxBanLevel: 3,
}
}
@@ -157,6 +158,187 @@ func (s *LoginFilterTestSuite) TestMetaChangeIsAllowedAfterWindowResets() {
s.Equal(1, s.filter.logged[pubKey].metaChangeCounter, "meta change counter should reset")
}
func (s *LoginFilterTestSuite) TestReconnectStormAfterQuietPeriodTriggersBan() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
s.filter.addLogin(pubKey, meta)
s.Require().Contains(s.filter.logged, pubKey)
s.filter.logged[pubKey].sessionStart = time.Now().Add(-(s.filter.cfg.reconnThreshold + time.Second))
s.filter.addLogin(pubKey, meta)
s.Equal(1, s.filter.logged[pubKey].sessionCounter, "expired window should restart the count")
for i := 1; i < limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.allowLogin(pubKey, meta))
s.False(s.filter.logged[pubKey].isBanned)
s.filter.addLogin(pubKey, meta)
s.False(s.filter.allowLogin(pubKey, meta))
s.True(s.filter.logged[pubKey].isBanned)
}
func (s *LoginFilterTestSuite) TestReconnectStormAfterBanExpiresTriggersBanAgain() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.Require().Contains(s.filter.logged, pubKey)
s.Require().True(s.filter.logged[pubKey].isBanned)
expired := time.Now().Add(-(s.filter.cfg.baseBlockDuration + time.Second))
s.filter.logged[pubKey].banExpiresAt = expired
s.filter.logged[pubKey].sessionStart = expired
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.logged[pubKey].isBanned)
s.Equal(2, s.filter.logged[pubKey].banLevel)
}
func (s *LoginFilterTestSuite) TestSlowReconnectsAcrossWindowsDoNotBan() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
for i := 0; i < limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.Require().Contains(s.filter.logged, pubKey)
s.filter.logged[pubKey].sessionStart = time.Now().Add(-(s.filter.cfg.reconnThreshold + time.Second))
for i := 0; i < limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.allowLogin(pubKey, meta))
s.False(s.filter.logged[pubKey].isBanned)
}
func (s *LoginFilterTestSuite) TestBanLevelEscalatesWhenStormResumesRightAfterBan() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
banTime := time.Now().Add(-3 * s.filter.cfg.baseBlockDuration)
s.filter.logged[pubKey] = &peerState{
currentHash: meta,
isBanned: true,
banLevel: 1,
banExpiresAt: time.Now().Add(-time.Millisecond),
sessionStart: banTime,
lastSeen: banTime,
}
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.logged[pubKey].isBanned)
s.Equal(2, s.filter.logged[pubKey].banLevel)
}
func (s *LoginFilterTestSuite) TestBanLevelResetsAfterQuietPeriodFollowingBan() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
quiet := 2*s.filter.cfg.baseBlockDuration + time.Second
s.filter.logged[pubKey] = &peerState{
currentHash: meta,
banLevel: 2,
banExpiresAt: time.Now().Add(-s.filter.cfg.baseBlockDuration),
lastSeen: time.Now().Add(-2 * quiet),
}
s.filter.addLogin(pubKey, meta)
s.Equal(2, s.filter.logged[pubKey].banLevel, "ban ended more recently than the quiet period")
s.filter.logged[pubKey].banExpiresAt = time.Now().Add(-quiet)
s.filter.logged[pubKey].lastSeen = time.Now().Add(-2 * quiet)
s.filter.addLogin(pubKey, meta)
s.Equal(0, s.filter.logged[pubKey].banLevel)
}
func (s *LoginFilterTestSuite) TestBanDurationIsCappedAtMaxLevel() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
maxLevel := s.filter.cfg.maxBanLevel
s.filter.logged[pubKey] = &peerState{
currentHash: meta,
banLevel: maxLevel,
sessionStart: time.Now(),
lastSeen: time.Now(),
}
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.logged[pubKey].isBanned)
s.Equal(maxLevel, s.filter.logged[pubKey].banLevel)
expected := s.filter.cfg.baseBlockDuration << (maxLevel - 1)
s.InDelta(expected, s.filter.logged[pubKey].banExpiresAt.Sub(s.filter.logged[pubKey].lastSeen), float64(time.Millisecond))
}
func (s *LoginFilterTestSuite) TestEstablishedPeerReconnectingOnceIsAllowed() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
longAgo := time.Now().Add(-time.Hour)
s.filter.logged[pubKey] = &peerState{
currentHash: meta,
sessionCounter: 1,
sessionStart: longAgo,
lastSeen: longAgo,
metaChangeWindowStart: longAgo,
metaChangeCounter: 1,
}
s.True(s.filter.allowLogin(pubKey, meta))
s.filter.addLogin(pubKey, meta)
s.True(s.filter.allowLogin(pubKey, meta))
s.False(s.filter.logged[pubKey].isBanned)
s.Equal(1, s.filter.logged[pubKey].sessionCounter)
}
func (s *LoginFilterTestSuite) TestLoginsDuringActiveBanDoNotExtendIt() {
pubKey := "PUB_KEY_A"
meta := uint64(1)
limit := s.filter.cfg.reconnLimitForBan
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.Require().Contains(s.filter.logged, pubKey)
s.Require().True(s.filter.logged[pubKey].isBanned)
expiresAt := time.Now().Add(time.Hour)
s.filter.logged[pubKey].banExpiresAt = expiresAt
lastSeen := s.filter.logged[pubKey].lastSeen
for i := 0; i <= limit; i++ {
s.filter.addLogin(pubKey, meta)
}
s.True(s.filter.logged[pubKey].isBanned)
s.Equal(1, s.filter.logged[pubKey].banLevel)
s.Equal(expiresAt, s.filter.logged[pubKey].banExpiresAt)
s.Equal(lastSeen, s.filter.logged[pubKey].lastSeen)
s.Equal(0, s.filter.logged[pubKey].sessionCounter)
}
func BenchmarkHashingMethods(b *testing.B) {
meta := nbpeer.PeerSystemMeta{
WtVersion: "1.25.1",
@@ -0,0 +1,135 @@
package grpc
import (
"context"
"time"
"github.com/netbirdio/netbird/encryption"
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
"github.com/netbirdio/netbird/management/server/telemetry"
"github.com/netbirdio/netbird/shared/management/proto"
log "github.com/sirupsen/logrus"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
)
func PeerUpdateHandlerFactory(
peerKey wgtypes.Key,
updates chan *network_map.UpdateMessage,
secretsManager SecretsManager,
srv proto.ManagementService_SyncServer,
cleanupfunc func()) *PeerUpdateHandler {
return &PeerUpdateHandler{
peerKey: peerKey,
updates: updates,
secretsManager: secretsManager,
srv: srv,
encrypter: encryption.DefaultEncrypter{},
debouncer: NewUpdateDebouncer(1000 * time.Millisecond),
cleanupFunc: cleanupfunc,
}
}
// PeerUpdateHandler sends updates to the connected peer until the updates channel is closed.
// It implements a backpressure mechanism that sends the first update immediately,
// then debounces subsequent rapid updates, ensuring only the latest update is sent
// after a quiet period.
type PeerUpdateHandler struct {
peerKey wgtypes.Key
updates chan *network_map.UpdateMessage
appMetrics telemetry.AppMetrics
secretsManager SecretsManager
srv syncSender
encrypter encryption.Encrypter
debouncer Debouncer
cleanupFunc func()
}
func (pu *PeerUpdateHandler) WithMetrics(appMetrics telemetry.AppMetrics) *PeerUpdateHandler {
pu.appMetrics = appMetrics
return pu
}
//go:generate go tool mockgen -source=./peer_update_handler.go -destination=./sync_sender_mock.go -package=grpc
type syncSender interface {
Send(*proto.EncryptedMessage) error
Context() context.Context
}
func (pu *PeerUpdateHandler) HandleUpdates(ctx context.Context) error {
log.WithContext(ctx).Tracef("starting to handle updates for peer %s", pu.peerKey.String())
defer pu.debouncer.Stop()
for {
select {
// condition when there are some updates
// todo set the updates channel size to 1
case update, open := <-pu.updates:
if pu.appMetrics != nil {
pu.appMetrics.GRPCMetrics().UpdateChannelQueueLength(len(pu.updates) + 1)
}
if !open {
log.WithContext(ctx).Debugf("updates channel for peer %s was closed", pu.peerKey.String())
pu.cleanupFunc()
return nil
}
log.WithContext(ctx).Tracef("received an update for peer %s", pu.peerKey.String())
if pu.debouncer.ProcessUpdate(update) {
// Send immediately (first update or after quiet period)
if err := pu.SendUpdate(ctx, update); err != nil {
log.WithContext(ctx).Debugf("error while sending an update to peer %s: %v", pu.peerKey.String(), err)
return err
}
}
// Timer expired - quiet period reached, send pending updates if any
case <-pu.debouncer.TimerChannel():
pendingUpdates := pu.debouncer.GetPendingUpdates()
if len(pendingUpdates) == 0 {
continue
}
log.WithContext(ctx).Debugf("sending %d debounced update(s) for peer %s", len(pendingUpdates), pu.peerKey.String())
for _, pendingUpdate := range pendingUpdates {
if err := pu.SendUpdate(ctx, pendingUpdate); err != nil {
log.WithContext(ctx).Debugf("error while sending an update to peer %s: %v", pu.peerKey.String(), err)
return err
}
}
// condition when client <-> server connection has been terminated
case <-pu.srv.Context().Done():
// happens when connection drops, e.g. client disconnects
log.WithContext(ctx).Debugf("stream of peer %s has been closed", pu.peerKey.String())
pu.cleanupFunc()
return pu.srv.Context().Err()
}
}
}
func (pu *PeerUpdateHandler) SendUpdate(ctx context.Context, update *network_map.UpdateMessage) error {
key, err := pu.secretsManager.GetWGKey()
if err != nil {
pu.cleanupFunc()
return status.Errorf(codes.Internal, "failed processing update message")
}
encryptedResp, err := pu.encrypter.EncryptMessage(pu.peerKey, key, update.Update)
if err != nil {
pu.cleanupFunc()
return status.Errorf(codes.Internal, "failed processing update message")
}
err = pu.srv.Send(&proto.EncryptedMessage{
WgPubKey: key.PublicKey().String(),
Body: encryptedResp,
})
if err != nil {
pu.cleanupFunc()
return status.Errorf(codes.Internal, "failed sending update message")
}
log.WithContext(ctx).Tracef("sent an update to peer %s", pu.peerKey.String())
return nil
}
@@ -0,0 +1,155 @@
package grpc
import (
"context"
"fmt"
"sync"
"testing"
"time"
pb "github.com/golang/protobuf/proto" //nolint
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
"github.com/netbirdio/netbird/shared/management/proto"
"github.com/stretchr/testify/assert"
"go.uber.org/mock/gomock"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
)
func TestSendPeerUpdates_FirstUpdate(t *testing.T) {
ctrl := gomock.NewController(t)
secretsManager := NewMockSecretsManager(ctrl)
updateDebouncer := NewMockDebouncer(ctrl)
syncSender := NewMocksyncSender(ctrl)
pu := PeerUpdateHandler{
peerKey: mustGenerateKey(t),
updates: make(chan *network_map.UpdateMessage),
secretsManager: secretsManager,
encrypter: testEncrypter{},
debouncer: updateDebouncer,
srv: syncSender,
cleanupFunc: func() {},
}
msg := network_map.UpdateMessage{
Update: &proto.SyncResponse{Version: 1},
}
timeCh := make(chan time.Time)
srvCtx := context.TODO()
srvKey := mustGenerateKey(t)
// mock a first update, should send it right away
updateDebouncer.EXPECT().ProcessUpdate(gomock.Eq(&msg)).Return(true)
updateDebouncer.EXPECT().TimerChannel().AnyTimes().Return(timeCh)
syncSender.EXPECT().Context().AnyTimes().Return(srvCtx)
secretsManager.EXPECT().GetWGKey().Return(srvKey, nil)
syncSender.EXPECT().Send(pbMatcher{x: &proto.EncryptedMessage{WgPubKey: srvKey.PublicKey().String(), Body: mustMarshal(t, &msg)}})
updateDebouncer.EXPECT().Stop()
var wg sync.WaitGroup
wg.Go(func() { pu.HandleUpdates(context.TODO()) }) //nolint:errcheck
pu.updates <- &msg
close(pu.updates)
wg.Wait()
}
func TestSendPeerUpdates_TimerUpdate(t *testing.T) {
ctrl := gomock.NewController(t)
secretsManager := NewMockSecretsManager(ctrl)
updateDebouncer := NewMockDebouncer(ctrl)
syncSender := NewMocksyncSender(ctrl)
pu := PeerUpdateHandler{
peerKey: mustGenerateKey(t),
updates: make(chan *network_map.UpdateMessage),
secretsManager: secretsManager,
encrypter: testEncrypter{},
debouncer: updateDebouncer,
srv: syncSender,
cleanupFunc: func() {},
}
msg := network_map.UpdateMessage{
Update: &proto.SyncResponse{Version: 1},
}
timeCh := make(chan time.Time)
srvCtx := context.TODO()
srvKey := mustGenerateKey(t)
updateDebouncer.EXPECT().GetPendingUpdates().Return([]*network_map.UpdateMessage{&msg})
updateDebouncer.EXPECT().TimerChannel().AnyTimes().Return(timeCh)
syncSender.EXPECT().Context().AnyTimes().Return(srvCtx)
secretsManager.EXPECT().GetWGKey().Return(srvKey, nil)
syncSender.EXPECT().Send(pbMatcher{x: &proto.EncryptedMessage{WgPubKey: srvKey.PublicKey().String(), Body: mustMarshal(t, &msg)}})
updateDebouncer.EXPECT().Stop()
var wg sync.WaitGroup
wg.Go(func() { pu.HandleUpdates(context.TODO()) }) //nolint:errcheck
timeCh <- time.Now()
close(pu.updates)
wg.Wait()
}
func TestSendPeerUpdates_ServerContextDone(t *testing.T) {
ctrl := gomock.NewController(t)
secretsManager := NewMockSecretsManager(ctrl)
updateDebouncer := NewMockDebouncer(ctrl)
syncSender := NewMocksyncSender(ctrl)
pu := PeerUpdateHandler{
peerKey: mustGenerateKey(t),
updates: make(chan *network_map.UpdateMessage),
secretsManager: secretsManager,
encrypter: testEncrypter{},
debouncer: updateDebouncer,
srv: syncSender,
cleanupFunc: func() {},
}
timeCh := make(chan time.Time)
srvCtx, cancel := context.WithCancel(context.TODO())
updateDebouncer.EXPECT().TimerChannel().AnyTimes().Return(timeCh)
syncSender.EXPECT().Context().AnyTimes().Return(srvCtx)
updateDebouncer.EXPECT().Stop()
var wg sync.WaitGroup
wg.Go(func() { pu.HandleUpdates(context.TODO()) }) //nolint:errcheck
cancel()
wg.Wait()
}
func mustGenerateKey(t *testing.T) wgtypes.Key {
t.Helper()
k, err := wgtypes.GenerateKey()
assert.NoError(t, err)
return k
}
func mustMarshal(t *testing.T, msg *network_map.UpdateMessage) []byte {
t.Helper()
r, err := pb.Marshal(msg.Update)
assert.NoError(t, err)
return r
}
type testEncrypter struct{}
func (testEncrypter) EncryptMessage(remotePubKey wgtypes.Key, ourPrivateKey wgtypes.Key, message pb.Message) ([]byte, error) {
return pb.Marshal(message)
}
type pbMatcher struct {
x pb.Message
}
func (pbm pbMatcher) Matches(x any) bool {
msg, ok := x.(pb.Message)
if !ok {
return false
}
return pb.Equal(pbm.x, msg)
}
func (pbm pbMatcher) String() string {
return fmt.Sprintf("is equal to %s (%T)", pbm.x, pbm.x)
}
@@ -1,54 +0,0 @@
package grpc
import (
"context"
"fmt"
"time"
"github.com/eko/gocache/lib/v4/cache"
"github.com/eko/gocache/lib/v4/store"
log "github.com/sirupsen/logrus"
)
// 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]
ctx context.Context
}
// NewPKCEVerifierStore creates a PKCE verifier store using the provided shared cache store.
func NewPKCEVerifierStore(ctx context.Context, cacheStore store.StoreInterface) *PKCEVerifierStore {
return &PKCEVerifierStore{
cache: cache.New[string](cacheStore),
ctx: ctx,
}
}
// Store saves a PKCE verifier associated with an OAuth state parameter.
// The verifier is stored with the specified TTL and will be automatically deleted after expiration.
func (s *PKCEVerifierStore) Store(state, verifier string, ttl time.Duration) error {
if err := s.cache.Set(s.ctx, state, verifier, store.WithExpiration(ttl)); err != nil {
return fmt.Errorf("failed to store PKCE verifier: %w", err)
}
log.Debugf("Stored PKCE verifier for state (expires in %s)", ttl)
return nil
}
// LoadAndDelete retrieves and removes a PKCE verifier for the given state.
// 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)
if err != nil {
log.Debugf("PKCE verifier not found for state")
return "", false
}
if err := s.cache.Delete(s.ctx, state); err != nil {
log.Warnf("Failed to delete PKCE verifier for state: %v", err)
}
return verifier, true
}
+90 -29
View File
@@ -27,8 +27,6 @@ import (
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
"github.com/netbirdio/netbird/shared/management/domain"
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
"github.com/netbirdio/netbird/management/internals/modules/peers"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
@@ -42,6 +40,7 @@ import (
"github.com/netbirdio/netbird/management/server/users"
proxyauth "github.com/netbirdio/netbird/proxy/auth"
"github.com/netbirdio/netbird/shared/hash/argon2id"
"github.com/netbirdio/netbird/shared/management/domain"
"github.com/netbirdio/netbird/shared/management/proto"
nbstatus "github.com/netbirdio/netbird/shared/management/status"
)
@@ -102,7 +101,8 @@ type ProxyServiceServer struct {
mu sync.RWMutex
// Manager for reverse proxy operations
serviceManager rpservice.Manager
serviceManager rpservice.Manager
credentialLimits credentialVerificationLimiter
// agentNetworkSynth produces synthesised reverse-proxy services from
// Agent Network state. Optional — when nil the snapshot path only ships
// persisted services.
@@ -141,8 +141,8 @@ type ProxyServiceServer struct {
// OIDC configuration for proxy authentication
oidcConfig ProxyOIDCConfig
// Store for PKCE verifiers
pkceVerifierStore *PKCEVerifierStore
// singleUseStore backs both PKCE verifiers and OIDC session exchange codes.
singleUseStore *SingleUseStore
// tokenTTL is the lifetime of one-time tokens generated for proxy
// authentication. Defaults to defaultProxyTokenTTL when zero.
@@ -157,6 +157,13 @@ type ProxyServiceServer struct {
const pkceVerifierTTL = 10 * time.Minute
const sessionCodeTTL = 60 * time.Second
const sessionCodeCacheNamespace = "proxy:session"
// The signed nonce binds the handoff mode without changing the state format.
const sessionCodeNoncePrefix = "code."
const defaultProxyTokenTTL = 5 * time.Minute
const defaultSnapshotBatchSize = 500
@@ -207,13 +214,13 @@ func enforceAccountScope(ctx context.Context, requestAccountID string) error {
}
// NewProxyServiceServer creates a new proxy service server.
func NewProxyServiceServer(accessLogMgr accesslogs.Manager, tokenStore *OneTimeTokenStore, pkceStore *PKCEVerifierStore, oidcConfig ProxyOIDCConfig, peersManager peers.Manager, usersManager users.Manager, idpManager idp.Manager, proxyMgr proxy.Manager, tokenChecker ProxyTokenChecker) *ProxyServiceServer {
func NewProxyServiceServer(accessLogMgr accesslogs.Manager, tokenStore *OneTimeTokenStore, singleUseStore *SingleUseStore, oidcConfig ProxyOIDCConfig, peersManager peers.Manager, usersManager users.Manager, idpManager idp.Manager, proxyMgr proxy.Manager, tokenChecker ProxyTokenChecker) *ProxyServiceServer {
ctx, cancel := context.WithCancel(context.Background())
s := &ProxyServiceServer{
accessLogManager: accessLogMgr,
oidcConfig: oidcConfig,
tokenStore: tokenStore,
pkceVerifierStore: pkceStore,
singleUseStore: singleUseStore,
peersManager: peersManager,
usersManager: usersManager,
idpManager: idpManager,
@@ -242,9 +249,10 @@ func (s *ProxyServiceServer) cleanupStaleProxies(ctx context.Context) {
}
}
// Close stops background goroutines.
// Close stops background goroutines and releases credential verification state.
func (s *ProxyServiceServer) Close() {
s.cancel()
s.credentialLimits.close()
}
// SetServiceManager sets the service manager. Must be called before serving.
@@ -304,6 +312,16 @@ func (s *ProxyServiceServer) proxyConnectAuthorizer() ProxyConnectAuthorizer {
return s.connectAuthorizer
}
// GenerateSessionCode creates a single-use code for the given session token.
func (s *ProxyServiceServer) GenerateSessionCode(sessionToken string) (code string, ok bool) {
code, err := s.singleUseStore.Generate(sessionCodeCacheNamespace, sessionToken, sessionCodeTTL)
if err != nil {
log.WithError(err).Error("failed to generate proxy session code")
return "", false
}
return code, true
}
// CheckLLMPolicyLimits is the pre-flight policy gate the proxy calls before
// forwarding an LLM request upstream. Delegates to the agent-network selector,
// which scores applicable policies by remaining headroom and returns the
@@ -412,6 +430,7 @@ func (s *ProxyServiceServer) SetProxyController(proxyController proxy.Controller
type proxyConnectParams struct {
proxyID string
address string
version string
capabilities *proto.ProxyCapabilities
}
@@ -422,6 +441,7 @@ func (s *ProxyServiceServer) GetMappingUpdate(req *proto.GetMappingUpdateRequest
return err
}
params.capabilities = req.GetCapabilities()
params.version = req.GetVersion()
conn, proxyRecord, err := s.registerProxyConnection(stream.Context(), params, &proxyConnection{
stream: stream,
@@ -455,6 +475,7 @@ func (s *ProxyServiceServer) SyncMappings(stream proto.ProxyService_SyncMappings
return err
}
params.capabilities = init.GetCapabilities()
params.version = init.GetVersion()
conn, proxyRecord, err := s.registerProxyConnection(stream.Context(), params, &proxyConnection{
syncStream: stream,
@@ -566,7 +587,7 @@ func (s *ProxyServiceServer) registerProxyConnection(ctx context.Context, params
}
}
proxyRecord, err := s.proxyManager.Connect(ctx, params.proxyID, sessionID, params.address, peerInfo, accountID, caps)
proxyRecord, err := s.proxyManager.Connect(ctx, params.proxyID, sessionID, params.address, peerInfo, params.version, accountID, caps)
if err != nil {
cancel()
if accountID != nil {
@@ -1223,6 +1244,7 @@ func shallowCloneMapping(m *proto.ProxyMapping) *proto.ProxyMapping {
}
}
// Authenticate verifies service credentials and issues a session token.
func (s *ProxyServiceServer) Authenticate(ctx context.Context, req *proto.AuthenticateRequest) (*proto.AuthenticateResponse, error) {
if err := enforceAccountScope(ctx, req.GetAccountId()); err != nil {
return nil, err
@@ -1234,6 +1256,14 @@ func (s *ProxyServiceServer) Authenticate(ctx context.Context, req *proto.Authen
return nil, status.Errorf(codes.FailedPrecondition, "get service from store: %v", err)
}
switch req.GetRequest().(type) {
case *proto.AuthenticateRequest_Pin, *proto.AuthenticateRequest_Password:
key := credentialVerificationKey{accountID: credentialAccountID(service.AccountID), serviceID: credentialServiceID(service.ID)}
if err := s.credentialLimits.allow(key); err != nil {
return nil, err
}
}
authenticated, userId, method := s.authenticateRequest(ctx, req, service)
// Non-OIDC schemes (PIN/Password/Header) authenticate against per-service
@@ -1522,18 +1552,20 @@ func (s *ProxyServiceServer) GetOIDCURL(ctx context.Context, req *proto.GetOIDCU
log.WithContext(ctx).Errorf("failed to get account services: %v", err)
return nil, status.Errorf(codes.FailedPrecondition, "get account services: %v", err)
}
var found bool
var matchedService *rpservice.Service
for _, service := range services {
if service.Domain == redirectURL.Hostname() {
found = true
matchedService = service
break
}
}
if !found {
if matchedService == nil {
log.WithContext(ctx).Debugf("OIDC redirect URL %q does not match any service domain", redirectURL.Hostname())
return nil, status.Errorf(codes.FailedPrecondition, "service not found in store")
}
useSessionCode := s.proxyManager.ClusterSupportsSessionCode(ctx, matchedService.ProxyCluster)
provider, err := oidc.NewProvider(ctx, s.oidcConfig.Issuer)
if err != nil {
log.WithContext(ctx).Errorf("failed to create OIDC provider: %v", err)
@@ -1553,15 +1585,18 @@ func (s *ProxyServiceServer) GetOIDCURL(ctx context.Context, req *proto.GetOIDCU
return nil, status.Errorf(codes.Internal, "generate nonce: %v", err)
}
nonceB64 := base64.URLEncoding.EncodeToString(nonce)
if useSessionCode {
nonceB64 = sessionCodeNoncePrefix + nonceB64
}
// Using an HMAC here to avoid redirection state being modified.
// State format: base64(redirectURL)|nonce|hmac(redirectURL|nonce)
// State format: base64(redirectURL)|[code.]nonce|hmac(redirectURL|nonce)
payload := redirectURL.String() + "|" + nonceB64
hmacSum := s.generateHMAC(payload)
state := fmt.Sprintf("%s|%s|%s", base64.URLEncoding.EncodeToString([]byte(redirectURL.String())), nonceB64, hmacSum)
codeVerifier := oauth2.GenerateVerifier()
if err := s.pkceVerifierStore.Store(state, codeVerifier, pkceVerifierTTL); err != nil {
if err := s.singleUseStore.Store(state, codeVerifier, pkceVerifierTTL); err != nil {
log.WithContext(ctx).Errorf("failed to store PKCE verifier: %v", err)
return nil, status.Errorf(codes.Internal, "store PKCE verifier: %v", err)
}
@@ -1598,15 +1633,12 @@ func (s *ProxyServiceServer) generateHMAC(input string) string {
return hex.EncodeToString(mac.Sum(nil))
}
// ValidateState validates the state parameter from an OAuth callback.
// Returns the original redirect URL if valid, or an error if invalid.
// The HMAC is verified before consuming the PKCE verifier to prevent
// an attacker from invalidating a legitimate user's auth flow.
func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL string, err error) {
// State format: base64(redirectURL)|nonce|hmac(redirectURL|nonce)
// ValidateState validates and consumes an OIDC state.
func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL string, useSessionCode bool, err error) {
// State format: base64(redirectURL)|[code.]nonce|hmac(redirectURL|nonce)
parts := strings.Split(state, "|")
if len(parts) != 3 {
return "", "", errors.New("invalid state format")
return "", "", false, errors.New("invalid state format")
}
encodedURL := parts[0]
@@ -1615,7 +1647,7 @@ func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL
redirectURLBytes, err := base64.URLEncoding.DecodeString(encodedURL)
if err != nil {
return "", "", fmt.Errorf("invalid state encoding: %w", err)
return "", "", false, fmt.Errorf("invalid state encoding: %w", err)
}
redirectURL = string(redirectURLBytes)
@@ -1623,16 +1655,17 @@ func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL
expectedHMAC := s.generateHMAC(payload)
if !hmac.Equal([]byte(providedHMAC), []byte(expectedHMAC)) {
return "", "", errors.New("invalid state signature")
return "", "", false, errors.New("invalid state signature")
}
useSessionCode = strings.HasPrefix(nonce, sessionCodeNoncePrefix)
// Consume the PKCE verifier only after HMAC validation passes.
verifier, ok := s.pkceVerifierStore.LoadAndDelete(state)
verifier, ok := s.singleUseStore.LoadAndDelete(state)
if !ok {
return "", "", errors.New("no verifier for state")
return "", "", false, errors.New("no verifier for state")
}
return verifier, redirectURL, nil
return verifier, redirectURL, useSessionCode, nil
}
// Denied reasons reported to the proxy when access is refused because of the
@@ -1651,6 +1684,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 +1726,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 +1765,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(
@@ -1823,7 +1870,20 @@ func (s *ProxyServiceServer) getAccountServiceByDomain(ctx context.Context, acco
// ValidateSession validates a session token and checks if the user has access to the domain.
func (s *ProxyServiceServer) ValidateSession(ctx context.Context, req *proto.ValidateSessionRequest) (*proto.ValidateSessionResponse, error) {
domain := req.GetDomain()
sessionToken := req.GetSessionToken()
sessionToken := req.GetSessionToken() //nolint:staticcheck
// A one-time code from the OIDC callback is redeemed here for the durable
// token, so the token never travels in a redirect URL. The redeemed token
// is returned to the proxy (mintedToken) to install as the session cookie.
mintedToken := ""
if code := req.GetSessionCode(); code != "" {
redeemed, found := s.singleUseStore.LoadAndDelete(singleUseCacheKey(sessionCodeCacheNamespace, code))
if !found {
return deniedSessionResponse("invalid or expired session code"), nil
}
sessionToken = redeemed
mintedToken = redeemed
}
if domain == "" || sessionToken == "" {
return deniedSessionResponse("missing domain or session_token"), nil
@@ -1896,6 +1956,7 @@ func (s *ProxyServiceServer) ValidateSession(ctx context.Context, req *proto.Val
UserEmail: user.Email,
PeerGroupIds: groupIDs,
PeerGroupNames: groupNames,
SessionToken: mintedToken,
}, nil
}
@@ -0,0 +1,93 @@
package grpc
import (
"context"
"testing"
"github.com/stretchr/testify/require"
"go.uber.org/mock/gomock"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
"github.com/netbirdio/netbird/shared/management/proto"
)
const (
versionTestProxyID = "proxy-a"
versionTestCluster = "cluster.example.com"
versionTestVersion = "0.60.0"
)
// hangupStream cancels its context on the first Send, emulating a proxy that
// disconnects right after receiving the initial snapshot. The legacy stream
// carries no proxy-to-management messages, so this is the only way for
// GetMappingUpdate to return.
type hangupStream struct {
recordingStream
ctx context.Context
cancel context.CancelFunc
}
func (s *hangupStream) Send(m *proto.GetMappingUpdateResponse) error {
s.cancel()
return s.recordingStream.Send(m)
}
func (s *hangupStream) Context() context.Context { return s.ctx }
// newVersionTestServer wires a server whose proxy manager only accepts a
// Connect carrying versionTestVersion, so a dropped or mangled version fails
// the test as an unexpected call.
func newVersionTestServer(t *testing.T) *ProxyServiceServer {
t.Helper()
ctrl := gomock.NewController(t)
svcMgr := rpservice.NewMockManager(ctrl)
svcMgr.EXPECT().GetGlobalServices(gomock.Any()).Return(nil, nil)
proxyMgr := proxy.NewMockManager(ctrl)
proxyMgr.EXPECT().
Connect(gomock.Any(), versionTestProxyID, gomock.Any(), versionTestCluster, gomock.Any(), versionTestVersion, gomock.Any(), gomock.Any()).
Return(&proxy.Proxy{ID: versionTestProxyID, Version: versionTestVersion}, nil)
proxyMgr.EXPECT().Disconnect(gomock.Any(), versionTestProxyID, gomock.Any()).Return(nil)
s := newSnapshotTestServer(t, 10)
s.serviceManager = svcMgr
s.proxyManager = proxyMgr
return s
}
func TestSyncMappings_ForwardsProxyVersion(t *testing.T) {
s := newVersionTestServer(t)
// The init carries the version, the ack acknowledges the empty snapshot,
// and the exhausted fake stream then ends the RPC.
stream := &syncRecordingStream{
recvMsgs: []*proto.SyncMappingsRequest{
{Msg: &proto.SyncMappingsRequest_Init{Init: &proto.SyncMappingsInit{
ProxyId: versionTestProxyID,
Address: versionTestCluster,
Version: versionTestVersion,
}}},
ackMsg(),
},
}
err := s.SyncMappings(stream)
require.ErrorContains(t, err, "no more recv messages")
}
func TestGetMappingUpdate_ForwardsProxyVersion(t *testing.T) {
s := newVersionTestServer(t)
ctx, cancel := context.WithCancel(context.Background())
t.Cleanup(cancel)
stream := &hangupStream{ctx: ctx, cancel: cancel}
err := s.GetMappingUpdate(&proto.GetMappingUpdateRequest{
ProxyId: versionTestProxyID,
Address: versionTestCluster,
Version: versionTestVersion,
}, stream)
require.ErrorIs(t, err, context.Canceled)
}
@@ -0,0 +1,101 @@
package grpc
import (
"sync"
"time"
"golang.org/x/time/rate"
"google.golang.org/genproto/googleapis/rpc/errdetails"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
"google.golang.org/protobuf/types/known/durationpb"
)
const (
credentialVerificationInterval = 6 * time.Second
credentialVerificationBurst = 5
credentialVerificationMaxServices = 4096
credentialVerificationIdleTimeout = 15 * time.Minute
credentialVerificationCleanupInterval = time.Minute
)
type credentialAccountID string
type credentialServiceID string
type credentialVerificationKey struct {
accountID credentialAccountID
serviceID credentialServiceID
}
type credentialVerificationBudget struct {
limiter *rate.Limiter
lastUsed time.Time
}
// The zero value is ready to use. Budgets are local to this Management process;
// proxy replicas reaching this process share a service's verification budget.
type credentialVerificationLimiter struct {
mu sync.Mutex
now func() time.Time
services map[credentialVerificationKey]*credentialVerificationBudget
nextCleanup time.Time
closed bool
}
func (l *credentialVerificationLimiter) allow(key credentialVerificationKey) error {
l.mu.Lock()
defer l.mu.Unlock()
if l.closed {
return status.Error(codes.Unavailable, "credential verification is closed")
}
now := time.Now()
if l.now != nil {
now = l.now()
}
l.cleanup(now)
budget := l.services[key]
if budget == nil {
if len(l.services) >= credentialVerificationMaxServices {
return credentialVerificationThrottled(credentialVerificationCleanupInterval)
}
if l.services == nil {
l.services = make(map[credentialVerificationKey]*credentialVerificationBudget)
}
budget = &credentialVerificationBudget{limiter: rate.NewLimiter(rate.Every(credentialVerificationInterval), credentialVerificationBurst)}
l.services[key] = budget
}
budget.lastUsed = now
if budget.limiter.AllowN(now, 1) {
return nil
}
delay := max(time.Nanosecond, time.Duration((1-budget.limiter.TokensAt(now))*float64(credentialVerificationInterval)))
return credentialVerificationThrottled(delay)
}
func (l *credentialVerificationLimiter) cleanup(now time.Time) {
if now.Before(l.nextCleanup) {
return
}
l.nextCleanup = now.Add(credentialVerificationCleanupInterval)
for key, budget := range l.services {
if now.Sub(budget.lastUsed) >= credentialVerificationIdleTimeout {
delete(l.services, key)
}
}
}
func (l *credentialVerificationLimiter) close() {
l.mu.Lock()
defer l.mu.Unlock()
l.closed = true
l.services = nil
}
func credentialVerificationThrottled(delay time.Duration) error {
s := status.New(codes.ResourceExhausted, "too many credential verification attempts")
withRetry, err := s.WithDetails(&errdetails.RetryInfo{RetryDelay: durationpb.New(delay)})
if err != nil {
return s.Err()
}
return withRetry.Err()
}
@@ -0,0 +1,79 @@
package grpc
import (
"strconv"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/genproto/googleapis/rpc/errdetails"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/status"
)
func TestCredentialVerificationRefillAndIsolation(t *testing.T) {
now := time.Now()
l := credentialVerificationLimiter{now: func() time.Time { return now }}
key := credentialVerificationKey{accountID: "account", serviceID: "service"}
for range credentialVerificationBurst {
require.NoError(t, l.allow(key))
}
err := l.allow(key)
require.Equal(t, codes.ResourceExhausted, status.Code(err), "the burst must be bounded")
now = now.Add(3 * time.Second)
err = l.allow(key)
require.Equal(t, codes.ResourceExhausted, status.Code(err), "a partially refilled token must not permit a check")
details := status.Convert(err).Details()
require.Len(t, details, 1, "throttling must provide RetryInfo")
retry, ok := details[0].(*errdetails.RetryInfo)
require.True(t, ok, "retry details must use the standard message")
assert.Equal(t, 3*time.Second, retry.RetryDelay.AsDuration(), "retry hint must reflect time until the next check")
now = now.Add(3 * time.Second)
require.NoError(t, l.allow(key))
assert.Equal(t, codes.ResourceExhausted, status.Code(l.allow(key)), "only one check must refill every six seconds")
require.NoError(t, l.allow(credentialVerificationKey{accountID: "other-account", serviceID: key.serviceID}))
require.NoError(t, l.allow(credentialVerificationKey{accountID: key.accountID, serviceID: "other-service"}))
}
func TestCredentialVerificationCapacityAndExpiry(t *testing.T) {
now := time.Now()
l := credentialVerificationLimiter{now: func() time.Time { return now }}
for i := range credentialVerificationMaxServices {
require.NoError(t, l.allow(credentialVerificationKey{accountID: "account", serviceID: credentialServiceID(strconv.Itoa(i))}))
}
key := credentialVerificationKey{accountID: "account", serviceID: "new-service"}
assert.Equal(t, codes.ResourceExhausted, status.Code(l.allow(key)), "capacity exhaustion must deny new checks")
now = now.Add(credentialVerificationIdleTimeout)
for range credentialVerificationBurst {
require.NoError(t, l.allow(key))
}
assert.Equal(t, codes.ResourceExhausted, status.Code(l.allow(key)), "expiry must retain the normal burst bound")
}
func TestCredentialVerificationConcurrentChecksAndClose(t *testing.T) {
var l credentialVerificationLimiter
key := credentialVerificationKey{accountID: "account", serviceID: "service"}
var admitted atomic.Int32
var wg sync.WaitGroup
for range 100 {
wg.Go(func() {
if err := l.allow(key); err == nil {
admitted.Add(1)
} else {
assert.Equal(t, codes.ResourceExhausted, status.Code(err), "excess checks must be throttled")
}
})
}
wg.Wait()
assert.EqualValues(t, credentialVerificationBurst, admitted.Load(), "concurrent checks must share the burst")
for range 10 {
wg.Go(l.close)
wg.Go(func() { assert.Error(t, l.allow(key)) })
}
wg.Wait()
assert.Empty(t, l.services, "closing must release retained budgets")
assert.Equal(t, codes.Unavailable, status.Code(l.allow(key)), "checks after close must fail closed")
}
@@ -0,0 +1,18 @@
# Reverse proxy credential verification
The `ProxyService.Authenticate` RPC limits PIN and password checks before
verifying their Argon2 hashes. Both methods share one budget per account and
service: a burst of five checks, replenishing one check every six seconds
(ten per minute). Successful and failed checks consume the budget. Account
scope and service lookup run before the limiter.
Excess checks receive gRPC `ResourceExhausted` with a standard `RetryInfo` delay.
Updated proxies translate it to HTTP 429 and `Retry-After`. Older proxies show
an authentication-service error but cannot bypass the Management limit.
Budgets are held in memory per Management process and reset on restart. Proxy
replicas reaching the same Management process share its budgets. Multiple
Management processes have independent budgets; this is not a cluster-wide
limit. At most 4,096 service budgets are retained, with idle entries expiring
after fifteen minutes. Capacity exhaustion denies new checks until entries
expire. Closing the server releases the retained state.
@@ -0,0 +1,131 @@
package grpc_test
import (
"context"
"net"
"net/netip"
"sync"
"sync/atomic"
"testing"
"time"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"google.golang.org/genproto/googleapis/rpc/errdetails"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/metadata"
"google.golang.org/grpc/peer"
"google.golang.org/grpc/status"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
servicemanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service/manager"
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/sessionkey"
nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/management/proto"
)
func credentialServer(t *testing.T) (*nbgrpc.ProxyServiceServer, context.Context, grpc.UnaryServerInterceptor) {
t.Helper()
ctx := context.Background()
s, err := store.NewStore(ctx, types.SqliteStoreEngine, t.TempDir(), nil, false)
require.NoError(t, err)
t.Cleanup(func() { assert.NoError(t, s.Close(ctx)) })
require.NoError(t, s.SaveAccount(ctx, &types.Account{Id: "account"}))
keys, err := sessionkey.GenerateKeyPair()
require.NoError(t, err)
for _, id := range []string{"service", "other-service"} {
svc := &service.Service{
ID: id, AccountID: "account", Name: id, Domain: id + ".example.com",
Enabled: true, SessionPrivateKey: keys.PrivateKey, SessionPublicKey: keys.PublicKey,
Auth: service.AuthConfig{
PinAuth: &service.PINAuthConfig{Enabled: true, Pin: "842716"},
PasswordAuth: &service.PasswordAuthConfig{Enabled: true, Password: "test-password"},
},
}
require.NoError(t, svc.Auth.HashSecrets())
require.NoError(t, s.CreateService(ctx, svc))
}
account := "account"
token, err := types.CreateNewProxyAccessToken("test proxy", time.Hour, &account, "admin")
require.NoError(t, err)
require.NoError(t, s.SaveProxyAccessToken(ctx, &token.ProxyAccessToken))
ctx = metadata.NewIncomingContext(ctx, metadata.Pairs("authorization", "Bearer "+string(token.PlainToken)))
ctx = peer.NewContext(ctx, &peer.Peer{Addr: net.TCPAddrFromAddrPort(netip.MustParseAddrPort("192.0.2.1:443"))})
server := nbgrpc.NewProxyServiceServer(nil, nil, nil, nbgrpc.ProxyOIDCConfig{}, nil, nil, nil, nil, nil)
t.Cleanup(server.Close)
server.SetServiceManager(servicemanager.NewManager(s, nil, nil, nil, nil, nil))
interceptor, _, closeInterceptor := nbgrpc.NewProxyAuthInterceptors(s)
t.Cleanup(closeInterceptor)
return server, ctx, interceptor
}
func TestAuthenticateCredentialRateLimit(t *testing.T) {
server, ctx, interceptor := credentialServer(t)
authenticate := func(req *proto.AuthenticateRequest) (*proto.AuthenticateResponse, error) {
response, err := interceptor(ctx, req, &grpc.UnaryServerInfo{FullMethod: "/management.ProxyService/Authenticate"}, func(ctx context.Context, req any) (any, error) {
return server.Authenticate(ctx, req.(*proto.AuthenticateRequest))
})
if err != nil {
return nil, err
}
return response.(*proto.AuthenticateResponse), nil
}
for i := range 5 {
req := &proto.AuthenticateRequest{AccountId: "account", Id: "service"}
if i%2 == 0 {
req.Request = &proto.AuthenticateRequest_Pin{Pin: &proto.PinRequest{Pin: "000000"}}
} else {
req.Request = &proto.AuthenticateRequest_Password{Password: &proto.PasswordRequest{Password: "wrong-password"}}
}
resp, err := authenticate(req)
require.NoError(t, err)
assert.False(t, resp.GetSuccess(), "incorrect PINs and passwords must be denied")
assert.Empty(t, resp.GetSessionToken(), "incorrect credentials must not issue a token")
}
req := &proto.AuthenticateRequest{AccountId: "account", Id: "service", Request: &proto.AuthenticateRequest_Pin{Pin: &proto.PinRequest{Pin: "842716"}}}
resp, err := authenticate(req)
assert.Nil(t, resp, "a throttled verification must not return a session")
require.Equal(t, codes.ResourceExhausted, status.Code(err), "PIN and password checks must share a service budget even with a valid proxy token")
details := status.Convert(err).Details()
require.Len(t, details, 1, "throttled responses must include a retry hint")
retry, ok := details[0].(*errdetails.RetryInfo)
require.True(t, ok, "the hint must use the standard RetryInfo message")
assert.Positive(t, retry.RetryDelay.AsDuration(), "the retry delay must be positive")
assert.LessOrEqual(t, retry.RetryDelay.AsDuration(), 6*time.Second, "the service must replenish one verification every six seconds")
req.AccountId = "another-account"
_, err = authenticate(req)
assert.Equal(t, codes.PermissionDenied, status.Code(err), "account scope must still be enforced before throttling")
req.AccountId = "account"
req.Id = "other-service"
resp, err = authenticate(req)
require.NoError(t, err)
assert.True(t, resp.GetSuccess(), "one service's throttle must not block another service")
assert.NotEmpty(t, resp.GetSessionToken(), "valid credentials on another service must issue a session")
}
func TestAuthenticateCredentialConcurrentLimit(t *testing.T) {
server, _, _ := credentialServer(t)
req := &proto.AuthenticateRequest{AccountId: "account", Id: "service", Request: &proto.AuthenticateRequest_Pin{Pin: &proto.PinRequest{Pin: "000000"}}}
var checked, throttled atomic.Int32
var wg sync.WaitGroup
for range 20 {
wg.Go(func() {
resp, err := server.Authenticate(context.Background(), req)
switch status.Code(err) {
case codes.OK:
checked.Add(1)
assert.False(t, resp.GetSuccess(), "incorrect credentials must be denied")
case codes.ResourceExhausted:
throttled.Add(1)
default:
assert.NoError(t, err)
}
})
}
wg.Wait()
assert.EqualValues(t, 5, checked.Load(), "only the burst budget may reach concurrent credential verification")
assert.EqualValues(t, 15, throttled.Load(), "excess concurrent checks must be throttled")
}
+44 -21
View File
@@ -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)
@@ -130,11 +129,11 @@ func drainEmpty(ch chan *proto.GetMappingUpdateResponse) bool {
func TestSendServiceUpdateToCluster_UniqueTokensPerProxy(t *testing.T) {
ctx := context.Background()
tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
tokenStore: tokenStore,
pkceVerifierStore: pkceStore,
tokenStore: tokenStore,
singleUseStore: singleUseStore,
}
s.SetProxyController(newTestProxyController())
@@ -187,11 +186,11 @@ func TestSendServiceUpdateToCluster_UniqueTokensPerProxy(t *testing.T) {
func TestSendServiceUpdateToCluster_DeleteNoToken(t *testing.T) {
ctx := context.Background()
tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
tokenStore: tokenStore,
pkceVerifierStore: pkceStore,
tokenStore: tokenStore,
singleUseStore: singleUseStore,
}
s.SetProxyController(newTestProxyController())
@@ -221,11 +220,11 @@ func TestSendServiceUpdateToCluster_DeleteNoToken(t *testing.T) {
func TestSendServiceUpdate_UniqueTokensPerProxy(t *testing.T) {
ctx := context.Background()
tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
tokenStore: tokenStore,
pkceVerifierStore: pkceStore,
tokenStore: tokenStore,
singleUseStore: singleUseStore,
}
s.SetProxyController(newTestProxyController())
@@ -273,13 +272,13 @@ func generateState(s *ProxyServiceServer, redirectURL string) string {
func TestOAuthState_NeverTheSame(t *testing.T) {
ctx := context.Background()
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
oidcConfig: ProxyOIDCConfig{
HMACKey: []byte("test-hmac-key"),
},
pkceVerifierStore: pkceStore,
singleUseStore: singleUseStore,
}
redirectURL := "https://app.example.com/callback"
@@ -301,20 +300,20 @@ func TestOAuthState_NeverTheSame(t *testing.T) {
func TestValidateState_RejectsOldTwoPartFormat(t *testing.T) {
ctx := context.Background()
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
oidcConfig: ProxyOIDCConfig{
HMACKey: []byte("test-hmac-key"),
},
pkceVerifierStore: pkceStore,
singleUseStore: singleUseStore,
}
// Old format had only 2 parts: base64(url)|hmac
err := s.pkceVerifierStore.Store("base64url|hmac", "test", 10*time.Minute)
err := s.singleUseStore.Store("base64url|hmac", "test", 10*time.Minute)
require.NoError(t, err)
_, _, err = s.ValidateState("base64url|hmac")
_, _, _, err = s.ValidateState("base64url|hmac")
require.Error(t, err)
assert.Contains(t, err.Error(), "invalid state format")
}
@@ -373,24 +372,48 @@ func TestEnforceAccountScope_AllowsNoTokenInContext(t *testing.T) {
func TestValidateState_RejectsInvalidHMAC(t *testing.T) {
ctx := context.Background()
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
s := &ProxyServiceServer{
oidcConfig: ProxyOIDCConfig{
HMACKey: []byte("test-hmac-key"),
},
pkceVerifierStore: pkceStore,
singleUseStore: singleUseStore,
}
// Store with tampered HMAC
err := s.pkceVerifierStore.Store("dGVzdA==|nonce|wrong-hmac", "test", 10*time.Minute)
err := s.singleUseStore.Store("dGVzdA==|nonce|wrong-hmac", "test", 10*time.Minute)
require.NoError(t, err)
_, _, err = s.ValidateState("dGVzdA==|nonce|wrong-hmac")
_, _, _, err = s.ValidateState("dGVzdA==|nonce|wrong-hmac")
require.Error(t, err)
assert.Contains(t, err.Error(), "invalid state signature")
}
func TestSessionCodeCannotConsumeOIDCState(t *testing.T) {
const verifier = "pkce-verifier"
store := NewSingleUseStore(context.Background(), testCacheStore(t))
server := &ProxyServiceServer{
oidcConfig: ProxyOIDCConfig{
HMACKey: []byte("test-hmac-key"),
},
singleUseStore: store,
}
state := generateState(server, "https://service.example.com/callback")
require.NoError(t, store.Store(state, verifier, time.Minute))
response, err := server.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
SessionCode: state,
})
require.NoError(t, err)
assert.False(t, response.GetValid())
gotVerifier, _, _, err := server.ValidateState(state)
require.NoError(t, err)
assert.Equal(t, verifier, gotVerifier)
}
func TestSendServiceUpdateToCluster_FiltersOnCapability(t *testing.T) {
tokenStore := NewOneTimeTokenStore(context.Background(), testCacheStore(t))
+10 -102
View File
@@ -42,10 +42,10 @@ import (
"github.com/netbirdio/netbird/management/server/auth"
nbContext "github.com/netbirdio/netbird/management/server/context"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/posture"
"github.com/netbirdio/netbird/management/server/settings"
"github.com/netbirdio/netbird/management/server/telemetry"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/management/networkmap/nmdata"
"github.com/netbirdio/netbird/shared/management/proto"
internalStatus "github.com/netbirdio/netbird/shared/management/status"
)
@@ -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()
}
@@ -346,7 +337,8 @@ func (s *Server) Sync(req *proto.EncryptedMessage, srv proto.ManagementService_S
s.syncSem.Add(-1)
return s.handleUpdates(ctx, accountID, peerKey, peer, updates, srv, syncStart)
return PeerUpdateHandlerFactory(peerKey, updates, s.secretsManager, srv, func() { s.cancelPeerRoutines(ctx, accountID, peer, syncStart) }).
WithMetrics(s.appMetrics).HandleUpdates(ctx)
}
func (s *Server) handleHandshake(ctx context.Context, srv proto.ManagementService_JobServer) (wgtypes.Key, error) {
@@ -413,91 +405,6 @@ func (s *Server) sendJobsLoop(ctx context.Context, accountID string, peerKey wgt
}
}
// handleUpdates sends updates to the connected peer until the updates channel is closed.
// It implements a backpressure mechanism that sends the first update immediately,
// then debounces subsequent rapid updates, ensuring only the latest update is sent
// after a quiet period.
func (s *Server) handleUpdates(ctx context.Context, accountID string, peerKey wgtypes.Key, peer *nbpeer.Peer, updates chan *network_map.UpdateMessage, srv proto.ManagementService_SyncServer, streamStartTime time.Time) error {
log.WithContext(ctx).Tracef("starting to handle updates for peer %s", peerKey.String())
// Create a debouncer for this peer connection
debouncer := NewUpdateDebouncer(1000 * time.Millisecond)
defer debouncer.Stop()
for {
select {
// condition when there are some updates
// todo set the updates channel size to 1
case update, open := <-updates:
if s.appMetrics != nil {
s.appMetrics.GRPCMetrics().UpdateChannelQueueLength(len(updates) + 1)
}
if !open {
log.WithContext(ctx).Debugf("updates channel for peer %s was closed", peerKey.String())
s.cancelPeerRoutines(ctx, accountID, peer, streamStartTime)
return nil
}
log.WithContext(ctx).Tracef("received an update for peer %s", peerKey.String())
if debouncer.ProcessUpdate(update) {
// Send immediately (first update or after quiet period)
if err := s.sendUpdate(ctx, accountID, peerKey, peer, update, srv, streamStartTime); err != nil {
log.WithContext(ctx).Debugf("error while sending an update to peer %s: %v", peerKey.String(), err)
return err
}
}
// Timer expired - quiet period reached, send pending updates if any
case <-debouncer.TimerChannel():
pendingUpdates := debouncer.GetPendingUpdates()
if len(pendingUpdates) == 0 {
continue
}
log.WithContext(ctx).Debugf("sending %d debounced update(s) for peer %s", len(pendingUpdates), peerKey.String())
for _, pendingUpdate := range pendingUpdates {
if err := s.sendUpdate(ctx, accountID, peerKey, peer, pendingUpdate, srv, streamStartTime); err != nil {
log.WithContext(ctx).Debugf("error while sending an update to peer %s: %v", peerKey.String(), err)
return err
}
}
// condition when client <-> server connection has been terminated
case <-srv.Context().Done():
// happens when connection drops, e.g. client disconnects
log.WithContext(ctx).Debugf("stream of peer %s has been closed", peerKey.String())
s.cancelPeerRoutines(ctx, accountID, peer, streamStartTime)
return srv.Context().Err()
}
}
}
// sendUpdate encrypts the update message using the peer key and the server's wireguard key,
// then sends the encrypted message to the connected peer via the sync server.
func (s *Server) sendUpdate(ctx context.Context, accountID string, peerKey wgtypes.Key, peer *nbpeer.Peer, update *network_map.UpdateMessage, srv proto.ManagementService_SyncServer, streamStartTime time.Time) error {
key, err := s.secretsManager.GetWGKey()
if err != nil {
s.cancelPeerRoutines(ctx, accountID, peer, streamStartTime)
return status.Errorf(codes.Internal, "failed processing update message")
}
encryptedResp, err := encryption.EncryptMessage(peerKey, key, update.Update)
if err != nil {
s.cancelPeerRoutines(ctx, accountID, peer, streamStartTime)
return status.Errorf(codes.Internal, "failed processing update message")
}
err = srv.Send(&proto.EncryptedMessage{
WgPubKey: key.PublicKey().String(),
Body: encryptedResp,
})
if err != nil {
s.cancelPeerRoutines(ctx, accountID, peer, streamStartTime)
return status.Errorf(codes.Internal, "failed sending update message")
}
log.WithContext(ctx).Tracef("sent an update to peer %s", peerKey.String())
return nil
}
// sendJob encrypts the update message using the peer key and the server's wireguard key,
// then sends the encrypted message to the connected peer via the sync server.
func (s *Server) sendJob(ctx context.Context, peerKey wgtypes.Key, job *job.Event, srv proto.ManagementService_JobServer) error {
@@ -676,6 +583,7 @@ func extractPeerMeta(ctx context.Context, meta *proto.PeerSystemMeta) nbpeer.Pee
RosenpassEnabled: meta.GetFlags().GetRosenpassEnabled(),
RosenpassPermissive: meta.GetFlags().GetRosenpassPermissive(),
ServerSSHAllowed: meta.GetFlags().GetServerSSHAllowed(),
RemoteJobsAllowed: meta.GetFlags().GetRemoteJobsAllowed(),
DisableClientRoutes: meta.GetFlags().GetDisableClientRoutes(),
DisableServerRoutes: meta.GetFlags().GetDisableServerRoutes(),
DisableDNS: meta.GetFlags().GetDisableDNS(),
@@ -902,7 +810,7 @@ func (s *Server) ExtendAuthSession(ctx context.Context, req *proto.EncryptedMess
}, nil
}
func (s *Server) prepareLoginResponse(ctx context.Context, peer *nbpeer.Peer, network *types.Network, postureChecks []*posture.Checks, enableSSH bool) (*proto.LoginResponse, error) {
func (s *Server) prepareLoginResponse(ctx context.Context, peer *nbpeer.Peer, network *types.Network, postureChecks []*nmdata.PostureChecks, enableSSH bool) (*proto.LoginResponse, error) {
var relayToken *Token
var err error
if s.config.Relay != nil && len(s.config.Relay.Addresses) > 0 {
@@ -990,7 +898,7 @@ func (s *Server) IsHealthy(ctx context.Context, req *proto.Empty) (*proto.Empty,
}
// sendInitialSync sends initial proto.SyncResponse to the peer requesting synchronization
func (s *Server) sendInitialSync(ctx context.Context, peerKey wgtypes.Key, peer *nbpeer.Peer, networkMap *types.NetworkMap, postureChecks []*posture.Checks, srv proto.ManagementService_SyncServer, dnsFwdPort int64) error {
func (s *Server) sendInitialSync(ctx context.Context, peerKey wgtypes.Key, peer *nbpeer.Peer, networkMap *types.NetworkMap, postureChecks []*nmdata.PostureChecks, srv proto.ManagementService_SyncServer, dnsFwdPort int64) error {
var err error
var turnToken *Token
@@ -1337,7 +1245,7 @@ func (s *Server) Logout(ctx context.Context, req *proto.EncryptedMessage) (*prot
}
// toProtocolChecks converts posture checks to protocol checks.
func toProtocolChecks(ctx context.Context, postureChecks []*posture.Checks) []*proto.Checks {
func toProtocolChecks(ctx context.Context, postureChecks []*nmdata.PostureChecks) []*proto.Checks {
protoChecks := make([]*proto.Checks, 0, len(postureChecks))
for _, postureCheck := range postureChecks {
check := toProtocolCheck(postureCheck)
@@ -1349,8 +1257,8 @@ func toProtocolChecks(ctx context.Context, postureChecks []*posture.Checks) []*p
return protoChecks
}
// toProtocolCheck converts a posture.Checks to a proto.Checks.
func toProtocolCheck(postureCheck *posture.Checks) *proto.Checks {
// toProtocolCheck converts posture checks to a proto.Checks.
func toProtocolCheck(postureCheck *nmdata.PostureChecks) *proto.Checks {
protoCheck := &proto.Checks{}
if check := postureCheck.Checks.ProcessCheck; check != nil {
@@ -0,0 +1,67 @@
package grpc
import (
"context"
"crypto/rand"
"encoding/base64"
"fmt"
"time"
"github.com/eko/gocache/lib/v4/store"
log "github.com/sirupsen/logrus"
nbcache "github.com/netbirdio/netbird/management/server/cache"
)
// SingleUseStore stores short-lived values that can be retrieved only once.
type SingleUseStore struct {
cache nbcache.Store
ctx context.Context
}
// NewSingleUseStore creates a single-use value store over the shared cache.
func NewSingleUseStore(ctx context.Context, cacheStore nbcache.Store) *SingleUseStore {
return &SingleUseStore{
cache: cacheStore,
ctx: ctx,
}
}
// Store saves value under key with the given TTL, after which it is evicted.
func (s *SingleUseStore) Store(key, value string, ttl time.Duration) error {
if err := s.cache.Set(s.ctx, key, value, store.WithExpiration(ttl)); err != nil {
return fmt.Errorf("store single-use value: %w", err)
}
return nil
}
// Generate stores a value under a namespaced random key and returns the random key.
func (s *SingleUseStore) Generate(namespace, value string, ttl time.Duration) (string, error) {
buf := make([]byte, 32)
if _, err := rand.Read(buf); err != nil {
return "", fmt.Errorf("generate single-use key: %w", err)
}
key := base64.RawURLEncoding.EncodeToString(buf)
if err := s.Store(singleUseCacheKey(namespace, key), value, ttl); err != nil {
return "", err
}
return key, nil
}
func singleUseCacheKey(namespace, key string) string {
return namespace + ":" + key
}
// LoadAndDelete retrieves and removes the value for a key.
func (s *SingleUseStore) LoadAndDelete(key string) (string, bool) {
value, found, err := s.cache.GetDel(s.ctx, key)
if err != nil {
log.Warnf("failed to consume single-use value: %v", err)
return "", false
}
if !found {
return "", false
}
return value, true
}
@@ -0,0 +1,122 @@
package grpc
import (
"context"
"testing"
"time"
)
func TestSingleUseStoreLoadAndDelete(t *testing.T) {
const (
state = "state"
verifier = "verifier"
attempts = 64
)
t.Run("exactly one concurrent caller consumes the verifier", func(t *testing.T) {
store := NewSingleUseStore(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 := NewSingleUseStore(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 := NewSingleUseStore(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 := NewSingleUseStore(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)
}
})
}
func TestSingleUseStore_GenerateAndConsumeOnce(t *testing.T) {
const namespace = "test"
s := NewSingleUseStore(context.Background(), testCacheStore(t))
key, err := s.Generate(namespace, "the-value", time.Minute)
if err != nil {
t.Fatalf("generate: %v", err)
}
if key == "" || key == "the-value" {
t.Fatalf("unexpected key %q", key)
}
value, found := s.LoadAndDelete(singleUseCacheKey(namespace, key))
if !found || value != "the-value" {
t.Fatalf("expected to load the stored value, got %q found=%v", value, found)
}
if _, found := s.LoadAndDelete(singleUseCacheKey(namespace, key)); found {
t.Fatal("value must be consumed on first LoadAndDelete")
}
}
func TestSingleUseStore_GenerateUniqueKeys(t *testing.T) {
s := NewSingleUseStore(context.Background(), testCacheStore(t))
a, err := s.Generate("test", "v", time.Minute)
if err != nil {
t.Fatalf("generate a: %v", err)
}
b, err := s.Generate("test", "v", time.Minute)
if err != nil {
t.Fatalf("generate b: %v", err)
}
if a == b {
t.Fatal("generated keys must be distinct")
}
}
@@ -0,0 +1,70 @@
// Code generated by MockGen. DO NOT EDIT.
// Source: ./peer_update_handler.go
//
// Generated by this command:
//
// mockgen -source=./peer_update_handler.go -destination=./sync_sender_mock.go -package=grpc
//
// Package grpc is a generated GoMock package.
package grpc
import (
context "context"
reflect "reflect"
proto "github.com/netbirdio/netbird/shared/management/proto"
gomock "go.uber.org/mock/gomock"
)
// MocksyncSender is a mock of syncSender interface.
type MocksyncSender struct {
ctrl *gomock.Controller
recorder *MocksyncSenderMockRecorder
isgomock struct{}
}
// MocksyncSenderMockRecorder is the mock recorder for MocksyncSender.
type MocksyncSenderMockRecorder struct {
mock *MocksyncSender
}
// NewMocksyncSender creates a new mock instance.
func NewMocksyncSender(ctrl *gomock.Controller) *MocksyncSender {
mock := &MocksyncSender{ctrl: ctrl}
mock.recorder = &MocksyncSenderMockRecorder{mock}
return mock
}
// EXPECT returns an object that allows the caller to indicate expected use.
func (m *MocksyncSender) EXPECT() *MocksyncSenderMockRecorder {
return m.recorder
}
// Context mocks base method.
func (m *MocksyncSender) Context() context.Context {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "Context")
ret0, _ := ret[0].(context.Context)
return ret0
}
// Context indicates an expected call of Context.
func (mr *MocksyncSenderMockRecorder) Context() *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Context", reflect.TypeOf((*MocksyncSender)(nil).Context))
}
// Send mocks base method.
func (m *MocksyncSender) Send(arg0 *proto.EncryptedMessage) error {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "Send", arg0)
ret0, _ := ret[0].(error)
return ret0
}
// Send indicates an expected call of Send.
func (mr *MocksyncSenderMockRecorder) Send(arg0 any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Send", reflect.TypeOf((*MocksyncSender)(nil).Send), arg0)
}
@@ -25,6 +25,8 @@ import (
const defaultDuration = 12 * time.Hour
// SecretsManager used to manage TURN and relay secrets
//
//go:generate go tool mockgen -source=./token_mgr.go -destination=./token_mgr_mock.go -package=grpc
type SecretsManager interface {
GenerateTurnToken() (*Token, error)
GenerateRelayToken() (*Token, error)
@@ -0,0 +1,111 @@
// Code generated by MockGen. DO NOT EDIT.
// Source: ./token_mgr.go
//
// Generated by this command:
//
// mockgen -source=./token_mgr.go -destination=./token_mgr_mock.go -package=grpc
//
// Package grpc is a generated GoMock package.
package grpc
import (
context "context"
reflect "reflect"
gomock "go.uber.org/mock/gomock"
wgtypes "golang.zx2c4.com/wireguard/wgctrl/wgtypes"
)
// MockSecretsManager is a mock of SecretsManager interface.
type MockSecretsManager struct {
ctrl *gomock.Controller
recorder *MockSecretsManagerMockRecorder
isgomock struct{}
}
// MockSecretsManagerMockRecorder is the mock recorder for MockSecretsManager.
type MockSecretsManagerMockRecorder struct {
mock *MockSecretsManager
}
// NewMockSecretsManager creates a new mock instance.
func NewMockSecretsManager(ctrl *gomock.Controller) *MockSecretsManager {
mock := &MockSecretsManager{ctrl: ctrl}
mock.recorder = &MockSecretsManagerMockRecorder{mock}
return mock
}
// EXPECT returns an object that allows the caller to indicate expected use.
func (m *MockSecretsManager) EXPECT() *MockSecretsManagerMockRecorder {
return m.recorder
}
// CancelRefresh mocks base method.
func (m *MockSecretsManager) CancelRefresh(peerKey string) {
m.ctrl.T.Helper()
m.ctrl.Call(m, "CancelRefresh", peerKey)
}
// CancelRefresh indicates an expected call of CancelRefresh.
func (mr *MockSecretsManagerMockRecorder) CancelRefresh(peerKey any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CancelRefresh", reflect.TypeOf((*MockSecretsManager)(nil).CancelRefresh), peerKey)
}
// GenerateRelayToken mocks base method.
func (m *MockSecretsManager) GenerateRelayToken() (*Token, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GenerateRelayToken")
ret0, _ := ret[0].(*Token)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GenerateRelayToken indicates an expected call of GenerateRelayToken.
func (mr *MockSecretsManagerMockRecorder) GenerateRelayToken() *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GenerateRelayToken", reflect.TypeOf((*MockSecretsManager)(nil).GenerateRelayToken))
}
// GenerateTurnToken mocks base method.
func (m *MockSecretsManager) GenerateTurnToken() (*Token, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GenerateTurnToken")
ret0, _ := ret[0].(*Token)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GenerateTurnToken indicates an expected call of GenerateTurnToken.
func (mr *MockSecretsManagerMockRecorder) GenerateTurnToken() *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GenerateTurnToken", reflect.TypeOf((*MockSecretsManager)(nil).GenerateTurnToken))
}
// GetWGKey mocks base method.
func (m *MockSecretsManager) GetWGKey() (wgtypes.Key, error) {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetWGKey")
ret0, _ := ret[0].(wgtypes.Key)
ret1, _ := ret[1].(error)
return ret0, ret1
}
// GetWGKey indicates an expected call of GetWGKey.
func (mr *MockSecretsManagerMockRecorder) GetWGKey() *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetWGKey", reflect.TypeOf((*MockSecretsManager)(nil).GetWGKey))
}
// SetupRefresh mocks base method.
func (m *MockSecretsManager) SetupRefresh(ctx context.Context, accountID, peerKey string) {
m.ctrl.T.Helper()
m.ctrl.Call(m, "SetupRefresh", ctx, accountID, peerKey)
}
// SetupRefresh indicates an expected call of SetupRefresh.
func (mr *MockSecretsManagerMockRecorder) SetupRefresh(ctx, accountID, peerKey any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetupRefresh", reflect.TypeOf((*MockSecretsManager)(nil).SetupRefresh), ctx, accountID, peerKey)
}
@@ -6,6 +6,14 @@ import (
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
)
//go:generate go tool mockgen -source=./update_debouncer.go -destination=./update_debouncer_mock.go -package=grpc
type Debouncer interface {
Stop()
TimerChannel() <-chan time.Time
ProcessUpdate(update *network_map.UpdateMessage) bool
GetPendingUpdates() []*network_map.UpdateMessage
}
// UpdateDebouncer implements a backpressure mechanism that:
// - Sends the first update immediately
// - Coalesces rapid subsequent network map updates (only latest matters)
@@ -0,0 +1,96 @@
// Code generated by MockGen. DO NOT EDIT.
// Source: ./update_debouncer.go
//
// Generated by this command:
//
// mockgen -source=./update_debouncer.go -destination=./update_debouncer_mock.go -package=grpc
//
// Package grpc is a generated GoMock package.
package grpc
import (
reflect "reflect"
time "time"
network_map "github.com/netbirdio/netbird/management/internals/controllers/network_map"
gomock "go.uber.org/mock/gomock"
)
// MockDebouncer is a mock of Debouncer interface.
type MockDebouncer struct {
ctrl *gomock.Controller
recorder *MockDebouncerMockRecorder
isgomock struct{}
}
// MockDebouncerMockRecorder is the mock recorder for MockDebouncer.
type MockDebouncerMockRecorder struct {
mock *MockDebouncer
}
// NewMockDebouncer creates a new mock instance.
func NewMockDebouncer(ctrl *gomock.Controller) *MockDebouncer {
mock := &MockDebouncer{ctrl: ctrl}
mock.recorder = &MockDebouncerMockRecorder{mock}
return mock
}
// EXPECT returns an object that allows the caller to indicate expected use.
func (m *MockDebouncer) EXPECT() *MockDebouncerMockRecorder {
return m.recorder
}
// GetPendingUpdates mocks base method.
func (m *MockDebouncer) GetPendingUpdates() []*network_map.UpdateMessage {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "GetPendingUpdates")
ret0, _ := ret[0].([]*network_map.UpdateMessage)
return ret0
}
// GetPendingUpdates indicates an expected call of GetPendingUpdates.
func (mr *MockDebouncerMockRecorder) GetPendingUpdates() *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPendingUpdates", reflect.TypeOf((*MockDebouncer)(nil).GetPendingUpdates))
}
// ProcessUpdate mocks base method.
func (m *MockDebouncer) ProcessUpdate(update *network_map.UpdateMessage) bool {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "ProcessUpdate", update)
ret0, _ := ret[0].(bool)
return ret0
}
// ProcessUpdate indicates an expected call of ProcessUpdate.
func (mr *MockDebouncerMockRecorder) ProcessUpdate(update any) *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ProcessUpdate", reflect.TypeOf((*MockDebouncer)(nil).ProcessUpdate), update)
}
// Stop mocks base method.
func (m *MockDebouncer) Stop() {
m.ctrl.T.Helper()
m.ctrl.Call(m, "Stop")
}
// Stop indicates an expected call of Stop.
func (mr *MockDebouncerMockRecorder) Stop() *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Stop", reflect.TypeOf((*MockDebouncer)(nil).Stop))
}
// TimerChannel mocks base method.
func (m *MockDebouncer) TimerChannel() <-chan time.Time {
m.ctrl.T.Helper()
ret := m.ctrl.Call(m, "TimerChannel")
ret0, _ := ret[0].(<-chan time.Time)
return ret0
}
// TimerChannel indicates an expected call of TimerChannel.
func (mr *MockDebouncerMockRecorder) TimerChannel() *gomock.Call {
mr.mock.ctrl.T.Helper()
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "TimerChannel", reflect.TypeOf((*MockDebouncer)(nil).TimerChannel))
}
@@ -40,9 +40,9 @@ func setupValidateSessionTest(t *testing.T) *validateSessionTestSetup {
proxyManager := &testValidateSessionProxyManager{}
tokenStore := NewOneTimeTokenStore(ctx, testCacheStore(t))
pkceStore := NewPKCEVerifierStore(ctx, testCacheStore(t))
singleUseStore := NewSingleUseStore(ctx, testCacheStore(t))
proxyService := NewProxyServiceServer(nil, tokenStore, pkceStore, ProxyOIDCConfig{}, nil, usersManager, nil, proxyManager, nil)
proxyService := NewProxyServiceServer(nil, tokenStore, singleUseStore, ProxyOIDCConfig{}, nil, usersManager, nil, proxyManager, nil)
proxyService.SetServiceManager(serviceManager)
createTestProxies(t, ctx, testStore)
@@ -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
}
@@ -519,7 +570,7 @@ func (m *testValidateSessionServiceManager) DeleteAccountCluster(_ context.Conte
type testValidateSessionProxyManager struct{}
func (m *testValidateSessionProxyManager) Connect(_ context.Context, _, _, _, _ string, _ *string, _ *proxy.Capabilities) (*proxy.Proxy, error) {
func (m *testValidateSessionProxyManager) Connect(_ context.Context, _, _, _, _, _ string, _ *string, _ *proxy.Capabilities) (*proxy.Proxy, error) {
return nil, nil
}
@@ -583,6 +634,10 @@ func (m *testValidateSessionProxyManager) ClusterSupportsPrivate(_ context.Conte
return nil
}
func (m *testValidateSessionProxyManager) ClusterSupportsSessionCode(_ context.Context, _ string) bool {
return false
}
type testValidateSessionUsersManager struct {
store store.Store
}
@@ -611,3 +666,47 @@ func (m *testValidateSessionUsersManager) GetUserWithGroups(ctx context.Context,
}
return user, groups, nil
}
func TestValidateSession_RedeemsSessionCode(t *testing.T) {
setup := setupValidateSessionTest(t)
defer setup.cleanup()
proxy, err := setup.store.GetServiceByID(context.Background(), store.LockingStrengthNone, "testAccountId", "testProxyId")
require.NoError(t, err)
token := createSessionToken(t, proxy.SessionPrivateKey, "allowedUserId", "test-proxy.example.com")
code, ok := setup.proxyService.GenerateSessionCode(token)
require.True(t, ok)
require.NotEqual(t, token, code, "code must not be the token itself")
resp, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
Domain: "test-proxy.example.com",
SessionCode: code,
})
require.NoError(t, err)
assert.True(t, resp.Valid, "redeemed code should authorize the user")
assert.Equal(t, "allowedUserId", resp.UserId)
assert.Equal(t, token, resp.GetSessionToken(), "response must carry the durable token for the cookie")
// Single-use: the same code must not redeem again.
resp2, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
Domain: "test-proxy.example.com",
SessionCode: code,
})
require.NoError(t, err)
assert.False(t, resp2.Valid, "a consumed code must be rejected")
assert.Empty(t, resp2.GetSessionToken())
}
func TestValidateSession_InvalidSessionCode(t *testing.T) {
setup := setupValidateSessionTest(t)
defer setup.cleanup()
resp, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
Domain: "test-proxy.example.com",
SessionCode: "does-not-exist",
})
require.NoError(t, err)
assert.False(t, resp.Valid)
assert.Empty(t, resp.GetSessionToken())
}