mirror of
https://github.com/netbirdio/netbird.git
synced 2026-08-19 22:21:29 +02:00
Compare commits
8 Commits
reverse-pr
...
dmitri-upd
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
ad44842161 | ||
|
|
a262ee0ee2 | ||
|
|
3c1b6afe8f | ||
|
|
b9a1999c0d | ||
|
|
79511d8caf | ||
|
|
9efa3c6579 | ||
|
|
77791b5858 | ||
|
|
ad98b99fc5 |
@@ -152,6 +152,7 @@ func NewClient(androidSDKVersion int, deviceName string, uiVersion string, tunAd
|
||||
execWorkaround(androidSDKVersion)
|
||||
|
||||
net.SetAndroidProtectSocketFn(tunAdapter.ProtectSocket)
|
||||
system.SetIFaceDiscover(iFaceDiscover)
|
||||
return &Client{
|
||||
deviceName: deviceName,
|
||||
uiVersion: uiVersion,
|
||||
|
||||
@@ -6,7 +6,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.opentelemetry.io/otel"
|
||||
"google.golang.org/grpc"
|
||||
|
||||
@@ -6,7 +6,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/grpc"
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/google/gopacket"
|
||||
"github.com/google/gopacket/layers"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
@@ -4,7 +4,7 @@ import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/google/gopacket/layers"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
@@ -4,7 +4,7 @@ import (
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/google/gopacket"
|
||||
"github.com/google/gopacket/layers"
|
||||
|
||||
|
||||
@@ -8,7 +8,7 @@ import (
|
||||
"net/netip"
|
||||
reflect "reflect"
|
||||
|
||||
gomock "github.com/golang/mock/gomock"
|
||||
gomock "go.uber.org/mock/gomock"
|
||||
)
|
||||
|
||||
// MockPacketFilter is a mock of PacketFilter interface.
|
||||
|
||||
@@ -8,7 +8,7 @@ import (
|
||||
os "os"
|
||||
reflect "reflect"
|
||||
|
||||
gomock "github.com/golang/mock/gomock"
|
||||
gomock "go.uber.org/mock/gomock"
|
||||
tun "golang.zx2c4.com/wireguard/tun"
|
||||
)
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ import (
|
||||
"net/netip"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@ package mocks
|
||||
import (
|
||||
reflect "reflect"
|
||||
|
||||
gomock "github.com/golang/mock/gomock"
|
||||
gomock "go.uber.org/mock/gomock"
|
||||
wgdevice "golang.zx2c4.com/wireguard/device"
|
||||
|
||||
"github.com/netbirdio/netbird/client/iface/device"
|
||||
|
||||
@@ -4,7 +4,7 @@ import (
|
||||
"net"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/google/gopacket"
|
||||
"github.com/google/gopacket/layers"
|
||||
"github.com/miekg/dns"
|
||||
|
||||
@@ -9,7 +9,7 @@ import (
|
||||
"os"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/miekg/dns"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
|
||||
|
||||
@@ -12,7 +12,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/google/uuid"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
@@ -37,23 +37,32 @@
|
||||
// Updater Process (Setup):
|
||||
//
|
||||
// 1. Receives parameters from service via command-line arguments
|
||||
// 2. Runs installer with appropriate silent/quiet flags:
|
||||
// 2. Terminates the UI so the installer does not have to replace a locked image
|
||||
// file, which would otherwise leave the install needing a reboot
|
||||
// 3. Runs installer with appropriate silent/quiet flags:
|
||||
// - Windows EXE: installer.exe /S
|
||||
// - Windows MSI: msiexec.exe /i installer.msi /quiet /qn /l*v msi.log
|
||||
// - Windows MSI: msiexec.exe /i installer.msi /qn /norestart REBOOT=ReallySuppress /l*v msi.log
|
||||
// - macOS PKG: installer -pkg installer.pkg -target /
|
||||
// - macOS Homebrew: brew upgrade netbirdio/tap/netbird
|
||||
// 3. Installer terminates daemon and UI processes
|
||||
// 4. Installer replaces binaries with new version
|
||||
// 5. Updater waits for installer to complete
|
||||
// 6. Updater restarts daemon:
|
||||
// 4. Installer terminates the daemon
|
||||
// 5. Installer replaces binaries with new version
|
||||
// 6. Updater waits for installer to complete. On Windows, MSI exit codes 3010
|
||||
// (ERROR_SUCCESS_REBOOT_REQUIRED) and 1641 (ERROR_SUCCESS_REBOOT_INITIATED)
|
||||
// are a pending-reboot outcome, not a failure: the install succeeded, but
|
||||
// some files are only replaced on the next restart (the reboot itself is
|
||||
// suppressed via /norestart and REBOOT=ReallySuppress), and the flow
|
||||
// continues as on success
|
||||
// 7. Updater restarts daemon:
|
||||
// - Windows: netbird.exe service start
|
||||
// - macOS/Linux: netbird service start
|
||||
// 7. Updater restarts UI:
|
||||
// - Windows: Launches netbird-ui.exe as active console user using CreateProcessAsUser
|
||||
// 8. Updater restarts UI:
|
||||
// - Windows: Launches netbird-ui.exe using CreateProcessAsUser in every
|
||||
// session it was terminated in, falling back to the active console session
|
||||
// - macOS: Uses launchctl asuser to launch NetBird.app for console user
|
||||
// - Linux: Not implemented (UI typically auto-starts)
|
||||
// 8. Updater writes result.json with success/error status
|
||||
// 9. Updater process exits
|
||||
// 9. Updater writes result.json with success/error status (a pending reboot is
|
||||
// recorded as success)
|
||||
// 10. Updater process exits
|
||||
//
|
||||
// # Result Communication
|
||||
//
|
||||
|
||||
@@ -2,6 +2,7 @@ package installer
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
@@ -22,6 +23,12 @@ const (
|
||||
|
||||
msiLogFile = "msi.log"
|
||||
|
||||
// ERROR_SUCCESS_REBOOT_REQUIRED and ERROR_SUCCESS_REBOOT_INITIATED
|
||||
msiRebootRequired = 3010
|
||||
msiRebootInitiated = 1641
|
||||
|
||||
processExitWait = 10 * time.Second
|
||||
|
||||
msiDownloadURL = "https://github.com/netbirdio/netbird/releases/download/v%version/netbird_installer_%version_windows_%arch.msi"
|
||||
exeDownloadURL = "https://github.com/netbirdio/netbird/releases/download/v%version/netbird_installer_%version_windows_%arch.exe"
|
||||
)
|
||||
@@ -38,6 +45,8 @@ var (
|
||||
func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string, daemonFolder string) (resultErr error) {
|
||||
resultHandler := NewResultHandler(u.tempDir)
|
||||
|
||||
var uiSessions []uint32
|
||||
|
||||
// Always ensure daemon and UI are restarted after setup
|
||||
defer func() {
|
||||
log.Infof("starting daemon back")
|
||||
@@ -46,7 +55,7 @@ func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string
|
||||
}
|
||||
|
||||
log.Infof("starting UI back")
|
||||
if err := u.startUIAsUser(daemonFolder); err != nil {
|
||||
if err := u.startUI(daemonFolder, uiSessions); err != nil {
|
||||
log.Errorf("failed to start UI: %v", err)
|
||||
}
|
||||
|
||||
@@ -75,6 +84,14 @@ func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string
|
||||
return
|
||||
}
|
||||
|
||||
// The UI holds an open handle on its own image. Left running, Restart Manager
|
||||
// cannot shut it down (msiexec runs as LocalSystem here, the UI as the
|
||||
// interactive user), so the MSI falls back to replacing the file on reboot and
|
||||
// marks the install as restart-required. The deferred close-application action
|
||||
// in the package runs too late to prevent that, it happens after
|
||||
// InstallValidate has already registered the file as in use.
|
||||
uiSessions = killUI()
|
||||
|
||||
var cmd *exec.Cmd
|
||||
switch installerType {
|
||||
case TypeExe:
|
||||
@@ -84,7 +101,9 @@ func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string
|
||||
installerDir := filepath.Dir(installerFile)
|
||||
logPath := filepath.Join(installerDir, msiLogFile)
|
||||
log.Infof("run msi installer: %s", installerFile)
|
||||
cmd = exec.CommandContext(ctx, "msiexec.exe", "/i", filepath.Base(installerFile), "/quiet", "/qn", "/l*v", logPath)
|
||||
// REBOOT=ReallySuppress: a silent install has no way to ask, so without it
|
||||
// msiexec reboots the machine on its own if it decides one is needed.
|
||||
cmd = exec.CommandContext(ctx, "msiexec.exe", "/i", filepath.Base(installerFile), "/qn", "/norestart", "REBOOT=ReallySuppress", "/l*v", logPath)
|
||||
}
|
||||
|
||||
cmd.Dir = filepath.Dir(installerFile)
|
||||
@@ -95,9 +114,13 @@ func (u *Installer) Setup(ctx context.Context, dryRun bool, installerFile string
|
||||
}
|
||||
|
||||
log.Infof("installer started with PID %d", cmd.Process.Pid)
|
||||
if resultErr = cmd.Wait(); resultErr != nil {
|
||||
log.Errorf("installer process finished with error: %v", resultErr)
|
||||
return
|
||||
if err := cmd.Wait(); err != nil {
|
||||
if !isRebootPending(err) {
|
||||
resultErr = err
|
||||
log.Errorf("installer process finished with error: %v", err)
|
||||
return
|
||||
}
|
||||
log.Warnf("installer completed but reported a pending reboot, some files will be replaced on the next restart")
|
||||
}
|
||||
|
||||
return nil
|
||||
@@ -117,16 +140,142 @@ func (u *Installer) startDaemon(daemonFolder string) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (u *Installer) startUIAsUser(daemonFolder string) error {
|
||||
func (u *Installer) startUI(daemonFolder string, sessionIDs []uint32) error {
|
||||
uiPath := filepath.Join(daemonFolder, uiName)
|
||||
log.Infof("starting netbird-ui: %s", uiPath)
|
||||
|
||||
// Get the active console session ID
|
||||
sessionID := windows.WTSGetActiveConsoleSessionId()
|
||||
if sessionID == 0xFFFFFFFF {
|
||||
return fmt.Errorf("no active user session found")
|
||||
if len(sessionIDs) == 0 {
|
||||
sessionID := windows.WTSGetActiveConsoleSessionId()
|
||||
if sessionID == 0xFFFFFFFF {
|
||||
return fmt.Errorf("no active user session found")
|
||||
}
|
||||
sessionIDs = []uint32{sessionID}
|
||||
}
|
||||
|
||||
var errs []error
|
||||
for _, sessionID := range sessionIDs {
|
||||
if err := startUIInSession(uiPath, sessionID); err != nil {
|
||||
errs = append(errs, fmt.Errorf("session %d: %w", sessionID, err))
|
||||
continue
|
||||
}
|
||||
log.Infof("netbird-ui started successfully in session %d", sessionID)
|
||||
}
|
||||
return errors.Join(errs...)
|
||||
}
|
||||
|
||||
// isRebootPending reports whether the installer exit code means it succeeded but
|
||||
// left work for the next restart. The reboot itself is suppressed, so this is not
|
||||
// a failure.
|
||||
func isRebootPending(err error) bool {
|
||||
var exitErr *exec.ExitError
|
||||
if !errors.As(err, &exitErr) {
|
||||
return false
|
||||
}
|
||||
|
||||
switch exitErr.ExitCode() {
|
||||
case msiRebootRequired, msiRebootInitiated:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// killUI terminates any running netbird-ui process and returns the IDs of the
|
||||
// interactive sessions the terminated processes belonged to. Setup starts the
|
||||
// UI again in those sessions once the installer is done.
|
||||
func killUI() []uint32 {
|
||||
pids, err := processIDsByName(uiName)
|
||||
if err != nil {
|
||||
log.Warnf("failed to look up %s processes: %v", uiName, err)
|
||||
return nil
|
||||
}
|
||||
|
||||
sessions := make(map[uint32]struct{})
|
||||
for _, pid := range pids {
|
||||
var sessionID uint32
|
||||
if err := windows.ProcessIdToSessionId(pid, &sessionID); err != nil {
|
||||
log.Warnf("failed to look up session of %s (PID %d): %v", uiName, pid, err)
|
||||
}
|
||||
|
||||
if err := terminateProcess(pid); err != nil {
|
||||
log.Warnf("failed to terminate %s (PID %d): %v", uiName, pid, err)
|
||||
continue
|
||||
}
|
||||
log.Infof("terminated %s (PID %d) in session %d", uiName, pid, sessionID)
|
||||
|
||||
if sessionID != 0 {
|
||||
sessions[sessionID] = struct{}{}
|
||||
}
|
||||
}
|
||||
|
||||
sessionIDs := make([]uint32, 0, len(sessions))
|
||||
for sessionID := range sessions {
|
||||
sessionIDs = append(sessionIDs, sessionID)
|
||||
}
|
||||
return sessionIDs
|
||||
}
|
||||
|
||||
func processIDsByName(name string) ([]uint32, error) {
|
||||
snapshot, err := windows.CreateToolhelp32Snapshot(windows.TH32CS_SNAPPROCESS, 0)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("create process snapshot: %w", err)
|
||||
}
|
||||
defer func() {
|
||||
if err := windows.CloseHandle(snapshot); err != nil {
|
||||
log.Warnf("failed to close process snapshot: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
var entry windows.ProcessEntry32
|
||||
entry.Size = uint32(unsafe.Sizeof(entry))
|
||||
|
||||
var pids []uint32
|
||||
for err = windows.Process32First(snapshot, &entry); err == nil; err = windows.Process32Next(snapshot, &entry) {
|
||||
if strings.EqualFold(windows.UTF16ToString(entry.ExeFile[:]), name) {
|
||||
pids = append(pids, entry.ProcessID)
|
||||
}
|
||||
}
|
||||
if !errors.Is(err, windows.ERROR_NO_MORE_FILES) {
|
||||
return nil, fmt.Errorf("enumerate processes: %w", err)
|
||||
}
|
||||
|
||||
return pids, nil
|
||||
}
|
||||
|
||||
func terminateProcess(pid uint32) error {
|
||||
handle, err := windows.OpenProcess(windows.PROCESS_TERMINATE|windows.SYNCHRONIZE, false, pid)
|
||||
if err != nil {
|
||||
// The process may have exited between enumeration and now.
|
||||
if errors.Is(err, windows.ERROR_INVALID_PARAMETER) {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("open process: %w", err)
|
||||
}
|
||||
defer func() {
|
||||
if err := windows.CloseHandle(handle); err != nil {
|
||||
log.Warnf("failed to close process handle: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
if err := windows.TerminateProcess(handle, 0); err != nil {
|
||||
return fmt.Errorf("terminate process: %w", err)
|
||||
}
|
||||
|
||||
// Wait for the handle to signal so the image file is released before the
|
||||
// installer tries to overwrite it. A timeout is reported through the returned
|
||||
// event, not through err, which stays nil unless the wait itself failed.
|
||||
event, err := windows.WaitForSingleObject(handle, uint32(processExitWait.Milliseconds()))
|
||||
if err != nil {
|
||||
return fmt.Errorf("wait for process exit: %w", err)
|
||||
}
|
||||
if event != windows.WAIT_OBJECT_0 {
|
||||
return fmt.Errorf("wait for process exit: unexpected wait result %#x", event)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func startUIInSession(uiPath string, sessionID uint32) error {
|
||||
// Get the user token for that session
|
||||
var userToken windows.Token
|
||||
err := windows.WTSQueryUserToken(sessionID, &userToken)
|
||||
@@ -158,6 +307,16 @@ func (u *Installer) startUIAsUser(daemonFolder string) error {
|
||||
}
|
||||
}()
|
||||
|
||||
var env *uint16
|
||||
if err := windows.CreateEnvironmentBlock(&env, primaryToken, false); err != nil {
|
||||
return fmt.Errorf("create environment block: %w", err)
|
||||
}
|
||||
defer func() {
|
||||
if err := windows.DestroyEnvironmentBlock(env); err != nil {
|
||||
log.Warnf("failed to destroy environment block: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
// Prepare startup info
|
||||
var si windows.StartupInfo
|
||||
si.Cb = uint32(unsafe.Sizeof(si))
|
||||
@@ -180,7 +339,7 @@ func (u *Installer) startUIAsUser(daemonFolder string) error {
|
||||
nil,
|
||||
false,
|
||||
creationFlags,
|
||||
nil,
|
||||
env,
|
||||
nil,
|
||||
&si,
|
||||
&pi,
|
||||
@@ -197,7 +356,6 @@ func (u *Installer) startUIAsUser(daemonFolder string) error {
|
||||
log.Warnf("failed to close thread handle: %v", err)
|
||||
}
|
||||
|
||||
log.Infof("netbird-ui started successfully in session %d", sessionID)
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
108
client/internal/updater/installer/installer_run_windows_test.go
Normal file
108
client/internal/updater/installer/installer_run_windows_test.go
Normal file
@@ -0,0 +1,108 @@
|
||||
package installer
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"os/exec"
|
||||
"slices"
|
||||
"strconv"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// exitErrorWithCode returns a real *exec.ExitError carrying the given exit code.
|
||||
func exitErrorWithCode(t *testing.T, code int) error {
|
||||
t.Helper()
|
||||
|
||||
err := exec.Command("cmd.exe", "/c", "exit "+strconv.Itoa(code)).Run()
|
||||
if err == nil {
|
||||
t.Fatalf("expected a non-zero exit for code %d", code)
|
||||
}
|
||||
return err
|
||||
}
|
||||
|
||||
func TestIsRebootPending(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
code int
|
||||
want bool
|
||||
}{
|
||||
{name: "reboot required", code: msiRebootRequired, want: true},
|
||||
{name: "reboot initiated", code: msiRebootInitiated, want: true},
|
||||
{name: "generic failure", code: 1603, want: false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := isRebootPending(exitErrorWithCode(t, tt.code)); got != tt.want {
|
||||
t.Errorf("isRebootPending(exit %d) = %v, want %v", tt.code, got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestProcessIDsByNameAndTerminate spawns a long-running system process, finds it
|
||||
// by name and terminates it, covering the path the updater uses to release the UI
|
||||
// image file before the installer replaces it.
|
||||
func TestProcessIDsByNameAndTerminate(t *testing.T) {
|
||||
cmd := exec.Command("ping.exe", "-n", "60", "127.0.0.1")
|
||||
if err := cmd.Start(); err != nil {
|
||||
t.Fatalf("start ping: %v", err)
|
||||
}
|
||||
|
||||
pid := uint32(cmd.Process.Pid)
|
||||
killed := false
|
||||
t.Cleanup(func() {
|
||||
if !killed {
|
||||
_ = cmd.Process.Kill()
|
||||
}
|
||||
_ = cmd.Wait()
|
||||
})
|
||||
|
||||
// Name matching must be case-insensitive: the snapshot reports PING.EXE.
|
||||
pids, err := processIDsByName("ping.exe")
|
||||
if err != nil {
|
||||
t.Fatalf("processIDsByName: %v", err)
|
||||
}
|
||||
|
||||
if !slices.Contains(pids, pid) {
|
||||
t.Fatalf("PID %d not among the ping.exe processes found: %v", pid, pids)
|
||||
}
|
||||
|
||||
if err := terminateProcess(pid); err != nil {
|
||||
t.Fatalf("terminateProcess: %v", err)
|
||||
}
|
||||
killed = true
|
||||
|
||||
// terminateProcess only returns once the handle has signalled, so the process
|
||||
// is already gone and Wait must not block. It exits with the code passed to
|
||||
// TerminateProcess, which is 0, so Wait reports no error.
|
||||
if err := cmd.Wait(); err != nil {
|
||||
t.Fatalf("wait for terminated ping: %v", err)
|
||||
}
|
||||
if !cmd.ProcessState.Exited() {
|
||||
t.Error("process did not exit after terminateProcess")
|
||||
}
|
||||
|
||||
remaining, err := processIDsByName("ping.exe")
|
||||
if err != nil {
|
||||
t.Fatalf("processIDsByName after terminate: %v", err)
|
||||
}
|
||||
if slices.Contains(remaining, pid) {
|
||||
t.Errorf("PID %d still listed after terminateProcess", pid)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProcessIDsByNameNoMatch(t *testing.T) {
|
||||
pids, err := processIDsByName("netbird-nonexistent-process.exe")
|
||||
if err != nil {
|
||||
t.Fatalf("processIDsByName: %v", err)
|
||||
}
|
||||
if len(pids) != 0 {
|
||||
t.Errorf("expected no matches, got %v", pids)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsRebootPendingNonExitError(t *testing.T) {
|
||||
if isRebootPending(errors.New("start installer: file not found")) {
|
||||
t.Error("a non-exit error must not be treated as a pending reboot")
|
||||
}
|
||||
}
|
||||
@@ -10,7 +10,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.opentelemetry.io/otel"
|
||||
|
||||
|
||||
@@ -30,6 +30,11 @@ func GetInfo(ctx context.Context) *Info {
|
||||
kernelVersion = osInfo[2]
|
||||
}
|
||||
|
||||
addrs, err := networkAddresses()
|
||||
if err != nil {
|
||||
log.Warnf("discover network addresses: %s", err)
|
||||
}
|
||||
|
||||
gio := &Info{
|
||||
GoOS: runtime.GOOS,
|
||||
Kernel: kernel,
|
||||
@@ -41,6 +46,7 @@ func GetInfo(ctx context.Context) *Info {
|
||||
NetbirdVersion: version.NetbirdVersion(),
|
||||
UIVersion: extractUIVersion(ctx),
|
||||
KernelVersion: kernelVersion,
|
||||
NetworkAddresses: addrs,
|
||||
SystemSerialNumber: serial(),
|
||||
SystemProductName: productModel(),
|
||||
SystemManufacturer: productManufacturer(),
|
||||
|
||||
@@ -1,4 +1,4 @@
|
||||
//go:build !ios
|
||||
//go:build !ios && !android
|
||||
|
||||
package system
|
||||
|
||||
|
||||
89
client/system/network_addr_android.go
Normal file
89
client/system/network_addr_android.go
Normal file
@@ -0,0 +1,89 @@
|
||||
package system
|
||||
|
||||
import (
|
||||
"net/netip"
|
||||
"strings"
|
||||
)
|
||||
|
||||
var iFaceDiscover IFaceDiscover
|
||||
|
||||
type IFaceDiscover interface {
|
||||
IFaces() (string, error)
|
||||
}
|
||||
|
||||
// SetIFaceDiscover configures the Android interface discovery provider.
|
||||
func SetIFaceDiscover(discover IFaceDiscover) {
|
||||
iFaceDiscover = discover
|
||||
}
|
||||
|
||||
func networkAddresses() ([]NetworkAddress, error) {
|
||||
if iFaceDiscover == nil {
|
||||
return nil, nil
|
||||
}
|
||||
ifaces, err := iFaceDiscover.IFaces()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
var netAddresses []NetworkAddress
|
||||
for _, line := range strings.Split(ifaces, "\n") {
|
||||
addresses, ok := interfaceAddresses(line)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
for _, address := range addresses {
|
||||
netAddr, ok := toNetworkAddress(address)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if isDuplicated(netAddresses, netAddr) {
|
||||
continue
|
||||
}
|
||||
netAddresses = append(netAddresses, netAddr)
|
||||
}
|
||||
}
|
||||
return netAddresses, nil
|
||||
}
|
||||
|
||||
func interfaceAddresses(line string) ([]string, bool) {
|
||||
parts := strings.Split(line, "|")
|
||||
if len(parts) != 2 {
|
||||
return nil, false
|
||||
}
|
||||
flags := strings.Fields(parts[0])
|
||||
if len(flags) != 8 {
|
||||
return nil, false
|
||||
}
|
||||
up, loopback := flags[3], flags[5]
|
||||
if up != "true" || loopback == "true" {
|
||||
return nil, false
|
||||
}
|
||||
return strings.Fields(parts[1]), true
|
||||
}
|
||||
|
||||
func toNetworkAddress(address string) (NetworkAddress, bool) {
|
||||
prefix, err := netip.ParsePrefix(address)
|
||||
if err != nil {
|
||||
return NetworkAddress{}, false
|
||||
}
|
||||
if prefix.Addr().Is4In6() {
|
||||
if prefix.Bits() < 96 {
|
||||
return NetworkAddress{}, false
|
||||
}
|
||||
prefix = netip.PrefixFrom(prefix.Addr().Unmap(), prefix.Bits()-96)
|
||||
}
|
||||
ip := prefix.Addr()
|
||||
if ip.IsLoopback() || ip.IsLinkLocalUnicast() || ip.IsMulticast() {
|
||||
return NetworkAddress{}, false
|
||||
}
|
||||
return NetworkAddress{NetIP: prefix}, true
|
||||
}
|
||||
|
||||
func isDuplicated(addresses []NetworkAddress, addr NetworkAddress) bool {
|
||||
for _, duplicated := range addresses {
|
||||
if duplicated.NetIP == addr.NetIP {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -1,4 +1,4 @@
|
||||
//go:build !ios
|
||||
//go:build !ios && !android
|
||||
|
||||
package system
|
||||
|
||||
|
||||
4
go.mod
4
go.mod
@@ -62,7 +62,6 @@ require (
|
||||
github.com/goccy/go-yaml v1.18.0
|
||||
github.com/godbus/dbus/v5 v5.2.2
|
||||
github.com/golang-jwt/jwt/v5 v5.3.1
|
||||
github.com/golang/mock v1.6.0
|
||||
github.com/google/go-cmp v0.7.0
|
||||
github.com/google/gopacket v1.1.19
|
||||
github.com/google/nftables v0.3.0
|
||||
@@ -217,6 +216,7 @@ require (
|
||||
github.com/gobwas/pool v0.2.1 // indirect
|
||||
github.com/gogo/protobuf v1.3.2 // indirect
|
||||
github.com/golang-jwt/jwt/v4 v4.5.2 // indirect
|
||||
github.com/golang/mock v1.6.0 // indirect
|
||||
github.com/google/btree v1.1.3 // indirect
|
||||
github.com/google/go-querystring v1.1.0 // indirect
|
||||
github.com/google/go-tpm v0.9.8 // indirect
|
||||
@@ -340,3 +340,5 @@ replace github.com/dexidp/dex/api/v2 => github.com/netbirdio/dex/api/v2 v2.0.0-2
|
||||
replace github.com/mailru/easyjson => github.com/netbirdio/easyjson v0.9.0
|
||||
|
||||
replace github.com/wailsapp/wails/v3 => github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260810103952-24e716aea4db
|
||||
|
||||
tool go.uber.org/mock/mockgen
|
||||
|
||||
@@ -281,8 +281,6 @@ read_enable_crowdsec() {
|
||||
echo "CrowdSec checks client IPs against a community threat intelligence database" > /dev/stderr
|
||||
echo "and blocks known malicious sources before they reach your services." > /dev/stderr
|
||||
echo "A local CrowdSec LAPI container will be added to your deployment." > /dev/stderr
|
||||
echo "It also enables the AppSec (WAF) endpoint, so services can inspect HTTP" > /dev/stderr
|
||||
echo "requests for exploits. Both stay off per service until you enable them." > /dev/stderr
|
||||
echo -n "Enable CrowdSec? [y/N]: " > /dev/stderr
|
||||
read -r CHOICE < /dev/tty
|
||||
|
||||
@@ -556,8 +554,7 @@ generate_configuration_files() {
|
||||
# TCP ServersTransport for PROXY protocol v2 to the proxy backend
|
||||
render_traefik_dynamic > traefik-dynamic.yaml
|
||||
if [[ "$ENABLE_CROWDSEC" == "true" ]]; then
|
||||
mkdir -p crowdsec/acquis.d
|
||||
render_crowdsec_appsec_acquis > crowdsec/acquis.d/appsec.yaml
|
||||
mkdir -p crowdsec
|
||||
fi
|
||||
fi
|
||||
;;
|
||||
@@ -592,23 +589,6 @@ generate_configuration_files() {
|
||||
return 0
|
||||
}
|
||||
|
||||
# The AppSec (WAF) listener only exists if an appsec acquisition datasource is
|
||||
# configured. One datasource is one listener carrying one merged rule set: the
|
||||
# protocol has no rule-set selector, so per-service rule variation would need
|
||||
# either a second datasource on another port or pre_eval hooks filtering on
|
||||
# req.Host.
|
||||
render_crowdsec_appsec_acquis() {
|
||||
cat <<EOF
|
||||
source: appsec
|
||||
listen_addr: 0.0.0.0:7422
|
||||
appsec_configs:
|
||||
- crowdsecurity/appsec-default
|
||||
labels:
|
||||
type: appsec
|
||||
EOF
|
||||
return 0
|
||||
}
|
||||
|
||||
start_services_and_show_instructions() {
|
||||
# For built-in Traefik, start containers immediately
|
||||
# For NPM, start containers first (NPM needs services running to create proxy)
|
||||
@@ -825,11 +805,7 @@ render_docker_compose_traefik_builtin() {
|
||||
restart: unless-stopped
|
||||
networks: [netbird]
|
||||
environment:
|
||||
# appsec-generic-rules is required alongside appsec-virtual-patching:
|
||||
# the appsec-default config references crowdsecurity/generic-* and
|
||||
# crowdsecurity/experimental-*, which only that collection provides, and
|
||||
# the engine exits at startup if they are missing.
|
||||
COLLECTIONS: crowdsecurity/linux crowdsecurity/appsec-virtual-patching crowdsecurity/appsec-generic-rules
|
||||
COLLECTIONS: crowdsecurity/linux
|
||||
volumes:
|
||||
- ./crowdsec:/etc/crowdsec
|
||||
- crowdsec_db:/var/lib/crowdsec/data
|
||||
@@ -1095,11 +1071,6 @@ EOF
|
||||
cat <<EOF
|
||||
NB_PROXY_CROWDSEC_API_URL=http://crowdsec:8080
|
||||
NB_PROXY_CROWDSEC_API_KEY=$CROWDSEC_BOUNCER_KEY
|
||||
# AppSec (WAF) request inspection. Separate endpoint from the LAPI above and
|
||||
# validated with the same bouncer key. Setting it makes the proxy advertise the
|
||||
# AppSec capability, which is what lets a service select appsec_mode; nothing is
|
||||
# inspected until a service opts in.
|
||||
NB_PROXY_CROWDSEC_APPSEC_URL=http://crowdsec:7422/
|
||||
EOF
|
||||
fi
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
package network_map
|
||||
|
||||
//go:generate go run go.uber.org/mock/mockgen -package network_map -destination=interface_mock.go -source=./interface.go -build_flags=-mod=mod
|
||||
//go:generate go tool mockgen -package network_map -destination=interface_mock.go -source=./interface.go -build_flags=-mod=mod
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
@@ -10,7 +10,7 @@ import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/gorilla/mux"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
@@ -6,7 +6,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@ import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -6,7 +6,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ import (
|
||||
"encoding/json"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -7,7 +7,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
package peers
|
||||
|
||||
//go:generate go run github.com/golang/mock/mockgen -package peers -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
|
||||
//go:generate go tool mockgen -package peers -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
@@ -1,5 +1,10 @@
|
||||
// Code generated by MockGen. DO NOT EDIT.
|
||||
// Source: ./manager.go
|
||||
//
|
||||
// Generated by this command:
|
||||
//
|
||||
// mockgen -package peers -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
|
||||
//
|
||||
|
||||
// Package peers is a generated GoMock package.
|
||||
package peers
|
||||
@@ -9,18 +14,19 @@ import (
|
||||
net "net"
|
||||
reflect "reflect"
|
||||
|
||||
gomock "github.com/golang/mock/gomock"
|
||||
network_map "github.com/netbirdio/netbird/management/internals/controllers/network_map"
|
||||
account "github.com/netbirdio/netbird/management/server/account"
|
||||
integrated_validator "github.com/netbirdio/netbird/management/server/integrations/integrated_validator"
|
||||
peer "github.com/netbirdio/netbird/management/server/peer"
|
||||
types "github.com/netbirdio/netbird/management/server/types"
|
||||
gomock "go.uber.org/mock/gomock"
|
||||
)
|
||||
|
||||
// MockManager is a mock of Manager interface.
|
||||
type MockManager struct {
|
||||
ctrl *gomock.Controller
|
||||
recorder *MockManagerMockRecorder
|
||||
isgomock struct{}
|
||||
}
|
||||
|
||||
// MockManagerMockRecorder is the mock recorder for MockManager.
|
||||
@@ -49,7 +55,7 @@ func (m *MockManager) CreateProxyPeer(ctx context.Context, accountID, peerKey, c
|
||||
}
|
||||
|
||||
// CreateProxyPeer indicates an expected call of CreateProxyPeer.
|
||||
func (mr *MockManagerMockRecorder) CreateProxyPeer(ctx, accountID, peerKey, cluster interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) CreateProxyPeer(ctx, accountID, peerKey, cluster any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CreateProxyPeer", reflect.TypeOf((*MockManager)(nil).CreateProxyPeer), ctx, accountID, peerKey, cluster)
|
||||
}
|
||||
@@ -63,7 +69,7 @@ func (m *MockManager) DeletePeers(ctx context.Context, accountID string, peerIDs
|
||||
}
|
||||
|
||||
// DeletePeers indicates an expected call of DeletePeers.
|
||||
func (mr *MockManagerMockRecorder) DeletePeers(ctx, accountID, peerIDs, userID, checkConnected interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) DeletePeers(ctx, accountID, peerIDs, userID, checkConnected any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeletePeers", reflect.TypeOf((*MockManager)(nil).DeletePeers), ctx, accountID, peerIDs, userID, checkConnected)
|
||||
}
|
||||
@@ -78,7 +84,7 @@ func (m *MockManager) GetAllPeers(ctx context.Context, accountID, userID string)
|
||||
}
|
||||
|
||||
// GetAllPeers indicates an expected call of GetAllPeers.
|
||||
func (mr *MockManagerMockRecorder) GetAllPeers(ctx, accountID, userID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetAllPeers(ctx, accountID, userID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAllPeers", reflect.TypeOf((*MockManager)(nil).GetAllPeers), ctx, accountID, userID)
|
||||
}
|
||||
@@ -93,7 +99,7 @@ func (m *MockManager) GetPeer(ctx context.Context, accountID, userID, peerID str
|
||||
}
|
||||
|
||||
// GetPeer indicates an expected call of GetPeer.
|
||||
func (mr *MockManagerMockRecorder) GetPeer(ctx, accountID, userID, peerID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetPeer(ctx, accountID, userID, peerID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeer", reflect.TypeOf((*MockManager)(nil).GetPeer), ctx, accountID, userID, peerID)
|
||||
}
|
||||
@@ -108,7 +114,7 @@ func (m *MockManager) GetPeerAccountID(ctx context.Context, peerID string) (stri
|
||||
}
|
||||
|
||||
// GetPeerAccountID indicates an expected call of GetPeerAccountID.
|
||||
func (mr *MockManagerMockRecorder) GetPeerAccountID(ctx, peerID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetPeerAccountID(ctx, peerID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeerAccountID", reflect.TypeOf((*MockManager)(nil).GetPeerAccountID), ctx, peerID)
|
||||
}
|
||||
@@ -123,7 +129,7 @@ func (m *MockManager) GetPeerByTunnelIP(ctx context.Context, accountID string, i
|
||||
}
|
||||
|
||||
// GetPeerByTunnelIP indicates an expected call of GetPeerByTunnelIP.
|
||||
func (mr *MockManagerMockRecorder) GetPeerByTunnelIP(ctx, accountID, ip interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetPeerByTunnelIP(ctx, accountID, ip any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeerByTunnelIP", reflect.TypeOf((*MockManager)(nil).GetPeerByTunnelIP), ctx, accountID, ip)
|
||||
}
|
||||
@@ -138,7 +144,7 @@ func (m *MockManager) GetPeerID(ctx context.Context, peerKey string) (string, er
|
||||
}
|
||||
|
||||
// GetPeerID indicates an expected call of GetPeerID.
|
||||
func (mr *MockManagerMockRecorder) GetPeerID(ctx, peerKey interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetPeerID(ctx, peerKey any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeerID", reflect.TypeOf((*MockManager)(nil).GetPeerID), ctx, peerKey)
|
||||
}
|
||||
@@ -154,7 +160,7 @@ func (m *MockManager) GetPeerWithGroups(ctx context.Context, accountID, peerID s
|
||||
}
|
||||
|
||||
// GetPeerWithGroups indicates an expected call of GetPeerWithGroups.
|
||||
func (mr *MockManagerMockRecorder) GetPeerWithGroups(ctx, accountID, peerID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetPeerWithGroups(ctx, accountID, peerID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeerWithGroups", reflect.TypeOf((*MockManager)(nil).GetPeerWithGroups), ctx, accountID, peerID)
|
||||
}
|
||||
@@ -169,7 +175,7 @@ func (m *MockManager) GetPeersByGroupIDs(ctx context.Context, accountID string,
|
||||
}
|
||||
|
||||
// GetPeersByGroupIDs indicates an expected call of GetPeersByGroupIDs.
|
||||
func (mr *MockManagerMockRecorder) GetPeersByGroupIDs(ctx, accountID, groupsIDs interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetPeersByGroupIDs(ctx, accountID, groupsIDs any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPeersByGroupIDs", reflect.TypeOf((*MockManager)(nil).GetPeersByGroupIDs), ctx, accountID, groupsIDs)
|
||||
}
|
||||
@@ -181,7 +187,7 @@ func (m *MockManager) SetAccountManager(accountManager account.Manager) {
|
||||
}
|
||||
|
||||
// SetAccountManager indicates an expected call of SetAccountManager.
|
||||
func (mr *MockManagerMockRecorder) SetAccountManager(accountManager interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) SetAccountManager(accountManager any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetAccountManager", reflect.TypeOf((*MockManager)(nil).SetAccountManager), accountManager)
|
||||
}
|
||||
@@ -193,7 +199,7 @@ func (m *MockManager) SetIntegratedPeerValidator(integratedPeerValidator integra
|
||||
}
|
||||
|
||||
// SetIntegratedPeerValidator indicates an expected call of SetIntegratedPeerValidator.
|
||||
func (mr *MockManagerMockRecorder) SetIntegratedPeerValidator(integratedPeerValidator interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) SetIntegratedPeerValidator(integratedPeerValidator any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetIntegratedPeerValidator", reflect.TypeOf((*MockManager)(nil).SetIntegratedPeerValidator), integratedPeerValidator)
|
||||
}
|
||||
@@ -205,7 +211,7 @@ func (m *MockManager) SetNetworkMapController(networkMapController network_map.C
|
||||
}
|
||||
|
||||
// SetNetworkMapController indicates an expected call of SetNetworkMapController.
|
||||
func (mr *MockManagerMockRecorder) SetNetworkMapController(networkMapController interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) SetNetworkMapController(networkMapController any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetNetworkMapController", reflect.TypeOf((*MockManager)(nil).SetNetworkMapController), networkMapController)
|
||||
}
|
||||
|
||||
@@ -5,7 +5,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -23,9 +23,6 @@ type Domain struct {
|
||||
// SupportsCrowdSec is populated at query time from proxy cluster capabilities.
|
||||
// Not persisted.
|
||||
SupportsCrowdSec *bool `gorm:"-"`
|
||||
// SupportsAppSec is populated at query time from proxy cluster capabilities.
|
||||
// Not persisted.
|
||||
SupportsAppSec *bool `gorm:"-"`
|
||||
// SupportsPrivate is populated at query time from proxy cluster capabilities. Not persisted.
|
||||
SupportsPrivate *bool `gorm:"-"`
|
||||
}
|
||||
|
||||
@@ -49,7 +49,6 @@ func domainToApi(d *domain.Domain) api.ReverseProxyDomain {
|
||||
SupportsCustomPorts: d.SupportsCustomPorts,
|
||||
RequireSubdomain: d.RequireSubdomain,
|
||||
SupportsCrowdsec: d.SupportsCrowdSec,
|
||||
SupportsAppsec: d.SupportsAppSec,
|
||||
SupportsPrivate: d.SupportsPrivate,
|
||||
}
|
||||
if d.TargetCluster != "" {
|
||||
|
||||
@@ -39,7 +39,6 @@ type proxyManager interface {
|
||||
ClusterSupportsCustomPorts(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterSupportsAppSec(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool
|
||||
}
|
||||
|
||||
@@ -99,7 +98,6 @@ func (m Manager) GetDomains(ctx context.Context, accountID, userID string) ([]*d
|
||||
d.SupportsCustomPorts = m.proxyManager.ClusterSupportsCustomPorts(ctx, cluster)
|
||||
d.RequireSubdomain = m.proxyManager.ClusterRequireSubdomain(ctx, cluster)
|
||||
d.SupportsCrowdSec = m.proxyManager.ClusterSupportsCrowdSec(ctx, cluster)
|
||||
d.SupportsAppSec = m.proxyManager.ClusterSupportsAppSec(ctx, cluster)
|
||||
d.SupportsPrivate = m.proxyManager.ClusterSupportsPrivate(ctx, cluster)
|
||||
ret = append(ret, d)
|
||||
}
|
||||
@@ -117,7 +115,6 @@ func (m Manager) GetDomains(ctx context.Context, accountID, userID string) ([]*d
|
||||
if d.TargetCluster != "" {
|
||||
cd.SupportsCustomPorts = m.proxyManager.ClusterSupportsCustomPorts(ctx, d.TargetCluster)
|
||||
cd.SupportsCrowdSec = m.proxyManager.ClusterSupportsCrowdSec(ctx, d.TargetCluster)
|
||||
cd.SupportsAppSec = m.proxyManager.ClusterSupportsAppSec(ctx, d.TargetCluster)
|
||||
cd.SupportsPrivate = m.proxyManager.ClusterSupportsPrivate(ctx, d.TargetCluster)
|
||||
}
|
||||
// Custom domains never require a subdomain by default since
|
||||
|
||||
@@ -46,10 +46,6 @@ func (m *mockProxyManager) ClusterSupportsCrowdSec(_ context.Context, _ string)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *mockProxyManager) ClusterSupportsAppSec(_ context.Context, _ string) *bool {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *mockProxyManager) ClusterSupportsPrivate(_ context.Context, _ string) *bool {
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
package proxy
|
||||
|
||||
//go:generate go run github.com/golang/mock/mockgen -package proxy -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
|
||||
//go:generate go tool mockgen -package proxy -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
|
||||
|
||||
import (
|
||||
"context"
|
||||
@@ -19,7 +19,6 @@ type Manager interface {
|
||||
ClusterSupportsCustomPorts(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterSupportsAppSec(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool
|
||||
CleanupStale(ctx context.Context, inactivityDuration time.Duration) error
|
||||
GetAccountProxy(ctx context.Context, accountID string) (*Proxy, error)
|
||||
|
||||
@@ -21,7 +21,6 @@ type store interface {
|
||||
GetClusterSupportsCustomPorts(ctx context.Context, clusterAddr string) *bool
|
||||
GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool
|
||||
GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool
|
||||
GetClusterSupportsAppSec(ctx context.Context, clusterAddr string) *bool
|
||||
GetClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool
|
||||
CleanupStaleProxies(ctx context.Context, inactivityDuration time.Duration) error
|
||||
GetProxyByAccountID(ctx context.Context, accountID string) (*proxy.Proxy, error)
|
||||
@@ -139,13 +138,6 @@ func (m Manager) ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string
|
||||
return m.store.GetClusterSupportsCrowdSec(ctx, clusterAddr)
|
||||
}
|
||||
|
||||
// ClusterSupportsAppSec returns whether all active proxies in the cluster have
|
||||
// a CrowdSec AppSec endpoint configured (unanimous). Returns nil when no proxy
|
||||
// has reported capabilities.
|
||||
func (m Manager) ClusterSupportsAppSec(ctx context.Context, clusterAddr string) *bool {
|
||||
return m.store.GetClusterSupportsAppSec(ctx, clusterAddr)
|
||||
}
|
||||
|
||||
// ClusterSupportsPrivate reports whether any active proxy claims the private capability (nil = unreported).
|
||||
func (m Manager) ClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool {
|
||||
return m.store.GetClusterSupportsPrivate(ctx, clusterAddr)
|
||||
|
||||
@@ -99,9 +99,6 @@ func (m *mockStore) GetClusterRequireSubdomain(_ context.Context, _ string) *boo
|
||||
func (m *mockStore) GetClusterSupportsCrowdSec(_ context.Context, _ string) *bool {
|
||||
return nil
|
||||
}
|
||||
func (m *mockStore) GetClusterSupportsAppSec(_ context.Context, _ string) *bool {
|
||||
return nil
|
||||
}
|
||||
func (m *mockStore) GetClusterSupportsPrivate(_ context.Context, _ string) *bool {
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -1,5 +1,10 @@
|
||||
// Code generated by MockGen. DO NOT EDIT.
|
||||
// Source: ./manager.go
|
||||
//
|
||||
// Generated by this command:
|
||||
//
|
||||
// mockgen -package proxy -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
|
||||
//
|
||||
|
||||
// Package proxy is a generated GoMock package.
|
||||
package proxy
|
||||
@@ -9,14 +14,15 @@ import (
|
||||
reflect "reflect"
|
||||
time "time"
|
||||
|
||||
gomock "github.com/golang/mock/gomock"
|
||||
proto "github.com/netbirdio/netbird/shared/management/proto"
|
||||
gomock "go.uber.org/mock/gomock"
|
||||
)
|
||||
|
||||
// MockManager is a mock of Manager interface.
|
||||
type MockManager struct {
|
||||
ctrl *gomock.Controller
|
||||
recorder *MockManagerMockRecorder
|
||||
isgomock struct{}
|
||||
}
|
||||
|
||||
// MockManagerMockRecorder is the mock recorder for MockManager.
|
||||
@@ -45,7 +51,7 @@ func (m *MockManager) CleanupStale(ctx context.Context, inactivityDuration time.
|
||||
}
|
||||
|
||||
// CleanupStale indicates an expected call of CleanupStale.
|
||||
func (mr *MockManagerMockRecorder) CleanupStale(ctx, inactivityDuration interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) CleanupStale(ctx, inactivityDuration any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CleanupStale", reflect.TypeOf((*MockManager)(nil).CleanupStale), ctx, inactivityDuration)
|
||||
}
|
||||
@@ -59,25 +65,11 @@ func (m *MockManager) ClusterRequireSubdomain(ctx context.Context, clusterAddr s
|
||||
}
|
||||
|
||||
// ClusterRequireSubdomain indicates an expected call of ClusterRequireSubdomain.
|
||||
func (mr *MockManagerMockRecorder) ClusterRequireSubdomain(ctx, clusterAddr interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) ClusterRequireSubdomain(ctx, clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterRequireSubdomain", reflect.TypeOf((*MockManager)(nil).ClusterRequireSubdomain), ctx, clusterAddr)
|
||||
}
|
||||
|
||||
// ClusterSupportsAppSec mocks base method.
|
||||
func (m *MockManager) ClusterSupportsAppSec(ctx context.Context, clusterAddr string) *bool {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "ClusterSupportsAppSec", ctx, clusterAddr)
|
||||
ret0, _ := ret[0].(*bool)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// ClusterSupportsAppSec indicates an expected call of ClusterSupportsAppSec.
|
||||
func (mr *MockManagerMockRecorder) ClusterSupportsAppSec(ctx, clusterAddr interface{}) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsAppSec", reflect.TypeOf((*MockManager)(nil).ClusterSupportsAppSec), ctx, clusterAddr)
|
||||
}
|
||||
|
||||
// ClusterSupportsCrowdSec mocks base method.
|
||||
func (m *MockManager) ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool {
|
||||
m.ctrl.T.Helper()
|
||||
@@ -87,7 +79,7 @@ func (m *MockManager) ClusterSupportsCrowdSec(ctx context.Context, clusterAddr s
|
||||
}
|
||||
|
||||
// ClusterSupportsCrowdSec indicates an expected call of ClusterSupportsCrowdSec.
|
||||
func (mr *MockManagerMockRecorder) ClusterSupportsCrowdSec(ctx, clusterAddr interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) ClusterSupportsCrowdSec(ctx, clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsCrowdSec", reflect.TypeOf((*MockManager)(nil).ClusterSupportsCrowdSec), ctx, clusterAddr)
|
||||
}
|
||||
@@ -101,7 +93,7 @@ func (m *MockManager) ClusterSupportsCustomPorts(ctx context.Context, clusterAdd
|
||||
}
|
||||
|
||||
// ClusterSupportsCustomPorts indicates an expected call of ClusterSupportsCustomPorts.
|
||||
func (mr *MockManagerMockRecorder) ClusterSupportsCustomPorts(ctx, clusterAddr interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) ClusterSupportsCustomPorts(ctx, clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsCustomPorts", reflect.TypeOf((*MockManager)(nil).ClusterSupportsCustomPorts), ctx, clusterAddr)
|
||||
}
|
||||
@@ -115,7 +107,7 @@ func (m *MockManager) ClusterSupportsPrivate(ctx context.Context, clusterAddr st
|
||||
}
|
||||
|
||||
// ClusterSupportsPrivate indicates an expected call of ClusterSupportsPrivate.
|
||||
func (mr *MockManagerMockRecorder) ClusterSupportsPrivate(ctx, clusterAddr interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) ClusterSupportsPrivate(ctx, clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ClusterSupportsPrivate", reflect.TypeOf((*MockManager)(nil).ClusterSupportsPrivate), ctx, clusterAddr)
|
||||
}
|
||||
@@ -130,7 +122,7 @@ func (m *MockManager) Connect(ctx context.Context, proxyID, sessionID, clusterAd
|
||||
}
|
||||
|
||||
// Connect indicates an expected call of Connect.
|
||||
func (mr *MockManagerMockRecorder) Connect(ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) Connect(ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Connect", reflect.TypeOf((*MockManager)(nil).Connect), ctx, proxyID, sessionID, clusterAddress, ipAddress, accountID, capabilities)
|
||||
}
|
||||
@@ -145,7 +137,7 @@ func (m *MockManager) CountAccountProxies(ctx context.Context, accountID string)
|
||||
}
|
||||
|
||||
// CountAccountProxies indicates an expected call of CountAccountProxies.
|
||||
func (mr *MockManagerMockRecorder) CountAccountProxies(ctx, accountID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) CountAccountProxies(ctx, accountID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CountAccountProxies", reflect.TypeOf((*MockManager)(nil).CountAccountProxies), ctx, accountID)
|
||||
}
|
||||
@@ -159,7 +151,7 @@ func (m *MockManager) DeleteAccountCluster(ctx context.Context, clusterAddress,
|
||||
}
|
||||
|
||||
// DeleteAccountCluster indicates an expected call of DeleteAccountCluster.
|
||||
func (mr *MockManagerMockRecorder) DeleteAccountCluster(ctx, clusterAddress, accountID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) DeleteAccountCluster(ctx, clusterAddress, accountID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAccountCluster", reflect.TypeOf((*MockManager)(nil).DeleteAccountCluster), ctx, clusterAddress, accountID)
|
||||
}
|
||||
@@ -173,7 +165,7 @@ func (m *MockManager) Disconnect(ctx context.Context, proxyID, sessionID string)
|
||||
}
|
||||
|
||||
// Disconnect indicates an expected call of Disconnect.
|
||||
func (mr *MockManagerMockRecorder) Disconnect(ctx, proxyID, sessionID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) Disconnect(ctx, proxyID, sessionID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Disconnect", reflect.TypeOf((*MockManager)(nil).Disconnect), ctx, proxyID, sessionID)
|
||||
}
|
||||
@@ -188,7 +180,7 @@ func (m *MockManager) GetAccountProxy(ctx context.Context, accountID string) (*P
|
||||
}
|
||||
|
||||
// GetAccountProxy indicates an expected call of GetAccountProxy.
|
||||
func (mr *MockManagerMockRecorder) GetAccountProxy(ctx, accountID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetAccountProxy(ctx, accountID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountProxy", reflect.TypeOf((*MockManager)(nil).GetAccountProxy), ctx, accountID)
|
||||
}
|
||||
@@ -203,7 +195,7 @@ func (m *MockManager) GetActiveClusterAddresses(ctx context.Context) ([]string,
|
||||
}
|
||||
|
||||
// GetActiveClusterAddresses indicates an expected call of GetActiveClusterAddresses.
|
||||
func (mr *MockManagerMockRecorder) GetActiveClusterAddresses(ctx interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetActiveClusterAddresses(ctx any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetActiveClusterAddresses", reflect.TypeOf((*MockManager)(nil).GetActiveClusterAddresses), ctx)
|
||||
}
|
||||
@@ -218,7 +210,7 @@ func (m *MockManager) GetActiveClusterAddressesForAccount(ctx context.Context, a
|
||||
}
|
||||
|
||||
// GetActiveClusterAddressesForAccount indicates an expected call of GetActiveClusterAddressesForAccount.
|
||||
func (mr *MockManagerMockRecorder) GetActiveClusterAddressesForAccount(ctx, accountID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetActiveClusterAddressesForAccount(ctx, accountID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetActiveClusterAddressesForAccount", reflect.TypeOf((*MockManager)(nil).GetActiveClusterAddressesForAccount), ctx, accountID)
|
||||
}
|
||||
@@ -232,7 +224,7 @@ func (m *MockManager) Heartbeat(ctx context.Context, p *Proxy) error {
|
||||
}
|
||||
|
||||
// Heartbeat indicates an expected call of Heartbeat.
|
||||
func (mr *MockManagerMockRecorder) Heartbeat(ctx, p interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) Heartbeat(ctx, p any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "Heartbeat", reflect.TypeOf((*MockManager)(nil).Heartbeat), ctx, p)
|
||||
}
|
||||
@@ -247,7 +239,7 @@ func (m *MockManager) IsClusterAddressAvailable(ctx context.Context, clusterAddr
|
||||
}
|
||||
|
||||
// IsClusterAddressAvailable indicates an expected call of IsClusterAddressAvailable.
|
||||
func (mr *MockManagerMockRecorder) IsClusterAddressAvailable(ctx, clusterAddress, accountID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) IsClusterAddressAvailable(ctx, clusterAddress, accountID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "IsClusterAddressAvailable", reflect.TypeOf((*MockManager)(nil).IsClusterAddressAvailable), ctx, clusterAddress, accountID)
|
||||
}
|
||||
@@ -256,6 +248,7 @@ func (mr *MockManagerMockRecorder) IsClusterAddressAvailable(ctx, clusterAddress
|
||||
type MockController struct {
|
||||
ctrl *gomock.Controller
|
||||
recorder *MockControllerMockRecorder
|
||||
isgomock struct{}
|
||||
}
|
||||
|
||||
// MockControllerMockRecorder is the mock recorder for MockController.
|
||||
@@ -298,7 +291,7 @@ func (m *MockController) GetProxiesForCluster(clusterAddr string) []string {
|
||||
}
|
||||
|
||||
// GetProxiesForCluster indicates an expected call of GetProxiesForCluster.
|
||||
func (mr *MockControllerMockRecorder) GetProxiesForCluster(clusterAddr interface{}) *gomock.Call {
|
||||
func (mr *MockControllerMockRecorder) GetProxiesForCluster(clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetProxiesForCluster", reflect.TypeOf((*MockController)(nil).GetProxiesForCluster), clusterAddr)
|
||||
}
|
||||
@@ -312,7 +305,7 @@ func (m *MockController) RegisterProxyToCluster(ctx context.Context, clusterAddr
|
||||
}
|
||||
|
||||
// RegisterProxyToCluster indicates an expected call of RegisterProxyToCluster.
|
||||
func (mr *MockControllerMockRecorder) RegisterProxyToCluster(ctx, clusterAddr, proxyID interface{}) *gomock.Call {
|
||||
func (mr *MockControllerMockRecorder) RegisterProxyToCluster(ctx, clusterAddr, proxyID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RegisterProxyToCluster", reflect.TypeOf((*MockController)(nil).RegisterProxyToCluster), ctx, clusterAddr, proxyID)
|
||||
}
|
||||
@@ -324,7 +317,7 @@ func (m *MockController) SendServiceUpdateToCluster(ctx context.Context, account
|
||||
}
|
||||
|
||||
// SendServiceUpdateToCluster indicates an expected call of SendServiceUpdateToCluster.
|
||||
func (mr *MockControllerMockRecorder) SendServiceUpdateToCluster(ctx, accountID, update, clusterAddr interface{}) *gomock.Call {
|
||||
func (mr *MockControllerMockRecorder) SendServiceUpdateToCluster(ctx, accountID, update, clusterAddr any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SendServiceUpdateToCluster", reflect.TypeOf((*MockController)(nil).SendServiceUpdateToCluster), ctx, accountID, update, clusterAddr)
|
||||
}
|
||||
@@ -338,7 +331,7 @@ func (m *MockController) UnregisterProxyFromCluster(ctx context.Context, cluster
|
||||
}
|
||||
|
||||
// UnregisterProxyFromCluster indicates an expected call of UnregisterProxyFromCluster.
|
||||
func (mr *MockControllerMockRecorder) UnregisterProxyFromCluster(ctx, clusterAddr, proxyID interface{}) *gomock.Call {
|
||||
func (mr *MockControllerMockRecorder) UnregisterProxyFromCluster(ctx, clusterAddr, proxyID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UnregisterProxyFromCluster", reflect.TypeOf((*MockController)(nil).UnregisterProxyFromCluster), ctx, clusterAddr, proxyID)
|
||||
}
|
||||
|
||||
@@ -20,9 +20,6 @@ type Capabilities struct {
|
||||
RequireSubdomain *bool
|
||||
// SupportsCrowdsec indicates whether this proxy has CrowdSec configured.
|
||||
SupportsCrowdsec *bool
|
||||
// SupportsAppsec indicates whether this proxy has a CrowdSec AppSec (WAF)
|
||||
// endpoint configured.
|
||||
SupportsAppsec *bool
|
||||
// Private indicates whether this proxy supports inbound access via Wireguard
|
||||
// tunnel and netbird-only authentication policies
|
||||
Private *bool
|
||||
@@ -77,6 +74,5 @@ type Cluster struct {
|
||||
SupportsCustomPorts *bool
|
||||
RequireSubdomain *bool
|
||||
SupportsCrowdSec *bool
|
||||
SupportsAppSec *bool
|
||||
Private *bool
|
||||
}
|
||||
|
||||
@@ -9,7 +9,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/gorilla/mux"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
package service
|
||||
|
||||
//go:generate go run github.com/golang/mock/mockgen -package service -destination=interface_mock.go -source=./interface.go -build_flags=-mod=mod
|
||||
//go:generate go tool mockgen -package service -destination=interface_mock.go -source=./interface.go -build_flags=-mod=mod
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
@@ -1,5 +1,10 @@
|
||||
// Code generated by MockGen. DO NOT EDIT.
|
||||
// Source: ./interface.go
|
||||
//
|
||||
// Generated by this command:
|
||||
//
|
||||
// mockgen -package service -destination=interface_mock.go -source=./interface.go -build_flags=-mod=mod
|
||||
//
|
||||
|
||||
// Package service is a generated GoMock package.
|
||||
package service
|
||||
@@ -8,14 +13,15 @@ import (
|
||||
context "context"
|
||||
reflect "reflect"
|
||||
|
||||
gomock "github.com/golang/mock/gomock"
|
||||
proxy "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
|
||||
gomock "go.uber.org/mock/gomock"
|
||||
)
|
||||
|
||||
// MockManager is a mock of Manager interface.
|
||||
type MockManager struct {
|
||||
ctrl *gomock.Controller
|
||||
recorder *MockManagerMockRecorder
|
||||
isgomock struct{}
|
||||
}
|
||||
|
||||
// MockManagerMockRecorder is the mock recorder for MockManager.
|
||||
@@ -45,7 +51,7 @@ func (m *MockManager) CreateService(ctx context.Context, accountID, userID strin
|
||||
}
|
||||
|
||||
// CreateService indicates an expected call of CreateService.
|
||||
func (mr *MockManagerMockRecorder) CreateService(ctx, accountID, userID, service interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) CreateService(ctx, accountID, userID, service any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CreateService", reflect.TypeOf((*MockManager)(nil).CreateService), ctx, accountID, userID, service)
|
||||
}
|
||||
@@ -60,7 +66,7 @@ func (m *MockManager) CreateServiceFromPeer(ctx context.Context, accountID, peer
|
||||
}
|
||||
|
||||
// CreateServiceFromPeer indicates an expected call of CreateServiceFromPeer.
|
||||
func (mr *MockManagerMockRecorder) CreateServiceFromPeer(ctx, accountID, peerID, req interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) CreateServiceFromPeer(ctx, accountID, peerID, req any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CreateServiceFromPeer", reflect.TypeOf((*MockManager)(nil).CreateServiceFromPeer), ctx, accountID, peerID, req)
|
||||
}
|
||||
@@ -74,7 +80,7 @@ func (m *MockManager) DeleteAccountCluster(ctx context.Context, accountID, userI
|
||||
}
|
||||
|
||||
// DeleteAccountCluster indicates an expected call of DeleteAccountCluster.
|
||||
func (mr *MockManagerMockRecorder) DeleteAccountCluster(ctx, accountID, userID, clusterAddress interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) DeleteAccountCluster(ctx, accountID, userID, clusterAddress any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAccountCluster", reflect.TypeOf((*MockManager)(nil).DeleteAccountCluster), ctx, accountID, userID, clusterAddress)
|
||||
}
|
||||
@@ -88,7 +94,7 @@ func (m *MockManager) DeleteAllServices(ctx context.Context, accountID, userID s
|
||||
}
|
||||
|
||||
// DeleteAllServices indicates an expected call of DeleteAllServices.
|
||||
func (mr *MockManagerMockRecorder) DeleteAllServices(ctx, accountID, userID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) DeleteAllServices(ctx, accountID, userID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAllServices", reflect.TypeOf((*MockManager)(nil).DeleteAllServices), ctx, accountID, userID)
|
||||
}
|
||||
@@ -102,7 +108,7 @@ func (m *MockManager) DeleteService(ctx context.Context, accountID, userID, serv
|
||||
}
|
||||
|
||||
// DeleteService indicates an expected call of DeleteService.
|
||||
func (mr *MockManagerMockRecorder) DeleteService(ctx, accountID, userID, serviceID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) DeleteService(ctx, accountID, userID, serviceID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteService", reflect.TypeOf((*MockManager)(nil).DeleteService), ctx, accountID, userID, serviceID)
|
||||
}
|
||||
@@ -117,7 +123,7 @@ func (m *MockManager) GetAccountServices(ctx context.Context, accountID string)
|
||||
}
|
||||
|
||||
// GetAccountServices indicates an expected call of GetAccountServices.
|
||||
func (mr *MockManagerMockRecorder) GetAccountServices(ctx, accountID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetAccountServices(ctx, accountID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAccountServices", reflect.TypeOf((*MockManager)(nil).GetAccountServices), ctx, accountID)
|
||||
}
|
||||
@@ -132,7 +138,7 @@ func (m *MockManager) GetAllServices(ctx context.Context, accountID, userID stri
|
||||
}
|
||||
|
||||
// GetAllServices indicates an expected call of GetAllServices.
|
||||
func (mr *MockManagerMockRecorder) GetAllServices(ctx, accountID, userID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetAllServices(ctx, accountID, userID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAllServices", reflect.TypeOf((*MockManager)(nil).GetAllServices), ctx, accountID, userID)
|
||||
}
|
||||
@@ -147,7 +153,7 @@ func (m *MockManager) GetClusters(ctx context.Context, accountID, userID string)
|
||||
}
|
||||
|
||||
// GetClusters indicates an expected call of GetClusters.
|
||||
func (mr *MockManagerMockRecorder) GetClusters(ctx, accountID, userID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetClusters(ctx, accountID, userID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetClusters", reflect.TypeOf((*MockManager)(nil).GetClusters), ctx, accountID, userID)
|
||||
}
|
||||
@@ -162,7 +168,7 @@ func (m *MockManager) GetGlobalServices(ctx context.Context) ([]*Service, error)
|
||||
}
|
||||
|
||||
// GetGlobalServices indicates an expected call of GetGlobalServices.
|
||||
func (mr *MockManagerMockRecorder) GetGlobalServices(ctx interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetGlobalServices(ctx any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetGlobalServices", reflect.TypeOf((*MockManager)(nil).GetGlobalServices), ctx)
|
||||
}
|
||||
@@ -177,7 +183,7 @@ func (m *MockManager) GetService(ctx context.Context, accountID, userID, service
|
||||
}
|
||||
|
||||
// GetService indicates an expected call of GetService.
|
||||
func (mr *MockManagerMockRecorder) GetService(ctx, accountID, userID, serviceID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetService(ctx, accountID, userID, serviceID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetService", reflect.TypeOf((*MockManager)(nil).GetService), ctx, accountID, userID, serviceID)
|
||||
}
|
||||
@@ -192,7 +198,7 @@ func (m *MockManager) GetServiceByDomain(ctx context.Context, domain string) (*S
|
||||
}
|
||||
|
||||
// GetServiceByDomain indicates an expected call of GetServiceByDomain.
|
||||
func (mr *MockManagerMockRecorder) GetServiceByDomain(ctx, domain interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetServiceByDomain(ctx, domain any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetServiceByDomain", reflect.TypeOf((*MockManager)(nil).GetServiceByDomain), ctx, domain)
|
||||
}
|
||||
@@ -207,7 +213,7 @@ func (m *MockManager) GetServiceByID(ctx context.Context, accountID, serviceID s
|
||||
}
|
||||
|
||||
// GetServiceByID indicates an expected call of GetServiceByID.
|
||||
func (mr *MockManagerMockRecorder) GetServiceByID(ctx, accountID, serviceID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetServiceByID(ctx, accountID, serviceID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetServiceByID", reflect.TypeOf((*MockManager)(nil).GetServiceByID), ctx, accountID, serviceID)
|
||||
}
|
||||
@@ -222,7 +228,7 @@ func (m *MockManager) GetServiceIDByTargetID(ctx context.Context, accountID, res
|
||||
}
|
||||
|
||||
// GetServiceIDByTargetID indicates an expected call of GetServiceIDByTargetID.
|
||||
func (mr *MockManagerMockRecorder) GetServiceIDByTargetID(ctx, accountID, resourceID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetServiceIDByTargetID(ctx, accountID, resourceID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetServiceIDByTargetID", reflect.TypeOf((*MockManager)(nil).GetServiceIDByTargetID), ctx, accountID, resourceID)
|
||||
}
|
||||
@@ -236,7 +242,7 @@ func (m *MockManager) ReloadAllServicesForAccount(ctx context.Context, accountID
|
||||
}
|
||||
|
||||
// ReloadAllServicesForAccount indicates an expected call of ReloadAllServicesForAccount.
|
||||
func (mr *MockManagerMockRecorder) ReloadAllServicesForAccount(ctx, accountID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) ReloadAllServicesForAccount(ctx, accountID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ReloadAllServicesForAccount", reflect.TypeOf((*MockManager)(nil).ReloadAllServicesForAccount), ctx, accountID)
|
||||
}
|
||||
@@ -250,7 +256,7 @@ func (m *MockManager) ReloadService(ctx context.Context, accountID, serviceID st
|
||||
}
|
||||
|
||||
// ReloadService indicates an expected call of ReloadService.
|
||||
func (mr *MockManagerMockRecorder) ReloadService(ctx, accountID, serviceID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) ReloadService(ctx, accountID, serviceID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ReloadService", reflect.TypeOf((*MockManager)(nil).ReloadService), ctx, accountID, serviceID)
|
||||
}
|
||||
@@ -264,7 +270,7 @@ func (m *MockManager) RenewServiceFromPeer(ctx context.Context, accountID, peerI
|
||||
}
|
||||
|
||||
// RenewServiceFromPeer indicates an expected call of RenewServiceFromPeer.
|
||||
func (mr *MockManagerMockRecorder) RenewServiceFromPeer(ctx, accountID, peerID, serviceID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) RenewServiceFromPeer(ctx, accountID, peerID, serviceID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RenewServiceFromPeer", reflect.TypeOf((*MockManager)(nil).RenewServiceFromPeer), ctx, accountID, peerID, serviceID)
|
||||
}
|
||||
@@ -278,7 +284,7 @@ func (m *MockManager) SetCertificateIssuedAt(ctx context.Context, accountID, ser
|
||||
}
|
||||
|
||||
// SetCertificateIssuedAt indicates an expected call of SetCertificateIssuedAt.
|
||||
func (mr *MockManagerMockRecorder) SetCertificateIssuedAt(ctx, accountID, serviceID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) SetCertificateIssuedAt(ctx, accountID, serviceID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetCertificateIssuedAt", reflect.TypeOf((*MockManager)(nil).SetCertificateIssuedAt), ctx, accountID, serviceID)
|
||||
}
|
||||
@@ -292,7 +298,7 @@ func (m *MockManager) SetStatus(ctx context.Context, accountID, serviceID string
|
||||
}
|
||||
|
||||
// SetStatus indicates an expected call of SetStatus.
|
||||
func (mr *MockManagerMockRecorder) SetStatus(ctx, accountID, serviceID, status interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) SetStatus(ctx, accountID, serviceID, status any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetStatus", reflect.TypeOf((*MockManager)(nil).SetStatus), ctx, accountID, serviceID, status)
|
||||
}
|
||||
@@ -304,7 +310,7 @@ func (m *MockManager) StartExposeReaper(ctx context.Context) {
|
||||
}
|
||||
|
||||
// StartExposeReaper indicates an expected call of StartExposeReaper.
|
||||
func (mr *MockManagerMockRecorder) StartExposeReaper(ctx interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) StartExposeReaper(ctx any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "StartExposeReaper", reflect.TypeOf((*MockManager)(nil).StartExposeReaper), ctx)
|
||||
}
|
||||
@@ -318,7 +324,7 @@ func (m *MockManager) StopServiceFromPeer(ctx context.Context, accountID, peerID
|
||||
}
|
||||
|
||||
// StopServiceFromPeer indicates an expected call of StopServiceFromPeer.
|
||||
func (mr *MockManagerMockRecorder) StopServiceFromPeer(ctx, accountID, peerID, serviceID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) StopServiceFromPeer(ctx, accountID, peerID, serviceID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "StopServiceFromPeer", reflect.TypeOf((*MockManager)(nil).StopServiceFromPeer), ctx, accountID, peerID, serviceID)
|
||||
}
|
||||
@@ -333,7 +339,7 @@ func (m *MockManager) UpdateService(ctx context.Context, accountID, userID strin
|
||||
}
|
||||
|
||||
// UpdateService indicates an expected call of UpdateService.
|
||||
func (mr *MockManagerMockRecorder) UpdateService(ctx, accountID, userID, service interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) UpdateService(ctx, accountID, userID, service any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateService", reflect.TypeOf((*MockManager)(nil).UpdateService), ctx, accountID, userID, service)
|
||||
}
|
||||
|
||||
@@ -204,7 +204,6 @@ func (h *handler) getClusters(w http.ResponseWriter, r *http.Request) {
|
||||
SupportsCustomPorts: c.SupportsCustomPorts,
|
||||
RequireSubdomain: c.RequireSubdomain,
|
||||
SupportsCrowdsec: c.SupportsCrowdSec,
|
||||
SupportsAppsec: c.SupportsAppSec,
|
||||
Private: c.Private,
|
||||
})
|
||||
}
|
||||
|
||||
@@ -6,7 +6,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -82,7 +82,6 @@ type CapabilityProvider interface {
|
||||
ClusterSupportsCustomPorts(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterSupportsAppSec(ctx context.Context, clusterAddr string) *bool
|
||||
ClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool
|
||||
}
|
||||
|
||||
@@ -138,7 +137,6 @@ func (m *Manager) GetClusters(ctx context.Context, accountID, userID string) ([]
|
||||
clusters[i].SupportsCustomPorts = m.capabilities.ClusterSupportsCustomPorts(ctx, clusters[i].Address)
|
||||
clusters[i].RequireSubdomain = m.capabilities.ClusterRequireSubdomain(ctx, clusters[i].Address)
|
||||
clusters[i].SupportsCrowdSec = m.capabilities.ClusterSupportsCrowdSec(ctx, clusters[i].Address)
|
||||
clusters[i].SupportsAppSec = m.capabilities.ClusterSupportsAppSec(ctx, clusters[i].Address)
|
||||
clusters[i].Private = m.capabilities.ClusterSupportsPrivate(ctx, clusters[i].Address)
|
||||
}
|
||||
|
||||
|
||||
@@ -8,7 +8,7 @@ import (
|
||||
"time"
|
||||
|
||||
cachestore "github.com/eko/gocache/lib/v4/store"
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"go.opentelemetry.io/otel/metric/noop"
|
||||
|
||||
@@ -165,22 +165,6 @@ type AccessRestrictions struct {
|
||||
AllowedCountries []string `json:"allowed_countries,omitempty" gorm:"serializer:json"`
|
||||
BlockedCountries []string `json:"blocked_countries,omitempty" gorm:"serializer:json"`
|
||||
CrowdSecMode string `json:"crowdsec_mode,omitempty" gorm:"serializer:json"`
|
||||
// AllowMatch controls how the allowlists combine: "" or "all" require
|
||||
// matching every allowlist (AND), "any" requires matching at least one (OR).
|
||||
// Empty is treated as "all" for backward compatibility with existing records.
|
||||
AllowMatch string `json:"allow_match,omitempty" gorm:"serializer:json"`
|
||||
// AppSecMode is the CrowdSec AppSec (WAF) request inspection mode: "",
|
||||
// "off", "enforce", or "observe". HTTP services only.
|
||||
AppSecMode string `json:"appsec_mode,omitempty" gorm:"serializer:json"`
|
||||
}
|
||||
|
||||
// isEmpty reports whether no restriction is configured. Both conversions drop
|
||||
// the object entirely in that case, so a field missing from this check is
|
||||
// silently discarded on the way to the API and the proxy.
|
||||
func (r AccessRestrictions) isEmpty() bool {
|
||||
return len(r.AllowedCIDRs) == 0 && len(r.BlockedCIDRs) == 0 &&
|
||||
len(r.AllowedCountries) == 0 && len(r.BlockedCountries) == 0 &&
|
||||
r.CrowdSecMode == "" && r.AllowMatch == "" && r.AppSecMode == ""
|
||||
}
|
||||
|
||||
// Copy returns a deep copy of the AccessRestrictions.
|
||||
@@ -191,8 +175,6 @@ func (r AccessRestrictions) Copy() AccessRestrictions {
|
||||
AllowedCountries: slices.Clone(r.AllowedCountries),
|
||||
BlockedCountries: slices.Clone(r.BlockedCountries),
|
||||
CrowdSecMode: r.CrowdSecMode,
|
||||
AllowMatch: r.AllowMatch,
|
||||
AppSecMode: r.AppSecMode,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -826,23 +808,13 @@ func restrictionsFromAPI(r *api.AccessRestrictions) (AccessRestrictions, error)
|
||||
}
|
||||
res.CrowdSecMode = string(*r.CrowdsecMode)
|
||||
}
|
||||
if r.AllowMatch != nil {
|
||||
if !r.AllowMatch.Valid() {
|
||||
return AccessRestrictions{}, fmt.Errorf("invalid allow_match %q", *r.AllowMatch)
|
||||
}
|
||||
res.AllowMatch = string(*r.AllowMatch)
|
||||
}
|
||||
if r.AppsecMode != nil {
|
||||
if !r.AppsecMode.Valid() {
|
||||
return AccessRestrictions{}, fmt.Errorf("invalid appsec_mode %q", *r.AppsecMode)
|
||||
}
|
||||
res.AppSecMode = string(*r.AppsecMode)
|
||||
}
|
||||
return res, nil
|
||||
}
|
||||
|
||||
func restrictionsToAPI(r AccessRestrictions) *api.AccessRestrictions {
|
||||
if r.isEmpty() {
|
||||
if len(r.AllowedCIDRs) == 0 && len(r.BlockedCIDRs) == 0 &&
|
||||
len(r.AllowedCountries) == 0 && len(r.BlockedCountries) == 0 &&
|
||||
r.CrowdSecMode == "" {
|
||||
return nil
|
||||
}
|
||||
res := &api.AccessRestrictions{}
|
||||
@@ -862,19 +834,13 @@ func restrictionsToAPI(r AccessRestrictions) *api.AccessRestrictions {
|
||||
mode := api.AccessRestrictionsCrowdsecMode(r.CrowdSecMode)
|
||||
res.CrowdsecMode = &mode
|
||||
}
|
||||
if r.AllowMatch != "" {
|
||||
match := api.AccessRestrictionsAllowMatch(r.AllowMatch)
|
||||
res.AllowMatch = &match
|
||||
}
|
||||
if r.AppSecMode != "" {
|
||||
mode := api.AccessRestrictionsAppsecMode(r.AppSecMode)
|
||||
res.AppsecMode = &mode
|
||||
}
|
||||
return res
|
||||
}
|
||||
|
||||
func restrictionsToProto(r AccessRestrictions) *proto.AccessRestrictions {
|
||||
if r.isEmpty() {
|
||||
if len(r.AllowedCIDRs) == 0 && len(r.BlockedCIDRs) == 0 &&
|
||||
len(r.AllowedCountries) == 0 && len(r.BlockedCountries) == 0 &&
|
||||
r.CrowdSecMode == "" {
|
||||
return nil
|
||||
}
|
||||
return &proto.AccessRestrictions{
|
||||
@@ -883,8 +849,6 @@ func restrictionsToProto(r AccessRestrictions) *proto.AccessRestrictions {
|
||||
AllowedCountries: r.AllowedCountries,
|
||||
BlockedCountries: r.BlockedCountries,
|
||||
CrowdsecMode: r.CrowdSecMode,
|
||||
AllowMatch: r.AllowMatch,
|
||||
AppsecMode: r.AppSecMode,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -910,11 +874,6 @@ func (s *Service) Validate() error {
|
||||
if err := validateAccessRestrictions(&s.Restrictions); err != nil {
|
||||
return err
|
||||
}
|
||||
// AppSec inspects HTTP requests, so it cannot apply to the L4 modes, which
|
||||
// forward opaque byte streams.
|
||||
if appSecEnabled(s.Restrictions.AppSecMode) && s.Mode != ModeHTTP {
|
||||
return fmt.Errorf("appsec_mode is only supported for HTTP services, got mode %q", s.Mode)
|
||||
}
|
||||
if err := s.validatePrivateRequirements(); err != nil {
|
||||
return err
|
||||
}
|
||||
@@ -1283,39 +1242,10 @@ func validateCrowdSecMode(mode string) error {
|
||||
}
|
||||
}
|
||||
|
||||
func validateAllowMatch(mode string) error {
|
||||
switch mode {
|
||||
case "", "all", "any":
|
||||
return nil
|
||||
default:
|
||||
return fmt.Errorf("allow_match %q is invalid", mode)
|
||||
}
|
||||
}
|
||||
|
||||
func validateAppSecMode(mode string) error {
|
||||
switch mode {
|
||||
case "", "off", "enforce", "observe":
|
||||
return nil
|
||||
default:
|
||||
return fmt.Errorf("appsec_mode %q is invalid", mode)
|
||||
}
|
||||
}
|
||||
|
||||
// appSecEnabled reports whether the mode asks for request inspection.
|
||||
func appSecEnabled(mode string) bool {
|
||||
return mode == "enforce" || mode == "observe"
|
||||
}
|
||||
|
||||
func validateAccessRestrictions(r *AccessRestrictions) error {
|
||||
if err := validateCrowdSecMode(r.CrowdSecMode); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validateAllowMatch(r.AllowMatch); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := validateAppSecMode(r.AppSecMode); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
if len(r.AllowedCIDRs) > maxCIDREntries {
|
||||
return fmt.Errorf("allowed_cidrs: exceeds maximum of %d entries", maxCIDREntries)
|
||||
|
||||
@@ -26,17 +26,6 @@ func validProxy() *Service {
|
||||
}
|
||||
}
|
||||
|
||||
// validL4Proxy returns a service that passes validation in one of the L4 modes.
|
||||
func validL4Proxy(mode string) *Service {
|
||||
rp := validProxy()
|
||||
rp.Mode = mode
|
||||
rp.ListenPort = 9000
|
||||
rp.Targets = []*Target{
|
||||
{TargetId: "peer-1", TargetType: TargetTypePeer, Host: "10.0.0.1", Port: 5432, Protocol: mode, Enabled: true},
|
||||
}
|
||||
return rp
|
||||
}
|
||||
|
||||
func TestValidate_Valid(t *testing.T) {
|
||||
require.NoError(t, validProxy().Validate())
|
||||
}
|
||||
@@ -1313,65 +1302,6 @@ func TestValidate_Private_AcceptsClusterTargetWithAccessGroups(t *testing.T) {
|
||||
require.NoError(t, rp.Validate())
|
||||
}
|
||||
|
||||
func TestRestrictions_AllowMatch_RoundTrip(t *testing.T) {
|
||||
anyMatch := api.AccessRestrictionsAllowMatchAny
|
||||
apiIn := &api.AccessRestrictions{
|
||||
AllowedCidrs: &[]string{"203.0.113.0/24"},
|
||||
AllowedCountries: &[]string{"US"},
|
||||
AllowMatch: &anyMatch,
|
||||
}
|
||||
|
||||
model, err := restrictionsFromAPI(apiIn)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "any", model.AllowMatch)
|
||||
|
||||
apiOut := restrictionsToAPI(model)
|
||||
require.NotNil(t, apiOut.AllowMatch)
|
||||
assert.Equal(t, api.AccessRestrictionsAllowMatchAny, *apiOut.AllowMatch)
|
||||
|
||||
protoOut := restrictionsToProto(model)
|
||||
require.NotNil(t, protoOut)
|
||||
assert.Equal(t, "any", protoOut.AllowMatch)
|
||||
}
|
||||
|
||||
func TestRestrictions_AllowMatch_EmptyDefaultsToAll(t *testing.T) {
|
||||
// A stored record with no allow_match (existing services) stays empty and
|
||||
// must not surface a value on the API, preserving AND behavior downstream.
|
||||
model, err := restrictionsFromAPI(&api.AccessRestrictions{
|
||||
AllowedCidrs: &[]string{"203.0.113.0/24"},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, model.AllowMatch, "unset allow_match stays empty")
|
||||
|
||||
apiOut := restrictionsToAPI(model)
|
||||
require.NotNil(t, apiOut)
|
||||
assert.Nil(t, apiOut.AllowMatch, "empty allow_match is omitted from the API response")
|
||||
}
|
||||
|
||||
func TestRestrictions_AllowMatchOnly_Preserved(t *testing.T) {
|
||||
// allow_match set without any list must not be dropped by the emptiness
|
||||
// guards, so it round-trips through both the API and proto conversions.
|
||||
model := AccessRestrictions{AllowMatch: "any"}
|
||||
|
||||
apiOut := restrictionsToAPI(model)
|
||||
require.NotNil(t, apiOut, "allow-match-only restriction must not be omitted from the API response")
|
||||
require.NotNil(t, apiOut.AllowMatch)
|
||||
assert.Equal(t, api.AccessRestrictionsAllowMatchAny, *apiOut.AllowMatch)
|
||||
|
||||
protoOut := restrictionsToProto(model)
|
||||
require.NotNil(t, protoOut, "allow-match-only restriction must not be omitted from the proto output")
|
||||
assert.Equal(t, "any", protoOut.AllowMatch)
|
||||
}
|
||||
|
||||
func TestValidate_RejectsInvalidAllowMatch(t *testing.T) {
|
||||
rp := validProxy()
|
||||
rp.Restrictions = AccessRestrictions{
|
||||
AllowedCIDRs: []string{"203.0.113.0/24"},
|
||||
AllowMatch: "sometimes",
|
||||
}
|
||||
assert.ErrorContains(t, rp.Validate(), "allow_match")
|
||||
}
|
||||
|
||||
func TestValidate_Private_RejectsNonHTTPMode(t *testing.T) {
|
||||
rp := validProxy()
|
||||
rp.Private = true
|
||||
@@ -1385,68 +1315,3 @@ func TestValidate_Private_RejectsNonHTTPMode(t *testing.T) {
|
||||
}}
|
||||
assert.ErrorContains(t, rp.Validate(), "HTTP")
|
||||
}
|
||||
|
||||
func TestRestrictions_AppSecMode_RoundTrip(t *testing.T) {
|
||||
mode := api.AccessRestrictionsAppsecModeEnforce
|
||||
apiIn := &api.AccessRestrictions{AppsecMode: &mode}
|
||||
|
||||
model, err := restrictionsFromAPI(apiIn)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "enforce", model.AppSecMode)
|
||||
|
||||
// appsec_mode alone must keep the restrictions object alive on both the API
|
||||
// and proto legs: it is meaningful without any CIDR or country entry.
|
||||
apiOut := restrictionsToAPI(model)
|
||||
require.NotNil(t, apiOut, "appsec_mode alone must not collapse the restrictions to nil")
|
||||
require.NotNil(t, apiOut.AppsecMode)
|
||||
assert.Equal(t, api.AccessRestrictionsAppsecModeEnforce, *apiOut.AppsecMode)
|
||||
|
||||
protoOut := restrictionsToProto(model)
|
||||
require.NotNil(t, protoOut, "appsec_mode alone must reach the proxy")
|
||||
assert.Equal(t, "enforce", protoOut.AppsecMode)
|
||||
}
|
||||
|
||||
func TestRestrictions_AppSecMode_EmptyIsOmitted(t *testing.T) {
|
||||
model, err := restrictionsFromAPI(&api.AccessRestrictions{
|
||||
AllowedCidrs: &[]string{"203.0.113.0/24"},
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, model.AppSecMode)
|
||||
|
||||
apiOut := restrictionsToAPI(model)
|
||||
require.NotNil(t, apiOut)
|
||||
assert.Nil(t, apiOut.AppsecMode, "empty appsec_mode is omitted from the API response")
|
||||
}
|
||||
|
||||
func TestRestrictions_AppSecMode_CopyIsDeep(t *testing.T) {
|
||||
original := AccessRestrictions{AppSecMode: "observe", CrowdSecMode: "enforce"}
|
||||
assert.Equal(t, original, original.Copy(), "Copy must carry every mode field")
|
||||
}
|
||||
|
||||
func TestValidate_RejectsInvalidAppSecMode(t *testing.T) {
|
||||
rp := validProxy()
|
||||
rp.Restrictions = AccessRestrictions{AppSecMode: "sometimes"}
|
||||
assert.ErrorContains(t, rp.Validate(), "appsec_mode")
|
||||
}
|
||||
|
||||
func TestValidate_RejectsAppSecOnL4Modes(t *testing.T) {
|
||||
// AppSec inspects HTTP requests, so the L4 modes cannot honor it. Accepting
|
||||
// the field there would report protection that never runs.
|
||||
for _, mode := range []string{ModeTCP, ModeUDP, ModeTLS} {
|
||||
t.Run(mode, func(t *testing.T) {
|
||||
rp := validL4Proxy(mode)
|
||||
rp.Restrictions = AccessRestrictions{AppSecMode: "enforce"}
|
||||
assert.ErrorContains(t, rp.Validate(), "appsec_mode is only supported for HTTP services")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestValidate_AllowsAppSecOffOnL4Modes(t *testing.T) {
|
||||
for _, mode := range []string{ModeTCP, ModeUDP, ModeTLS} {
|
||||
t.Run(mode, func(t *testing.T) {
|
||||
rp := validL4Proxy(mode)
|
||||
rp.Restrictions = AccessRestrictions{AppSecMode: "off"}
|
||||
require.NoError(t, rp.Validate(), "an explicit off must not be rejected on L4 services")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,7 +5,7 @@ import (
|
||||
"fmt"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ import (
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
nbconfig "github.com/netbirdio/netbird/management/internals/server/config"
|
||||
|
||||
@@ -562,7 +562,6 @@ func (s *ProxyServiceServer) registerProxyConnection(ctx context.Context, params
|
||||
SupportsCustomPorts: c.SupportsCustomPorts,
|
||||
RequireSubdomain: c.RequireSubdomain,
|
||||
SupportsCrowdsec: c.SupportsCrowdsec,
|
||||
SupportsAppsec: c.SupportsAppsec,
|
||||
Private: c.Private,
|
||||
}
|
||||
}
|
||||
|
||||
@@ -5,7 +5,7 @@ import (
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/grpc/codes"
|
||||
|
||||
@@ -6,7 +6,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/grpc"
|
||||
|
||||
@@ -7,7 +7,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/grpc"
|
||||
|
||||
@@ -10,7 +10,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
package account
|
||||
|
||||
//go:generate go run github.com/golang/mock/mockgen -package account -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
|
||||
//go:generate go tool mockgen -package account -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -14,7 +14,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/prometheus/client_golang/prometheus/push"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
@@ -6,7 +6,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
nbdns "github.com/netbirdio/netbird/dns"
|
||||
|
||||
@@ -11,7 +11,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/google/uuid"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
@@ -12,7 +12,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/gorilla/mux"
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
|
||||
@@ -10,7 +10,7 @@ import (
|
||||
"net/mail"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/gorilla/mux"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
@@ -13,9 +13,8 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"github.com/gorilla/mux"
|
||||
ugomock "go.uber.org/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"golang.org/x/exp/maps"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/controllers/network_map"
|
||||
@@ -106,7 +105,7 @@ func initTestMetaData(t *testing.T, peers ...*nbpeer.Peer) *Handler {
|
||||
},
|
||||
}
|
||||
|
||||
ctrl := ugomock.NewController(t)
|
||||
ctrl := gomock.NewController(t)
|
||||
|
||||
networkMapController := network_map.NewMockController(ctrl)
|
||||
networkMapController.EXPECT().
|
||||
|
||||
@@ -10,7 +10,7 @@ import (
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/gorilla/mux"
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
|
||||
@@ -10,7 +10,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@ import (
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -12,7 +12,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
|
||||
|
||||
@@ -10,7 +10,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
pb "github.com/golang/protobuf/proto" //nolint
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
@@ -6,7 +6,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
|
||||
@@ -4,7 +4,7 @@ import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
reverseproxy "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
|
||||
@@ -16,7 +16,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/rs/xid"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/stretchr/testify/assert"
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
package permissions
|
||||
|
||||
//go:generate go run github.com/golang/mock/mockgen -package permissions -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
|
||||
//go:generate go tool mockgen -package permissions -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
@@ -1,5 +1,10 @@
|
||||
// Code generated by MockGen. DO NOT EDIT.
|
||||
// Source: ./manager.go
|
||||
//
|
||||
// Generated by this command:
|
||||
//
|
||||
// mockgen -package permissions -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
|
||||
//
|
||||
|
||||
// Package permissions is a generated GoMock package.
|
||||
package permissions
|
||||
@@ -8,18 +13,19 @@ import (
|
||||
context "context"
|
||||
reflect "reflect"
|
||||
|
||||
gomock "github.com/golang/mock/gomock"
|
||||
account "github.com/netbirdio/netbird/management/server/account"
|
||||
modules "github.com/netbirdio/netbird/management/server/permissions/modules"
|
||||
operations "github.com/netbirdio/netbird/management/server/permissions/operations"
|
||||
roles "github.com/netbirdio/netbird/management/server/permissions/roles"
|
||||
types "github.com/netbirdio/netbird/management/server/types"
|
||||
gomock "go.uber.org/mock/gomock"
|
||||
)
|
||||
|
||||
// MockManager is a mock of Manager interface.
|
||||
type MockManager struct {
|
||||
ctrl *gomock.Controller
|
||||
recorder *MockManagerMockRecorder
|
||||
isgomock struct{}
|
||||
}
|
||||
|
||||
// MockManagerMockRecorder is the mock recorder for MockManager.
|
||||
@@ -49,7 +55,7 @@ func (m *MockManager) GetPermissionsByRole(ctx context.Context, role types.UserR
|
||||
}
|
||||
|
||||
// GetPermissionsByRole indicates an expected call of GetPermissionsByRole.
|
||||
func (mr *MockManagerMockRecorder) GetPermissionsByRole(ctx, role interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetPermissionsByRole(ctx, role any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetPermissionsByRole", reflect.TypeOf((*MockManager)(nil).GetPermissionsByRole), ctx, role)
|
||||
}
|
||||
@@ -61,7 +67,7 @@ func (m *MockManager) SetAccountManager(accountManager account.Manager) {
|
||||
}
|
||||
|
||||
// SetAccountManager indicates an expected call of SetAccountManager.
|
||||
func (mr *MockManagerMockRecorder) SetAccountManager(accountManager interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) SetAccountManager(accountManager any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "SetAccountManager", reflect.TypeOf((*MockManager)(nil).SetAccountManager), accountManager)
|
||||
}
|
||||
@@ -76,7 +82,7 @@ func (m *MockManager) ValidateAccountAccess(ctx context.Context, accountID strin
|
||||
}
|
||||
|
||||
// ValidateAccountAccess indicates an expected call of ValidateAccountAccess.
|
||||
func (mr *MockManagerMockRecorder) ValidateAccountAccess(ctx, accountID, user, allowOwnerAndAdmin interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) ValidateAccountAccess(ctx, accountID, user, allowOwnerAndAdmin any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ValidateAccountAccess", reflect.TypeOf((*MockManager)(nil).ValidateAccountAccess), ctx, accountID, user, allowOwnerAndAdmin)
|
||||
}
|
||||
@@ -90,7 +96,7 @@ func (m *MockManager) ValidateRoleModuleAccess(ctx context.Context, accountID st
|
||||
}
|
||||
|
||||
// ValidateRoleModuleAccess indicates an expected call of ValidateRoleModuleAccess.
|
||||
func (mr *MockManagerMockRecorder) ValidateRoleModuleAccess(ctx, accountID, role, module, operation interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) ValidateRoleModuleAccess(ctx, accountID, role, module, operation any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ValidateRoleModuleAccess", reflect.TypeOf((*MockManager)(nil).ValidateRoleModuleAccess), ctx, accountID, role, module, operation)
|
||||
}
|
||||
@@ -106,7 +112,7 @@ func (m *MockManager) ValidateUserPermissions(ctx context.Context, accountID, us
|
||||
}
|
||||
|
||||
// ValidateUserPermissions indicates an expected call of ValidateUserPermissions.
|
||||
func (mr *MockManagerMockRecorder) ValidateUserPermissions(ctx, accountID, userID, module, operation interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) ValidateUserPermissions(ctx, accountID, userID, module, operation any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "ValidateUserPermissions", reflect.TypeOf((*MockManager)(nil).ValidateUserPermissions), ctx, accountID, userID, module, operation)
|
||||
}
|
||||
|
||||
@@ -6,7 +6,7 @@ import (
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"go.uber.org/mock/gomock"
|
||||
"github.com/rs/xid"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
package settings
|
||||
|
||||
//go:generate go run github.com/golang/mock/mockgen -package settings -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
|
||||
//go:generate go tool mockgen -package settings -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
@@ -1,5 +1,10 @@
|
||||
// Code generated by MockGen. DO NOT EDIT.
|
||||
// Source: ./manager.go
|
||||
//
|
||||
// Generated by this command:
|
||||
//
|
||||
// mockgen -package settings -destination=manager_mock.go -source=./manager.go -build_flags=-mod=mod
|
||||
//
|
||||
|
||||
// Package settings is a generated GoMock package.
|
||||
package settings
|
||||
@@ -9,15 +14,16 @@ import (
|
||||
netip "net/netip"
|
||||
reflect "reflect"
|
||||
|
||||
gomock "github.com/golang/mock/gomock"
|
||||
extra_settings "github.com/netbirdio/netbird/management/server/integrations/extra_settings"
|
||||
types "github.com/netbirdio/netbird/management/server/types"
|
||||
gomock "go.uber.org/mock/gomock"
|
||||
)
|
||||
|
||||
// MockManager is a mock of Manager interface.
|
||||
type MockManager struct {
|
||||
ctrl *gomock.Controller
|
||||
recorder *MockManagerMockRecorder
|
||||
isgomock struct{}
|
||||
}
|
||||
|
||||
// MockManagerMockRecorder is the mock recorder for MockManager.
|
||||
@@ -37,6 +43,22 @@ func (m *MockManager) EXPECT() *MockManagerMockRecorder {
|
||||
return m.recorder
|
||||
}
|
||||
|
||||
// GetEffectiveNetworkRanges mocks base method.
|
||||
func (m *MockManager) GetEffectiveNetworkRanges(ctx context.Context, accountID string) (netip.Prefix, netip.Prefix, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetEffectiveNetworkRanges", ctx, accountID)
|
||||
ret0, _ := ret[0].(netip.Prefix)
|
||||
ret1, _ := ret[1].(netip.Prefix)
|
||||
ret2, _ := ret[2].(error)
|
||||
return ret0, ret1, ret2
|
||||
}
|
||||
|
||||
// GetEffectiveNetworkRanges indicates an expected call of GetEffectiveNetworkRanges.
|
||||
func (mr *MockManagerMockRecorder) GetEffectiveNetworkRanges(ctx, accountID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetEffectiveNetworkRanges", reflect.TypeOf((*MockManager)(nil).GetEffectiveNetworkRanges), ctx, accountID)
|
||||
}
|
||||
|
||||
// GetExtraSettings mocks base method.
|
||||
func (m *MockManager) GetExtraSettings(ctx context.Context, accountID string) (*types.ExtraSettings, error) {
|
||||
m.ctrl.T.Helper()
|
||||
@@ -47,7 +69,7 @@ func (m *MockManager) GetExtraSettings(ctx context.Context, accountID string) (*
|
||||
}
|
||||
|
||||
// GetExtraSettings indicates an expected call of GetExtraSettings.
|
||||
func (mr *MockManagerMockRecorder) GetExtraSettings(ctx, accountID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetExtraSettings(ctx, accountID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetExtraSettings", reflect.TypeOf((*MockManager)(nil).GetExtraSettings), ctx, accountID)
|
||||
}
|
||||
@@ -76,7 +98,7 @@ func (m *MockManager) GetSettings(ctx context.Context, accountID, userID string)
|
||||
}
|
||||
|
||||
// GetSettings indicates an expected call of GetSettings.
|
||||
func (mr *MockManagerMockRecorder) GetSettings(ctx, accountID, userID interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) GetSettings(ctx, accountID, userID any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetSettings", reflect.TypeOf((*MockManager)(nil).GetSettings), ctx, accountID, userID)
|
||||
}
|
||||
@@ -91,23 +113,7 @@ func (m *MockManager) UpdateExtraSettings(ctx context.Context, accountID, userID
|
||||
}
|
||||
|
||||
// UpdateExtraSettings indicates an expected call of UpdateExtraSettings.
|
||||
func (mr *MockManagerMockRecorder) UpdateExtraSettings(ctx, accountID, userID, extraSettings interface{}) *gomock.Call {
|
||||
func (mr *MockManagerMockRecorder) UpdateExtraSettings(ctx, accountID, userID, extraSettings any) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateExtraSettings", reflect.TypeOf((*MockManager)(nil).UpdateExtraSettings), ctx, accountID, userID, extraSettings)
|
||||
}
|
||||
|
||||
// GetEffectiveNetworkRanges mocks base method.
|
||||
func (m *MockManager) GetEffectiveNetworkRanges(ctx context.Context, accountID string) (netip.Prefix, netip.Prefix, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetEffectiveNetworkRanges", ctx, accountID)
|
||||
ret0, _ := ret[0].(netip.Prefix)
|
||||
ret1, _ := ret[1].(netip.Prefix)
|
||||
ret2, _ := ret[2].(error)
|
||||
return ret0, ret1, ret2
|
||||
}
|
||||
|
||||
// GetEffectiveNetworkRanges indicates an expected call of GetEffectiveNetworkRanges.
|
||||
func (mr *MockManagerMockRecorder) GetEffectiveNetworkRanges(ctx, accountID interface{}) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetEffectiveNetworkRanges", reflect.TypeOf((*MockManager)(nil).GetEffectiveNetworkRanges), ctx, accountID)
|
||||
}
|
||||
|
||||
@@ -6491,7 +6491,6 @@ var validCapabilityColumns = map[string]struct{}{
|
||||
"supports_custom_ports": {},
|
||||
"require_subdomain": {},
|
||||
"supports_crowdsec": {},
|
||||
"supports_appsec": {},
|
||||
"private": {},
|
||||
}
|
||||
|
||||
@@ -6522,14 +6521,6 @@ func (s *SqlStore) GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr s
|
||||
return s.getClusterUnanimousCapability(ctx, clusterAddr, "supports_crowdsec")
|
||||
}
|
||||
|
||||
// GetClusterSupportsAppSec returns whether all active proxies in the cluster
|
||||
// have a CrowdSec AppSec endpoint configured. Returns nil when no proxy
|
||||
// reported the capability. Unanimous for the same reason as CrowdSec: a single
|
||||
// proxy without AppSec would let requests through uninspected.
|
||||
func (s *SqlStore) GetClusterSupportsAppSec(ctx context.Context, clusterAddr string) *bool {
|
||||
return s.getClusterUnanimousCapability(ctx, clusterAddr, "supports_appsec")
|
||||
}
|
||||
|
||||
// getClusterUnanimousCapability returns an aggregated boolean capability
|
||||
// requiring all active proxies in the cluster to report true.
|
||||
func (s *SqlStore) getClusterUnanimousCapability(ctx context.Context, clusterAddr, column string) *bool {
|
||||
|
||||
@@ -1,119 +0,0 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"runtime"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
|
||||
)
|
||||
|
||||
// Capabilities travel proxy → gRPC → embedded gorm columns → aggregation → API.
|
||||
// A field dropped at any of those hops reads as "capability absent", which is
|
||||
// indistinguishable from a proxy that never reported it: the dashboard simply
|
||||
// hides the feature and nothing fails. These assertions cover the persistence
|
||||
// and aggregation hops.
|
||||
func TestSqlStore_ClusterCapabilityAggregation(t *testing.T) {
|
||||
if os.Getenv("CI") == "true" && (runtime.GOOS == "darwin" || runtime.GOOS == "windows") {
|
||||
t.Skip("skip CI tests on darwin and windows")
|
||||
}
|
||||
|
||||
yes, no := true, false
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
reported []*bool // one entry per connected proxy in the cluster
|
||||
wantAppSec *bool
|
||||
wantAssertion string
|
||||
}{
|
||||
{
|
||||
name: "unreported stays unknown",
|
||||
reported: []*bool{nil},
|
||||
wantAppSec: nil,
|
||||
wantAssertion: "an unreported capability must not read as false",
|
||||
},
|
||||
{
|
||||
name: "single proxy reporting true",
|
||||
reported: []*bool{&yes},
|
||||
wantAppSec: &yes,
|
||||
wantAssertion: "a reported capability must survive persistence",
|
||||
},
|
||||
{
|
||||
name: "one proxy without it disables the cluster",
|
||||
reported: []*bool{&yes, &no},
|
||||
wantAppSec: &no,
|
||||
wantAssertion: "capability must be unanimous, so a rolling upgrade cannot leave traffic uninspected",
|
||||
},
|
||||
{
|
||||
name: "one proxy yet to report disables the cluster",
|
||||
reported: []*bool{&yes, nil},
|
||||
wantAppSec: &no,
|
||||
wantAssertion: "a proxy that has not reported must not count as capable",
|
||||
},
|
||||
}
|
||||
|
||||
runTestForAllEngines(t, "", func(t *testing.T, store Store) {
|
||||
ctx := context.Background()
|
||||
for i, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
cluster := fmt.Sprintf("cluster-%d.proxy.example", i)
|
||||
for j, reported := range tt.reported {
|
||||
require.NoError(t, store.SaveProxy(ctx, &proxy.Proxy{
|
||||
ID: fmt.Sprintf("proxy-%d-%d", i, j),
|
||||
ClusterAddress: cluster,
|
||||
Status: proxy.StatusConnected,
|
||||
LastSeen: time.Now(),
|
||||
Capabilities: proxy.Capabilities{SupportsAppsec: reported},
|
||||
}))
|
||||
}
|
||||
|
||||
got := store.GetClusterSupportsAppSec(ctx, cluster)
|
||||
if tt.wantAppSec == nil {
|
||||
assert.Nil(t, got, tt.wantAssertion)
|
||||
return
|
||||
}
|
||||
require.NotNil(t, got, tt.wantAssertion)
|
||||
assert.Equal(t, *tt.wantAppSec, *got, tt.wantAssertion)
|
||||
})
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// AppSec and IP reputation are separate endpoints, so a cluster can have either
|
||||
// without the other. Gating one on the other would silently disable a feature
|
||||
// the operator configured.
|
||||
func TestSqlStore_ClusterCapabilitiesAreIndependent(t *testing.T) {
|
||||
if os.Getenv("CI") == "true" && (runtime.GOOS == "darwin" || runtime.GOOS == "windows") {
|
||||
t.Skip("skip CI tests on darwin and windows")
|
||||
}
|
||||
|
||||
runTestForAllEngines(t, "", func(t *testing.T, store Store) {
|
||||
ctx := context.Background()
|
||||
const cluster = "independent.proxy.example"
|
||||
yes, no := true, false
|
||||
|
||||
require.NoError(t, store.SaveProxy(ctx, &proxy.Proxy{
|
||||
ID: "proxy-independent",
|
||||
ClusterAddress: cluster,
|
||||
Status: proxy.StatusConnected,
|
||||
LastSeen: time.Now(),
|
||||
Capabilities: proxy.Capabilities{
|
||||
SupportsAppsec: &yes,
|
||||
SupportsCrowdsec: &no,
|
||||
},
|
||||
}))
|
||||
|
||||
appsec := store.GetClusterSupportsAppSec(ctx, cluster)
|
||||
crowdsec := store.GetClusterSupportsCrowdSec(ctx, cluster)
|
||||
require.NotNil(t, appsec)
|
||||
require.NotNil(t, crowdsec)
|
||||
assert.True(t, *appsec, "AppSec must not be gated on CrowdSec")
|
||||
assert.False(t, *crowdsec, "CrowdSec must not be implied by AppSec")
|
||||
})
|
||||
}
|
||||
@@ -133,46 +133,3 @@ func TestSqlStore_GetAccount_ServiceTargetOptionsRoundtrip(t *testing.T) {
|
||||
assert.True(t, tg.Options.DisableAccessLog, "options disable access log")
|
||||
})
|
||||
}
|
||||
|
||||
func TestSqlStore_GetAccount_ServiceRestrictionsRoundtrip(t *testing.T) {
|
||||
if os.Getenv("CI") == "true" && (runtime.GOOS == "darwin" || runtime.GOOS == "windows") {
|
||||
t.Skip("skip CI tests on darwin and windows")
|
||||
}
|
||||
|
||||
runTestForAllEngines(t, "", func(t *testing.T, store Store) {
|
||||
ctx := context.Background()
|
||||
account := newAccountWithId(ctx, "account_svc_restrictions", "testuser", "")
|
||||
require.NoError(t, store.SaveAccount(ctx, account))
|
||||
|
||||
svc := &rpservice.Service{
|
||||
ID: "svc-restrictions",
|
||||
AccountID: account.Id,
|
||||
Name: "restricted-svc",
|
||||
Domain: "restricted.example",
|
||||
Enabled: true,
|
||||
Mode: rpservice.ModeHTTP,
|
||||
Restrictions: rpservice.AccessRestrictions{
|
||||
AllowedCIDRs: []string{"203.0.113.0/24"},
|
||||
AllowedCountries: []string{"US"},
|
||||
AllowMatch: "any",
|
||||
CrowdSecMode: "observe",
|
||||
AppSecMode: "enforce",
|
||||
},
|
||||
}
|
||||
require.NoError(t, store.CreateService(ctx, svc))
|
||||
|
||||
loaded, err := store.GetAccount(ctx, account.Id)
|
||||
require.NoError(t, err)
|
||||
require.Len(t, loaded.Services, 1)
|
||||
|
||||
// Restrictions are stored as a JSON blob; confirm the whole struct,
|
||||
// including allow_match and the CrowdSec modes, survives the read path
|
||||
// (Postgres pgx path included via runTestForAllEngines).
|
||||
got := loaded.Services[0].Restrictions
|
||||
assert.Equal(t, []string{"203.0.113.0/24"}, got.AllowedCIDRs)
|
||||
assert.Equal(t, []string{"US"}, got.AllowedCountries)
|
||||
assert.Equal(t, "any", got.AllowMatch)
|
||||
assert.Equal(t, "observe", got.CrowdSecMode)
|
||||
assert.Equal(t, "enforce", got.AppSecMode)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
package store
|
||||
|
||||
//go:generate go run github.com/golang/mock/mockgen -package store -destination=store_mock.go -source=./store.go -build_flags=-mod=mod
|
||||
//go:generate go tool mockgen -package store -destination=store_mock.go -source=./store.go -build_flags=-mod=mod
|
||||
|
||||
import (
|
||||
"context"
|
||||
@@ -329,7 +329,6 @@ type Store interface {
|
||||
GetClusterSupportsCustomPorts(ctx context.Context, clusterAddr string) *bool
|
||||
GetClusterRequireSubdomain(ctx context.Context, clusterAddr string) *bool
|
||||
GetClusterSupportsCrowdSec(ctx context.Context, clusterAddr string) *bool
|
||||
GetClusterSupportsAppSec(ctx context.Context, clusterAddr string) *bool
|
||||
GetClusterSupportsPrivate(ctx context.Context, clusterAddr string) *bool
|
||||
CleanupStaleProxies(ctx context.Context, inactivityDuration time.Duration) error
|
||||
GetAllProxies(ctx context.Context) ([]*proxy.Proxy, error)
|
||||
|
||||
File diff suppressed because it is too large
Load Diff
@@ -79,11 +79,6 @@ var (
|
||||
geoDataDir string
|
||||
crowdsecAPIURL string
|
||||
crowdsecAPIKey string
|
||||
appsecURL string
|
||||
appsecTimeout time.Duration
|
||||
appsecMaxBodyBytes int64
|
||||
captureBudgetBytes int64
|
||||
appsecMaxConcurrent int
|
||||
)
|
||||
|
||||
var rootCmd = &cobra.Command{
|
||||
@@ -130,11 +125,6 @@ func init() {
|
||||
rootCmd.Flags().StringVar(&geoDataDir, "geo-data-dir", envStringOrDefault("NB_PROXY_GEO_DATA_DIR", "/var/lib/netbird/geolocation"), "Directory for the GeoLite2 MMDB file (auto-downloaded if missing)")
|
||||
rootCmd.Flags().StringVar(&crowdsecAPIURL, "crowdsec-api-url", envStringOrDefault("NB_PROXY_CROWDSEC_API_URL", ""), "CrowdSec LAPI URL for IP reputation checks")
|
||||
rootCmd.Flags().StringVar(&crowdsecAPIKey, "crowdsec-api-key", envStringOrDefault("NB_PROXY_CROWDSEC_API_KEY", ""), "CrowdSec bouncer API key")
|
||||
rootCmd.Flags().StringVar(&appsecURL, "crowdsec-appsec-url", envStringOrDefault("NB_PROXY_CROWDSEC_APPSEC_URL", ""), "CrowdSec AppSec (WAF) endpoint for HTTP request inspection, e.g. http://127.0.0.1:7422/ (reuses the bouncer API key)")
|
||||
rootCmd.Flags().DurationVar(&appsecTimeout, "crowdsec-appsec-timeout", envDurationOrDefault("NB_PROXY_CROWDSEC_APPSEC_TIMEOUT", 0), "Timeout for a single AppSec inspection call (0 = 200ms)")
|
||||
rootCmd.Flags().IntVar(&appsecMaxConcurrent, "crowdsec-appsec-max-concurrent", int(envInt64OrDefault("NB_PROXY_CROWDSEC_APPSEC_MAX_CONCURRENT", 0)), "Cap on AppSec inspections in flight; further requests are denied in enforce mode rather than queued (0 = 256, negative = no cap)")
|
||||
rootCmd.Flags().Int64Var(&captureBudgetBytes, "capture-budget-bytes", envInt64OrDefault("NB_PROXY_CAPTURE_BUDGET_BYTES", 0), "Total in-flight request-body buffering across the proxy, shared by AppSec inspection and agent-network capture (0 = 256MiB)")
|
||||
rootCmd.Flags().Int64Var(&appsecMaxBodyBytes, "crowdsec-appsec-max-body-bytes", envInt64OrDefault("NB_PROXY_CROWDSEC_APPSEC_MAX_BODY_BYTES", 0), "Cap on the request body mirrored to AppSec (0 = 64KiB, negative = headers and URI only)")
|
||||
}
|
||||
|
||||
// Execute runs the root command.
|
||||
@@ -228,59 +218,47 @@ func runServer(cmd *cobra.Command, args []string) error {
|
||||
ctx, stop := signal.NotifyContext(context.Background(), syscall.SIGTERM, syscall.SIGINT)
|
||||
defer stop()
|
||||
|
||||
srv := proxy.New(ctx, serverConfig(logger, proxyToken, parsedTrustedProxies, perf))
|
||||
srv := proxy.New(ctx, proxy.Config{
|
||||
ListenAddr: addr,
|
||||
Logger: logger,
|
||||
Version: Version,
|
||||
ManagementAddress: mgmtAddr,
|
||||
ProxyURL: proxyDomain,
|
||||
ProxyToken: proxyToken,
|
||||
CertificateDirectory: certDir,
|
||||
CertificateFile: certFile,
|
||||
CertificateKeyFile: certKeyFile,
|
||||
GenerateACMECertificates: acmeCerts,
|
||||
ACMEChallengeAddress: acmeAddr,
|
||||
ACMEDirectory: acmeDir,
|
||||
ACMEEABKID: acmeEABKID,
|
||||
ACMEEABHMACKey: acmeEABHMACKey,
|
||||
ACMEChallengeType: acmeChallengeType,
|
||||
DebugEndpointEnabled: debugEndpoint,
|
||||
DebugEndpointAddress: debugEndpointAddr,
|
||||
HealthAddr: healthAddr,
|
||||
ForwardedProto: forwardedProto,
|
||||
TrustedProxies: parsedTrustedProxies,
|
||||
CertLockMethod: nbacme.CertLockMethod(certLockMethod),
|
||||
WildcardCertDir: wildcardCertDir,
|
||||
WireguardPort: wgPort,
|
||||
Performance: perf,
|
||||
ProxyProtocol: proxyProtocol,
|
||||
PreSharedKey: preSharedKey,
|
||||
SupportsCustomPorts: supportsCustomPorts,
|
||||
RequireSubdomain: requireSubdomain,
|
||||
Private: private,
|
||||
MaxDialTimeout: maxDialTimeout,
|
||||
MaxSessionIdleTimeout: maxSessionIdleTimeout,
|
||||
MappingBatchWatchdog: envDurationOrDefault("NB_PROXY_MAPPING_BATCH_WATCHDOG", 0),
|
||||
GeoDataDir: geoDataDir,
|
||||
CrowdSecAPIURL: crowdsecAPIURL,
|
||||
CrowdSecAPIKey: crowdsecAPIKey,
|
||||
})
|
||||
|
||||
return srv.ListenAndServe(ctx, addr)
|
||||
}
|
||||
|
||||
// serverConfig maps the parsed flags and environment onto the proxy config.
|
||||
// Kept separate from runServer so registering a new flag does not grow the
|
||||
// startup path.
|
||||
func serverConfig(logger *log.Logger, proxyToken string, trustedProxyList *trustedproxy.List, perf embed.Performance) proxy.Config {
|
||||
return proxy.Config{
|
||||
ListenAddr: addr,
|
||||
Logger: logger,
|
||||
Version: Version,
|
||||
ManagementAddress: mgmtAddr,
|
||||
ProxyURL: proxyDomain,
|
||||
ProxyToken: proxyToken,
|
||||
CertificateDirectory: certDir,
|
||||
CertificateFile: certFile,
|
||||
CertificateKeyFile: certKeyFile,
|
||||
GenerateACMECertificates: acmeCerts,
|
||||
ACMEChallengeAddress: acmeAddr,
|
||||
ACMEDirectory: acmeDir,
|
||||
ACMEEABKID: acmeEABKID,
|
||||
ACMEEABHMACKey: acmeEABHMACKey,
|
||||
ACMEChallengeType: acmeChallengeType,
|
||||
DebugEndpointEnabled: debugEndpoint,
|
||||
DebugEndpointAddress: debugEndpointAddr,
|
||||
HealthAddr: healthAddr,
|
||||
ForwardedProto: forwardedProto,
|
||||
TrustedProxies: trustedProxyList,
|
||||
CertLockMethod: nbacme.CertLockMethod(certLockMethod),
|
||||
WildcardCertDir: wildcardCertDir,
|
||||
WireguardPort: wgPort,
|
||||
Performance: perf,
|
||||
ProxyProtocol: proxyProtocol,
|
||||
PreSharedKey: preSharedKey,
|
||||
SupportsCustomPorts: supportsCustomPorts,
|
||||
RequireSubdomain: requireSubdomain,
|
||||
Private: private,
|
||||
MaxDialTimeout: maxDialTimeout,
|
||||
MaxSessionIdleTimeout: maxSessionIdleTimeout,
|
||||
MappingBatchWatchdog: envDurationOrDefault("NB_PROXY_MAPPING_BATCH_WATCHDOG", 0),
|
||||
GeoDataDir: geoDataDir,
|
||||
CrowdSecAPIURL: crowdsecAPIURL,
|
||||
CrowdSecAPIKey: crowdsecAPIKey,
|
||||
CrowdSecAppSecURL: appsecURL,
|
||||
CrowdSecAppSecTimeout: appsecTimeout,
|
||||
CrowdSecAppSecMaxBodyBytes: appsecMaxBodyBytes,
|
||||
CrowdSecAppSecMaxConcurrent: appsecMaxConcurrent,
|
||||
MiddlewareCaptureBudgetBytes: captureBudgetBytes,
|
||||
}
|
||||
}
|
||||
|
||||
func envBoolOrDefault(key string, def bool) bool {
|
||||
v, exists := os.LookupEnv(key)
|
||||
if !exists {
|
||||
@@ -315,19 +293,6 @@ func envUint16OrDefault(key string, def uint16) uint16 {
|
||||
return uint16(parsed)
|
||||
}
|
||||
|
||||
func envInt64OrDefault(key string, def int64) int64 {
|
||||
v, exists := os.LookupEnv(key)
|
||||
if !exists {
|
||||
return def
|
||||
}
|
||||
parsed, err := strconv.ParseInt(v, 10, 64)
|
||||
if err != nil {
|
||||
log.Warnf("parse %s=%q: %v, using default %d", key, v, err, def)
|
||||
return def
|
||||
}
|
||||
return parsed
|
||||
}
|
||||
|
||||
func envDurationOrDefault(key string, def time.Duration) time.Duration {
|
||||
v, exists := os.LookupEnv(key)
|
||||
if !exists {
|
||||
|
||||
@@ -1,163 +0,0 @@
|
||||
package appsec
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"errors"
|
||||
"io"
|
||||
"mime"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"slices"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// bufferBody reads up to limit+1 bytes from r.Body and always restores r.Body so
|
||||
// the request stays forwardable. oversize reports that the body exceeded limit, in
|
||||
// which case the returned prefix must not be used for inspection: the bytes are
|
||||
// only read so they can be replayed to the backend.
|
||||
func bufferBody(r *http.Request, limit int64) (body []byte, oversize bool, err error) {
|
||||
original := r.Body
|
||||
buf, readErr := io.ReadAll(io.LimitReader(original, limit+1))
|
||||
if readErr != nil && !errors.Is(readErr, io.EOF) {
|
||||
// Restore what was read so a downstream retry sees a consistent stream,
|
||||
// then surface the failure.
|
||||
r.Body = replay(buf, original)
|
||||
return nil, false, readErr
|
||||
}
|
||||
|
||||
if int64(len(buf)) > limit {
|
||||
r.Body = replay(buf, original)
|
||||
return nil, true, nil
|
||||
}
|
||||
|
||||
// The whole body is buffered, so the original is drained and can be closed.
|
||||
// A close error on a drained read-only body does not invalidate the bytes.
|
||||
_ = original.Close()
|
||||
r.Body = io.NopCloser(bytes.NewReader(buf))
|
||||
// Framing is deliberately left as the client sent it. Rewriting a chunked
|
||||
// request to a fixed Content-Length here would be invisible to the client
|
||||
// but not to the rest of the chain: a later body capture with a smaller cap
|
||||
// sees a known length over its cap and skips capture entirely, where an
|
||||
// unknown length would have given it a truncated prefix. Inspecting a
|
||||
// request must not change what any other layer gets to inspect.
|
||||
return buf, false, nil
|
||||
}
|
||||
|
||||
// replay returns a ReadCloser that yields the already-read prefix followed by
|
||||
// the remainder of the original stream, and closes the original.
|
||||
func replay(prefix []byte, rest io.ReadCloser) io.ReadCloser {
|
||||
return struct {
|
||||
io.Reader
|
||||
io.Closer
|
||||
}{
|
||||
Reader: io.MultiReader(bytes.NewReader(prefix), rest),
|
||||
Closer: rest,
|
||||
}
|
||||
}
|
||||
|
||||
// redactedPlaceholder replaces a credential value in the mirrored body. It is
|
||||
// inert for rule matching, and its fixed length leaks nothing about the secret.
|
||||
const redactedPlaceholder = "redacted"
|
||||
|
||||
// redactFormFields returns the body to mirror for a URL-encoded form, with the
|
||||
// values of the named fields replaced. The proxy's own password / PIN login
|
||||
// form posts to the service path itself, so without this the plaintext
|
||||
// credential would reach the Security Engine.
|
||||
//
|
||||
// Only the credential values are removed, never the whole body: dropping the
|
||||
// body outright would let a caller exempt any payload from inspection just by
|
||||
// appending a field named "password". Everything else in the form stays
|
||||
// inspectable, which is the point.
|
||||
//
|
||||
// Returns body unchanged when it is not a URL-encoded form or carries none of
|
||||
// the fields.
|
||||
//
|
||||
// Substitution happens on the raw bytes rather than by re-encoding parsed
|
||||
// values. Re-encoding would drop pairs that url.ParseQuery rejects, so a
|
||||
// payload hidden in a malformed pair alongside a credential-named field would
|
||||
// never be inspected while a tolerant backend parser still acted on it. Working
|
||||
// byte-wise also avoids reordering keys and normalizing escapes, so the engine
|
||||
// sees the same bytes the backend will.
|
||||
//
|
||||
// Field names match case-sensitively, on purpose: the caller passes the exact
|
||||
// names the login handler reads via r.FormValue, and that lookup is itself
|
||||
// case-sensitive. A "Password" field is therefore never a credential as far as
|
||||
// the proxy is concerned, and redacting it would only blind the WAF to a value
|
||||
// the proxy does not own.
|
||||
func redactFormFields(contentType string, body []byte, fields []string) []byte {
|
||||
if len(fields) == 0 || len(body) == 0 {
|
||||
return body
|
||||
}
|
||||
media, _, err := mime.ParseMediaType(contentType)
|
||||
if err != nil || media != "application/x-www-form-urlencoded" {
|
||||
return body
|
||||
}
|
||||
return redactURLEncoded(body, fields)
|
||||
}
|
||||
|
||||
// redactURLEncoded replaces the values of the named keys in a URL-encoded
|
||||
// key/value sequence, the shared syntax of a query string and a form body.
|
||||
func redactURLEncoded(raw []byte, fields []string) []byte {
|
||||
// Split on "&" only, matching how Go's form parser delimits pairs.
|
||||
segments := bytes.Split(raw, []byte("&"))
|
||||
redacted := false
|
||||
for i, segment := range segments {
|
||||
rawKey, _, hasValue := bytes.Cut(segment, []byte("="))
|
||||
if !hasValue {
|
||||
continue
|
||||
}
|
||||
// Compare the decoded name, so an escaped spelling of the field
|
||||
// ("pass%77ord") is redacted too: the reader decodes before looking it
|
||||
// up. A key that fails to decode never reaches that reader either,
|
||||
// since the parser drops the pair.
|
||||
name, err := url.QueryUnescape(string(rawKey))
|
||||
if err != nil || !slices.Contains(fields, name) {
|
||||
continue
|
||||
}
|
||||
// Keep the key bytes as sent and replace only the value. Assigning a
|
||||
// fresh slice leaves raw untouched, which matters: the caller restored
|
||||
// the request body from the same buffer.
|
||||
segments[i] = []byte(string(rawKey) + "=" + redactedPlaceholder)
|
||||
redacted = true
|
||||
}
|
||||
if !redacted {
|
||||
return raw
|
||||
}
|
||||
return bytes.Join(segments, []byte("&"))
|
||||
}
|
||||
|
||||
// redactQuery replaces the values of the named query parameters in a raw query
|
||||
// string, leaving every other byte as sent.
|
||||
func redactQuery(rawQuery string, params []string) string {
|
||||
if len(params) == 0 || rawQuery == "" {
|
||||
return rawQuery
|
||||
}
|
||||
return string(redactURLEncoded([]byte(rawQuery), params))
|
||||
}
|
||||
|
||||
// redactCookieHeader replaces the values of the named cookies in a Cookie
|
||||
// header, keeping the others intact: cookies are a zone WAF rules match on, so
|
||||
// dropping the whole header would cost real coverage.
|
||||
func redactCookieHeader(value string, names []string) string {
|
||||
if len(names) == 0 || value == "" {
|
||||
return value
|
||||
}
|
||||
parts := strings.Split(value, ";")
|
||||
redacted := false
|
||||
for i, part := range parts {
|
||||
name, _, hasValue := strings.Cut(part, "=")
|
||||
if !hasValue {
|
||||
continue
|
||||
}
|
||||
// Cookie names are case-sensitive and are not percent-decoded.
|
||||
if !slices.Contains(names, strings.TrimSpace(name)) {
|
||||
continue
|
||||
}
|
||||
parts[i] = name + "=" + redactedPlaceholder
|
||||
redacted = true
|
||||
}
|
||||
if !redacted {
|
||||
return value
|
||||
}
|
||||
return strings.Join(parts, ";")
|
||||
}
|
||||
@@ -1,571 +0,0 @@
|
||||
// Package appsec implements the CrowdSec AppSec (WAF) side of the remediation
|
||||
// component protocol: each inspected HTTP request is mirrored to the Security
|
||||
// Engine's AppSec endpoint, which replies with an allow / ban / captcha verdict
|
||||
// for that request.
|
||||
//
|
||||
// This is a separate endpoint from the LAPI decision stream used by the
|
||||
// crowdsec package: LAPI answers "is this IP known bad", AppSec answers "is
|
||||
// this request an attack". The two are configured and enabled independently.
|
||||
package appsec
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/netutil"
|
||||
"github.com/netbirdio/netbird/proxy/internal/restrict"
|
||||
)
|
||||
|
||||
// Header names the AppSec component reads off the mirrored request. IP, URI and
|
||||
// Verb are mandatory: the engine answers 500 when any of them is missing.
|
||||
const (
|
||||
headerAPIKey = "X-Crowdsec-Appsec-Api-Key" //nolint:gosec // G101: a header name, not a credential
|
||||
headerIP = "X-Crowdsec-Appsec-Ip"
|
||||
headerURI = "X-Crowdsec-Appsec-Uri"
|
||||
headerVerb = "X-Crowdsec-Appsec-Verb"
|
||||
headerHost = "X-Crowdsec-Appsec-Host"
|
||||
headerUserAgent = "X-Crowdsec-Appsec-User-Agent"
|
||||
headerHTTPVersion = "X-Crowdsec-Appsec-Http-Version"
|
||||
headerTransactionID = "X-Crowdsec-Appsec-Transaction-Id"
|
||||
)
|
||||
|
||||
// headerPrefix covers every protocol header. Any client-supplied header in this
|
||||
// namespace is dropped before forwarding so a caller cannot influence the
|
||||
// engine's view of its own address, or replay an API key.
|
||||
const headerPrefix = "X-Crowdsec-Appsec-"
|
||||
|
||||
// Remediation actions the engine can return.
|
||||
const (
|
||||
actionAllow = "allow"
|
||||
actionBan = "ban"
|
||||
actionCaptcha = "captcha"
|
||||
)
|
||||
|
||||
const (
|
||||
// DefaultTimeout matches the 200ms budget CrowdSec's remediation component
|
||||
// spec sets for the blocking AppSec call.
|
||||
DefaultTimeout = 200 * time.Millisecond
|
||||
// MinTimeout and MaxTimeout bound the configured inspection timeout.
|
||||
// Inspection is synchronous, so the upper bound is what keeps a
|
||||
// mis-set value from parking every request to an inspected service on a
|
||||
// slow engine; the lower bound keeps the call from timing out before the
|
||||
// engine can realistically answer. Mirrors the per-middleware bounds the
|
||||
// proxy already applies to in-path calls.
|
||||
MinTimeout = 10 * time.Millisecond
|
||||
MaxTimeout = 5 * time.Second
|
||||
// DefaultMaxBodyBytes caps the request body mirrored to the engine.
|
||||
// Requests with a larger body are inspected on headers and URI only.
|
||||
DefaultMaxBodyBytes int64 = 64 << 10
|
||||
// DefaultMaxConcurrent bounds inspections in flight toward the engine. The
|
||||
// point is to fail fast instead of parking a goroutine per request for the
|
||||
// whole timeout once the engine is saturated: a slow engine otherwise turns
|
||||
// a traffic burst into a pile of waiters that all time out anyway. Sized so
|
||||
// a healthy engine (single-digit milliseconds per call) never reaches it.
|
||||
DefaultMaxConcurrent = 256
|
||||
// MaxConcurrentLimit is the ceiling for that bound.
|
||||
MaxConcurrentLimit = 4096
|
||||
// MaxBodyBytesLimit is the ceiling for that cap. A single request can hold
|
||||
// this much in memory; the shared Budget is what bounds the total across
|
||||
// concurrent requests. Matches the proxy-wide body-capture ceiling.
|
||||
MaxBodyBytesLimit int64 = 8 << 20
|
||||
// maxResponseBytes bounds how much of a verdict response is read. The
|
||||
// engine answers with a two-field JSON object, so anything beyond this is
|
||||
// not a response we can act on.
|
||||
maxResponseBytes int64 = 4 << 10
|
||||
)
|
||||
|
||||
// Reasons the request body was not mirrored. Reported so an access-log reader
|
||||
// can distinguish "inspected and clean" from "never inspected", and so an
|
||||
// oversize opt-out is visible rather than silent.
|
||||
const (
|
||||
BypassOversize = "oversize"
|
||||
BypassUpgrade = "upgrade"
|
||||
BypassDisabled = "disabled"
|
||||
BypassBudget = "budget_exhausted"
|
||||
)
|
||||
|
||||
// ErrUnavailable reports that the engine could not produce a verdict: the call
|
||||
// failed, timed out, or the engine rejected it (401 bad key, 500 malformed).
|
||||
// Distinguished from a block verdict so the caller can apply the per-service
|
||||
// mode: enforce fails closed, observe allows.
|
||||
var ErrUnavailable = errors.New("appsec engine unavailable")
|
||||
|
||||
// Config configures a Client.
|
||||
type Config struct {
|
||||
// URL is the AppSec endpoint, e.g. http://127.0.0.1:7422/.
|
||||
URL string
|
||||
// APIKey is the CrowdSec bouncer API key. The AppSec component validates it
|
||||
// against LAPI, so the same key used for the decision stream works here.
|
||||
APIKey string
|
||||
// Timeout bounds a single inspection call. Zero means DefaultTimeout.
|
||||
Timeout time.Duration
|
||||
// MaxBodyBytes caps the mirrored request body. Zero means
|
||||
// DefaultMaxBodyBytes; negative disables body forwarding entirely.
|
||||
MaxBodyBytes int64
|
||||
// MaxConcurrent bounds inspections in flight toward the engine. Zero means
|
||||
// DefaultMaxConcurrent; negative disables the bound.
|
||||
MaxConcurrent int
|
||||
// Budget bounds the total body buffering in flight across all inspected
|
||||
// requests. Nil disables that ceiling, which leaves the worst case at
|
||||
// MaxBodyBytes times the concurrent request count; callers serving
|
||||
// untrusted traffic should share the proxy-wide capture budget here.
|
||||
Budget Budget
|
||||
Logger *log.Entry
|
||||
}
|
||||
|
||||
// Budget is the shared allowance for in-flight body buffering. Acquire reports
|
||||
// whether n bytes could be reserved; every successful Acquire is matched by a
|
||||
// Release of the same n. Satisfied by the proxy's capture budget, so AppSec and
|
||||
// the middleware body tap draw down one pool rather than two independent ones.
|
||||
type Budget interface {
|
||||
Acquire(n int64) bool
|
||||
Release(n int64)
|
||||
}
|
||||
|
||||
// Client mirrors HTTP requests to a CrowdSec AppSec endpoint. It holds no
|
||||
// per-service state and is safe for concurrent use.
|
||||
type Client struct {
|
||||
url string
|
||||
apiKey string
|
||||
maxBodyBytes int64
|
||||
// sem bounds in-flight inspections. Nil when the bound is disabled.
|
||||
sem chan struct{}
|
||||
budget Budget
|
||||
http *http.Client
|
||||
logger *log.Entry
|
||||
}
|
||||
|
||||
// New validates the config and returns a Client. The endpoint is not contacted
|
||||
// here: the engine may come up after the proxy.
|
||||
func New(cfg Config) (*Client, error) {
|
||||
if cfg.URL == "" {
|
||||
return nil, errors.New("appsec url is empty")
|
||||
}
|
||||
if cfg.APIKey == "" {
|
||||
return nil, errors.New("appsec api key is empty")
|
||||
}
|
||||
parsed, err := url.Parse(cfg.URL)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("parse appsec url: %w", err)
|
||||
}
|
||||
if parsed.Scheme != "http" && parsed.Scheme != "https" {
|
||||
return nil, fmt.Errorf("appsec url scheme %q is not http(s)", parsed.Scheme)
|
||||
}
|
||||
if parsed.Host == "" {
|
||||
return nil, errors.New("appsec url has no host")
|
||||
}
|
||||
|
||||
logger := cfg.Logger
|
||||
if logger == nil {
|
||||
logger = log.NewEntry(log.StandardLogger())
|
||||
}
|
||||
|
||||
timeout := cfg.Timeout
|
||||
switch {
|
||||
case timeout <= 0:
|
||||
timeout = DefaultTimeout
|
||||
case timeout < MinTimeout:
|
||||
logger.Warnf("appsec timeout %s is below the minimum, using %s", timeout, MinTimeout)
|
||||
timeout = MinTimeout
|
||||
case timeout > MaxTimeout:
|
||||
logger.Warnf("appsec timeout %s exceeds the maximum, using %s", timeout, MaxTimeout)
|
||||
timeout = MaxTimeout
|
||||
}
|
||||
|
||||
// A negative cap is meaningful: forward no body at all.
|
||||
maxBody := cfg.MaxBodyBytes
|
||||
switch {
|
||||
case maxBody == 0:
|
||||
maxBody = DefaultMaxBodyBytes
|
||||
case maxBody > MaxBodyBytesLimit:
|
||||
logger.Warnf("appsec max body %d exceeds the maximum, using %d", maxBody, MaxBodyBytesLimit)
|
||||
maxBody = MaxBodyBytesLimit
|
||||
}
|
||||
|
||||
maxConcurrent := cfg.MaxConcurrent
|
||||
switch {
|
||||
case maxConcurrent == 0:
|
||||
maxConcurrent = DefaultMaxConcurrent
|
||||
case maxConcurrent > MaxConcurrentLimit:
|
||||
logger.Warnf("appsec max concurrent %d exceeds the maximum, using %d", maxConcurrent, MaxConcurrentLimit)
|
||||
maxConcurrent = MaxConcurrentLimit
|
||||
}
|
||||
var sem chan struct{}
|
||||
if maxConcurrent > 0 {
|
||||
sem = make(chan struct{}, maxConcurrent)
|
||||
}
|
||||
|
||||
return &Client{
|
||||
url: cfg.URL,
|
||||
apiKey: cfg.APIKey,
|
||||
maxBodyBytes: maxBody,
|
||||
sem: sem,
|
||||
budget: cfg.Budget,
|
||||
logger: logger,
|
||||
http: &http.Client{
|
||||
Timeout: timeout,
|
||||
Transport: &http.Transport{
|
||||
MaxIdleConns: 100,
|
||||
MaxIdleConnsPerHost: 32,
|
||||
IdleConnTimeout: 90 * time.Second,
|
||||
},
|
||||
},
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Request is one inspection request.
|
||||
type Request struct {
|
||||
// HTTP is the in-flight client request. Inspect buffers and restores its
|
||||
// body, so the request stays forwardable afterwards.
|
||||
HTTP *http.Request
|
||||
// ClientIP is the resolved client address (after trusted-proxy handling).
|
||||
ClientIP netip.Addr
|
||||
// TransactionID correlates the engine's alert with the proxy's access log
|
||||
// entry. Empty lets the engine generate its own UUID.
|
||||
TransactionID string
|
||||
// RedactBodyFields lists form fields whose values are replaced before the
|
||||
// body is mirrored. Used to keep credentials submitted to the proxy's own
|
||||
// login form out of the engine while still inspecting the rest.
|
||||
RedactBodyFields []string
|
||||
// RedactHeaders, RedactCookies and RedactQueryParams name the credentials
|
||||
// the proxy already withholds from backends: the header-auth values, its
|
||||
// session cookie, and the OIDC session token. The engine logs and alerts on
|
||||
// what it inspects, so mirroring them there would reintroduce the leak the
|
||||
// upstream strippers exist to prevent. Only the values are replaced, so the
|
||||
// surrounding headers, cookies and query stay inspectable.
|
||||
RedactHeaders []string
|
||||
RedactCookies []string
|
||||
RedactQueryParams []string
|
||||
}
|
||||
|
||||
// Result is the outcome of an inspection.
|
||||
type Result struct {
|
||||
Verdict restrict.Verdict
|
||||
// BodyBypass names why the request body was not mirrored, empty when it
|
||||
// was (or when the request had none). The engine still saw the headers and
|
||||
// URI, so this is a coverage note, not a failure.
|
||||
BodyBypass string
|
||||
// Release returns the buffered body's budget reservation. Never nil, so it
|
||||
// is always safe to defer. It must run only once the request has been
|
||||
// served, not when Inspect returns: the buffer stays alive as r.Body for
|
||||
// the backend to read, so releasing earlier would let the budget admit
|
||||
// buffering that is still resident.
|
||||
Release func()
|
||||
}
|
||||
|
||||
// noopRelease is the Release for inspections that reserved no budget.
|
||||
func noopRelease() {}
|
||||
|
||||
// Inspect mirrors r to the AppSec engine and returns its verdict. A nil error
|
||||
// with restrict.Allow means the request passed. On failure it returns
|
||||
// DenyAppSecUnavailable wrapped with ErrUnavailable; the caller decides whether
|
||||
// that blocks, based on the per-service mode.
|
||||
func (c *Client) Inspect(ctx context.Context, req Request) (Result, error) {
|
||||
if c == nil {
|
||||
return Result{Verdict: restrict.DenyAppSecUnavailable, Release: noopRelease}, ErrUnavailable
|
||||
}
|
||||
if req.HTTP == nil {
|
||||
return Result{Verdict: restrict.DenyAppSecUnavailable, Release: noopRelease}, fmt.Errorf("%w: nil request", ErrUnavailable)
|
||||
}
|
||||
|
||||
// release is carried out to the caller rather than deferred here: the
|
||||
// buffered body outlives this call as r.Body.
|
||||
if !c.acquireSlot() {
|
||||
// Deny rather than wave through: a flood must not be a way to switch
|
||||
// inspection off. Enforce blocks, observe logs and allows, exactly as
|
||||
// for an unreachable engine.
|
||||
return Result{Verdict: restrict.DenyAppSecUnavailable, Release: noopRelease},
|
||||
fmt.Errorf("%w: %d inspections already in flight", ErrUnavailable, cap(c.sem))
|
||||
}
|
||||
defer c.releaseSlot()
|
||||
|
||||
body, bypass, release, err := c.readBody(req)
|
||||
if err != nil {
|
||||
return Result{Verdict: restrict.DenyAppSecUnavailable, Release: release}, fmt.Errorf("%w: read body: %w", ErrUnavailable, err)
|
||||
}
|
||||
|
||||
outbound, err := c.buildRequest(ctx, req, body)
|
||||
if err != nil {
|
||||
return Result{Verdict: restrict.DenyAppSecUnavailable, BodyBypass: bypass, Release: release}, fmt.Errorf("%w: %w", ErrUnavailable, err)
|
||||
}
|
||||
|
||||
resp, err := c.http.Do(outbound)
|
||||
if err != nil {
|
||||
return Result{Verdict: restrict.DenyAppSecUnavailable, BodyBypass: bypass, Release: release}, fmt.Errorf("%w: %w", ErrUnavailable, err)
|
||||
}
|
||||
defer func() {
|
||||
// Drain before closing. net/http only returns a connection to the idle
|
||||
// pool once its body is read to EOF; closing with bytes outstanding
|
||||
// discards it. Every verdict carries a JSON body, so skipping this
|
||||
// would cost a fresh handshake per inspected request, inside the
|
||||
// timeout budget.
|
||||
if _, err := io.Copy(io.Discard, io.LimitReader(resp.Body, maxResponseBytes)); err != nil {
|
||||
c.logger.Tracef("drain appsec response body: %v", err)
|
||||
}
|
||||
if err := resp.Body.Close(); err != nil {
|
||||
c.logger.Tracef("close appsec response body: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
verdict, err := c.verdict(resp)
|
||||
return Result{Verdict: verdict, BodyBypass: bypass, Release: release}, err
|
||||
}
|
||||
|
||||
// acquireSlot takes an in-flight slot without blocking, reporting false when
|
||||
// the engine is already at capacity.
|
||||
func (c *Client) acquireSlot() bool {
|
||||
if c.sem == nil {
|
||||
return true
|
||||
}
|
||||
select {
|
||||
case c.sem <- struct{}{}:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// releaseSlot returns the slot. Scoped to the engine call, not the request: the
|
||||
// buffered body outlives the call but the engine's attention does not.
|
||||
func (c *Client) releaseSlot() {
|
||||
if c.sem == nil {
|
||||
return
|
||||
}
|
||||
select {
|
||||
case <-c.sem:
|
||||
default:
|
||||
}
|
||||
}
|
||||
|
||||
// readBody buffers the body so it can be mirrored, always restoring it on the
|
||||
// original request. Returns nil when there is no body to forward: no body at
|
||||
// all, an upgrade request, or a body over the cap. A login form is forwarded
|
||||
// with its credential values redacted rather than suppressed.
|
||||
// release is never nil; the caller invokes it once the request has been served.
|
||||
func (c *Client) readBody(req Request) (body []byte, bypass string, release func(), err error) {
|
||||
r := req.HTTP
|
||||
if r.Body == nil || r.Body == http.NoBody {
|
||||
return nil, "", noopRelease, nil
|
||||
}
|
||||
if c.maxBodyBytes < 0 {
|
||||
return nil, BypassDisabled, noopRelease, nil
|
||||
}
|
||||
// A genuine upgrade request carries no body to inspect (net/http hands us
|
||||
// http.NoBody, caught above); the hijacked stream is reached through
|
||||
// Hijacker, never r.Body. The test has to be the forwarder's own, because a
|
||||
// looser one would skip inspection for requests the forwarder still
|
||||
// delivers to the backend with their body intact.
|
||||
if netutil.IsUpgradeRequest(r.Header) {
|
||||
return nil, BypassUpgrade, noopRelease, nil
|
||||
}
|
||||
// A Content-Length over the cap is known to be too large before reading.
|
||||
if r.ContentLength > c.maxBodyBytes {
|
||||
return nil, BypassOversize, noopRelease, nil
|
||||
}
|
||||
|
||||
// Reserve the whole cap rather than the eventual length: the reservation
|
||||
// has to be made before the body is read, and until then the only bound
|
||||
// known is the cap. Skipping inspection when the pool is drained keeps a
|
||||
// burst of large bodies from being an out-of-memory lever; the bypass is
|
||||
// recorded so the gap in coverage is visible.
|
||||
release = noopRelease
|
||||
if c.budget != nil {
|
||||
if !c.budget.Acquire(c.maxBodyBytes) {
|
||||
c.logger.Debugf("appsec buffer budget exhausted, inspecting headers and URI only")
|
||||
return nil, BypassBudget, noopRelease, nil
|
||||
}
|
||||
var once sync.Once
|
||||
release = func() { once.Do(func() { c.budget.Release(c.maxBodyBytes) }) }
|
||||
}
|
||||
|
||||
buffered, oversize, err := bufferBody(r, c.maxBodyBytes)
|
||||
if err != nil {
|
||||
// bufferBody restored r.Body from the bytes it did read, so the
|
||||
// reservation stays held until the caller releases it.
|
||||
return nil, "", release, err
|
||||
}
|
||||
// An oversize body was only partially read: a truncated prefix changes the
|
||||
// engine's verdict in both directions, so inspect headers and URI only.
|
||||
if oversize {
|
||||
return nil, BypassOversize, release, nil
|
||||
}
|
||||
return redactFormFields(r.Header.Get("Content-Type"), buffered, req.RedactBodyFields), "", release, nil
|
||||
}
|
||||
|
||||
// buildRequest assembles the mirrored request. Per the protocol it is a GET
|
||||
// when there is no body and a POST otherwise; bytes.Reader gives the outbound
|
||||
// request an accurate Content-Length, which the engine relies on to read the
|
||||
// body at all.
|
||||
func (c *Client) buildRequest(ctx context.Context, req Request, body []byte) (*http.Request, error) {
|
||||
method := http.MethodGet
|
||||
var payload io.Reader
|
||||
if len(body) > 0 {
|
||||
method = http.MethodPost
|
||||
payload = bytes.NewReader(body)
|
||||
}
|
||||
|
||||
outbound, err := http.NewRequestWithContext(ctx, method, c.url, payload)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("build appsec request: %w", err)
|
||||
}
|
||||
|
||||
r := req.HTTP
|
||||
copyInspectableHeaders(outbound.Header, r.Header)
|
||||
redactSecrets(outbound.Header, req)
|
||||
|
||||
outbound.Header.Set(headerAPIKey, c.apiKey)
|
||||
outbound.Header.Set(headerIP, req.ClientIP.Unmap().String())
|
||||
outbound.Header.Set(headerURI, mirroredURI(r.URL, req.RedactQueryParams))
|
||||
outbound.Header.Set(headerVerb, r.Method)
|
||||
outbound.Header.Set(headerHost, r.Host)
|
||||
if ua := r.UserAgent(); ua != "" {
|
||||
outbound.Header.Set(headerUserAgent, ua)
|
||||
}
|
||||
outbound.Header.Set(headerHTTPVersion, httpVersion(r))
|
||||
if req.TransactionID != "" {
|
||||
outbound.Header.Set(headerTransactionID, req.TransactionID)
|
||||
}
|
||||
return outbound, nil
|
||||
}
|
||||
|
||||
// verdict maps the engine's response to a restrict.Verdict. 200 is a pass and
|
||||
// 401/500 are engine-side failures; every other status carries a remediation in
|
||||
// the body. The blocked status code is operator-configurable
|
||||
// (blocked_http_code), so the action field decides, not the status.
|
||||
func (c *Client) verdict(resp *http.Response) (restrict.Verdict, error) {
|
||||
switch resp.StatusCode {
|
||||
case http.StatusUnauthorized:
|
||||
return restrict.DenyAppSecUnavailable, fmt.Errorf("%w: rejected api key", ErrUnavailable)
|
||||
case http.StatusInternalServerError:
|
||||
return restrict.DenyAppSecUnavailable, fmt.Errorf("%w: engine error", ErrUnavailable)
|
||||
}
|
||||
|
||||
// Every status, 200 included, has to carry a decodable remediation. Taking a
|
||||
// bare 200 as a pass would mean a URL pointing at anything that answers 200
|
||||
// (a health endpoint, a load balancer's default page) silently allows every
|
||||
// request while the service reports itself as enforcing.
|
||||
|
||||
var decoded struct {
|
||||
Action string `json:"action"`
|
||||
}
|
||||
if err := json.NewDecoder(io.LimitReader(resp.Body, maxResponseBytes)).Decode(&decoded); err != nil {
|
||||
// Every remediation carries a decodable action, so a response without
|
||||
// one is not a verdict: most often the URL points at something that is
|
||||
// not the AppSec endpoint, which answers 404 with HTML. Reported as
|
||||
// unavailable rather than a ban so the access log names the real fault
|
||||
// instead of sending an operator hunting for a rule that never fired.
|
||||
// Enforce still blocks either way; only the recorded reason differs.
|
||||
return restrict.DenyAppSecUnavailable, fmt.Errorf("%w: undecodable response (status %d): %w", ErrUnavailable, resp.StatusCode, err)
|
||||
}
|
||||
|
||||
switch decoded.Action {
|
||||
case actionAllow:
|
||||
return restrict.Allow, nil
|
||||
case actionCaptcha:
|
||||
return restrict.DenyAppSecCaptcha, nil
|
||||
case actionBan:
|
||||
return restrict.DenyAppSecBan, nil
|
||||
case "":
|
||||
// Decodable JSON without a remediation is not a verdict either: the
|
||||
// endpoint answered, but not as the engine. Same reasoning as an
|
||||
// undecodable body, and the same reason to point at configuration.
|
||||
return restrict.DenyAppSecUnavailable,
|
||||
fmt.Errorf("%w: response carried no remediation (status %d)", ErrUnavailable, resp.StatusCode)
|
||||
default:
|
||||
// A remediation we do not implement still means the engine flagged the
|
||||
// request, so deny.
|
||||
c.logger.Debugf("unknown appsec action %q (status %d), treating as ban", decoded.Action, resp.StatusCode)
|
||||
return restrict.DenyAppSecBan, nil
|
||||
}
|
||||
}
|
||||
|
||||
// copyInspectableHeaders copies the client's headers, which are what the WAF
|
||||
// rules actually match on, dropping hop-by-hop headers that describe the
|
||||
// proxy-to-engine connection rather than the client request, and any header in
|
||||
// the AppSec protocol namespace.
|
||||
func copyInspectableHeaders(dst, src http.Header) {
|
||||
for name, values := range src {
|
||||
if hopByHopHeaders[http.CanonicalHeaderKey(name)] {
|
||||
continue
|
||||
}
|
||||
if strings.HasPrefix(http.CanonicalHeaderKey(name), headerPrefix) {
|
||||
continue
|
||||
}
|
||||
dst[http.CanonicalHeaderKey(name)] = append([]string(nil), values...)
|
||||
}
|
||||
// Content-Length describes the mirrored payload, not the client's: net/http
|
||||
// sets it from the body we actually attach. Content-Type is kept either way
|
||||
// so rules matching on it still fire when the body was not forwarded.
|
||||
dst.Del("Content-Length")
|
||||
}
|
||||
|
||||
// redactSecrets replaces the credential values the proxy withholds from
|
||||
// backends, so the mirrored copy does not carry them either.
|
||||
func redactSecrets(dst http.Header, req Request) {
|
||||
for _, name := range req.RedactHeaders {
|
||||
// Presence, not Get: a header whose first value is empty still carries
|
||||
// its later values to the engine, while the upstream strip deletes the
|
||||
// name outright. Set collapses every value into the placeholder.
|
||||
if len(dst.Values(name)) > 0 {
|
||||
dst.Set(name, redactedPlaceholder)
|
||||
}
|
||||
}
|
||||
// Every Cookie line, not just the first: a client may send several, and Get
|
||||
// would leave the session cookie in any later one mirrored in the clear.
|
||||
if cookies := dst.Values("Cookie"); len(cookies) > 0 {
|
||||
redacted := make([]string, len(cookies))
|
||||
for i, cookie := range cookies {
|
||||
redacted[i] = redactCookieHeader(cookie, req.RedactCookies)
|
||||
}
|
||||
dst["Cookie"] = redacted
|
||||
}
|
||||
}
|
||||
|
||||
// mirroredURI renders the request target for the URI header, with the named
|
||||
// query parameter values replaced.
|
||||
func mirroredURI(u *url.URL, redactParams []string) string {
|
||||
uri := u.RequestURI()
|
||||
if u.RawQuery == "" || len(redactParams) == 0 {
|
||||
return uri
|
||||
}
|
||||
redacted := redactQuery(u.RawQuery, redactParams)
|
||||
if redacted == u.RawQuery {
|
||||
return uri
|
||||
}
|
||||
// RequestURI is path + "?" + RawQuery; swap only the query part so the
|
||||
// path keeps its original encoding.
|
||||
return strings.TrimSuffix(uri, u.RawQuery) + redacted
|
||||
}
|
||||
|
||||
var hopByHopHeaders = map[string]bool{
|
||||
"Connection": true,
|
||||
"Keep-Alive": true,
|
||||
"Proxy-Authenticate": true,
|
||||
"Proxy-Authorization": true,
|
||||
"Proxy-Connection": true,
|
||||
"Te": true,
|
||||
"Trailer": true,
|
||||
"Transfer-Encoding": true,
|
||||
"Upgrade": true,
|
||||
}
|
||||
|
||||
// httpVersion renders the two-digit form the engine parses ("11", "20").
|
||||
func httpVersion(r *http.Request) string {
|
||||
major, minor := r.ProtoMajor, r.ProtoMinor
|
||||
if major < 0 || major > 9 || minor < 0 || minor > 9 {
|
||||
return ""
|
||||
}
|
||||
return fmt.Sprintf("%d%d", major, minor)
|
||||
}
|
||||
File diff suppressed because it is too large
Load Diff
@@ -1,306 +0,0 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"crypto/ed25519"
|
||||
"encoding/base64"
|
||||
"io"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/appsec"
|
||||
"github.com/netbirdio/netbird/proxy/internal/proxy"
|
||||
"github.com/netbirdio/netbird/proxy/internal/restrict"
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
)
|
||||
|
||||
// appsecEngine is a stub AppSec component returning a fixed remediation.
|
||||
func appsecEngine(t *testing.T, status int, body string) *httptest.Server {
|
||||
t.Helper()
|
||||
srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(status)
|
||||
if body != "" {
|
||||
_, _ = w.Write([]byte(body))
|
||||
}
|
||||
}))
|
||||
t.Cleanup(srv.Close)
|
||||
return srv
|
||||
}
|
||||
|
||||
// serveWithAppSec runs a request through the middleware for a domain configured
|
||||
// with the given AppSec mode, returning the response and the captured metadata.
|
||||
func serveWithAppSec(t *testing.T, mode restrict.AppSecMode, client *appsec.Client, r *http.Request) (*httptest.ResponseRecorder, map[string]string, bool) {
|
||||
t.Helper()
|
||||
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
mw.SetAppSec(client)
|
||||
require.NoError(t, mw.AddDomain("svc.example.com", DomainSettings{
|
||||
AccountID: types.AccountID("acct-1"),
|
||||
ServiceID: types.ServiceID("svc-1"),
|
||||
AppSecMode: mode,
|
||||
}))
|
||||
|
||||
reached := false
|
||||
handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
reached = true
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
|
||||
cd := proxy.NewCapturedData("req-1")
|
||||
r = r.WithContext(proxy.WithCapturedData(r.Context(), cd))
|
||||
rec := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rec, r)
|
||||
|
||||
return rec, cd.GetMetadata(), reached
|
||||
}
|
||||
|
||||
func appsecRequest() *http.Request {
|
||||
r := httptest.NewRequest(http.MethodGet, "http://svc.example.com/?x=/etc/passwd", nil)
|
||||
r.Host = "svc.example.com"
|
||||
r.RemoteAddr = "203.0.113.7:44444"
|
||||
return r
|
||||
}
|
||||
|
||||
func TestCheckAppSec_EnforceBlocksBannedRequest(t *testing.T) {
|
||||
srv := appsecEngine(t, http.StatusForbidden, `{"action":"ban","http_status":403}`)
|
||||
client, err := appsec.New(appsec.Config{URL: srv.URL, APIKey: "k"})
|
||||
require.NoError(t, err)
|
||||
|
||||
rec, meta, reached := serveWithAppSec(t, restrict.AppSecEnforce, client, appsecRequest())
|
||||
|
||||
assert.Equal(t, http.StatusForbidden, rec.Code)
|
||||
assert.False(t, reached, "a banned request must not reach the backend")
|
||||
assert.Equal(t, "appsec_ban", meta["appsec_verdict"])
|
||||
assert.NotContains(t, meta, "appsec_mode", "enforce is the default, only observe is annotated")
|
||||
}
|
||||
|
||||
func TestCheckAppSec_ObserveAllowsAndRecordsVerdict(t *testing.T) {
|
||||
srv := appsecEngine(t, http.StatusForbidden, `{"action":"ban","http_status":403}`)
|
||||
client, err := appsec.New(appsec.Config{URL: srv.URL, APIKey: "k"})
|
||||
require.NoError(t, err)
|
||||
|
||||
rec, meta, reached := serveWithAppSec(t, restrict.AppSecObserve, client, appsecRequest())
|
||||
|
||||
assert.Equal(t, http.StatusOK, rec.Code)
|
||||
assert.True(t, reached, "observe mode must not block")
|
||||
assert.Equal(t, "appsec_ban", meta["appsec_verdict"])
|
||||
assert.Equal(t, "observe", meta["appsec_mode"])
|
||||
}
|
||||
|
||||
func TestCheckAppSec_AllowedRequestPasses(t *testing.T) {
|
||||
srv := appsecEngine(t, http.StatusOK, `{"action":"allow","http_status":200}`)
|
||||
client, err := appsec.New(appsec.Config{URL: srv.URL, APIKey: "k"})
|
||||
require.NoError(t, err)
|
||||
|
||||
rec, meta, reached := serveWithAppSec(t, restrict.AppSecEnforce, client, appsecRequest())
|
||||
|
||||
assert.Equal(t, http.StatusOK, rec.Code)
|
||||
assert.True(t, reached)
|
||||
assert.NotContains(t, meta, "appsec_verdict", "a clean request records no verdict")
|
||||
}
|
||||
|
||||
func TestCheckAppSec_OffSkipsInspection(t *testing.T) {
|
||||
// An engine that would ban everything; the mode must keep us away from it.
|
||||
srv := appsecEngine(t, http.StatusForbidden, `{"action":"ban"}`)
|
||||
client, err := appsec.New(appsec.Config{URL: srv.URL, APIKey: "k"})
|
||||
require.NoError(t, err)
|
||||
|
||||
rec, meta, reached := serveWithAppSec(t, restrict.AppSecOff, client, appsecRequest())
|
||||
|
||||
assert.Equal(t, http.StatusOK, rec.Code)
|
||||
assert.True(t, reached)
|
||||
assert.Empty(t, meta)
|
||||
}
|
||||
|
||||
func TestCheckAppSec_EnforceFailsClosedWithoutClient(t *testing.T) {
|
||||
rec, meta, reached := serveWithAppSec(t, restrict.AppSecEnforce, nil, appsecRequest())
|
||||
|
||||
assert.Equal(t, http.StatusForbidden, rec.Code,
|
||||
"enforce with no configured endpoint must deny rather than pass traffic uninspected")
|
||||
assert.False(t, reached)
|
||||
assert.Equal(t, "appsec_unavailable", meta["appsec_verdict"])
|
||||
}
|
||||
|
||||
func TestCheckAppSec_ObserveAllowsWithoutClient(t *testing.T) {
|
||||
rec, meta, reached := serveWithAppSec(t, restrict.AppSecObserve, nil, appsecRequest())
|
||||
|
||||
assert.Equal(t, http.StatusOK, rec.Code)
|
||||
assert.True(t, reached)
|
||||
assert.Equal(t, "appsec_unavailable", meta["appsec_verdict"])
|
||||
assert.Equal(t, "observe", meta["appsec_mode"])
|
||||
}
|
||||
|
||||
func TestCheckAppSec_EnforceFailsClosedWhenEngineUnreachable(t *testing.T) {
|
||||
client, err := appsec.New(appsec.Config{URL: "http://127.0.0.1:1/", APIKey: "k"})
|
||||
require.NoError(t, err)
|
||||
|
||||
rec, meta, reached := serveWithAppSec(t, restrict.AppSecEnforce, client, appsecRequest())
|
||||
|
||||
assert.Equal(t, http.StatusForbidden, rec.Code)
|
||||
assert.False(t, reached)
|
||||
assert.Equal(t, "appsec_unavailable", meta["appsec_verdict"])
|
||||
}
|
||||
|
||||
func TestCheckAppSec_InspectsOverlayTraffic(t *testing.T) {
|
||||
srv := appsecEngine(t, http.StatusForbidden, `{"action":"ban"}`)
|
||||
client, err := appsec.New(appsec.Config{URL: srv.URL, APIKey: "k"})
|
||||
require.NoError(t, err)
|
||||
|
||||
// Requests arriving over the WireGuard overlay skip the geo and IP-reputation
|
||||
// checks, but request content is just as inspectable.
|
||||
r := appsecRequest()
|
||||
r = r.WithContext(types.WithOverlayOrigin(r.Context()))
|
||||
|
||||
rec, meta, reached := serveWithAppSec(t, restrict.AppSecEnforce, client, r)
|
||||
|
||||
assert.Equal(t, http.StatusForbidden, rec.Code, "overlay traffic must still be inspected")
|
||||
assert.False(t, reached)
|
||||
assert.Equal(t, "appsec_ban", meta["appsec_verdict"])
|
||||
}
|
||||
|
||||
func TestCheckAppSec_UnresolvableClientIPFailsClosed(t *testing.T) {
|
||||
srv := appsecEngine(t, http.StatusOK, `{"action":"allow"}`)
|
||||
client, err := appsec.New(appsec.Config{URL: srv.URL, APIKey: "k"})
|
||||
require.NoError(t, err)
|
||||
|
||||
r := appsecRequest()
|
||||
r.RemoteAddr = "not-an-address"
|
||||
|
||||
rec, meta, reached := serveWithAppSec(t, restrict.AppSecEnforce, client, r)
|
||||
|
||||
assert.Equal(t, http.StatusForbidden, rec.Code,
|
||||
"the engine requires a client address; a request we cannot attribute must not pass")
|
||||
assert.False(t, reached)
|
||||
assert.Equal(t, "appsec_unavailable", meta["appsec_verdict"])
|
||||
}
|
||||
|
||||
// The redaction sets are resolved from the domain's schemes at registration, so
|
||||
// what AppSec withholds cannot drift from what those schemes actually accept.
|
||||
func TestAddDomain_ResolvesRedactionSetsFromSchemes(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
require.NoError(t, mw.AddDomain("svc.example.com", DomainSettings{
|
||||
Schemes: []Scheme{
|
||||
NewPassword(nil, "svc-1", "acct-1"),
|
||||
NewHeader(nil, "svc-1", "acct-1", "X-Api-Key"),
|
||||
},
|
||||
SessionPublicKey: base64.StdEncoding.EncodeToString(make([]byte, ed25519.PublicKeySize)),
|
||||
SessionExpiration: time.Hour,
|
||||
AppSecMode: restrict.AppSecEnforce,
|
||||
}))
|
||||
|
||||
mw.domainsMux.RLock()
|
||||
config := mw.domains["svc.example.com"]
|
||||
mw.domainsMux.RUnlock()
|
||||
|
||||
assert.Equal(t, []string{"password"}, config.redactBodyFields)
|
||||
assert.Equal(t, []string{"X-Api-Key"}, config.redactHeaders)
|
||||
// r.FormValue merges the query into the form, so a credential passed there
|
||||
// authenticates and must be redacted alongside the OIDC session token.
|
||||
assert.Equal(t, []string{"session_token", "password"}, config.redactQueryParams)
|
||||
}
|
||||
|
||||
// countingBudget records reservations so a test can observe when the
|
||||
// middleware hands them back.
|
||||
type countingBudget struct {
|
||||
mu sync.Mutex
|
||||
total int64
|
||||
used int64
|
||||
maxAtOnce int64
|
||||
}
|
||||
|
||||
func (b *countingBudget) Acquire(n int64) bool {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
if b.used+n > b.total {
|
||||
return false
|
||||
}
|
||||
b.used += n
|
||||
if b.used > b.maxAtOnce {
|
||||
b.maxAtOnce = b.used
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
func (b *countingBudget) Release(n int64) {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
b.used -= n
|
||||
}
|
||||
|
||||
func (b *countingBudget) inUse() int64 {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
return b.used
|
||||
}
|
||||
|
||||
// The buffered body stays alive as r.Body until the backend has read it, so
|
||||
// Protect must hold the reservation for the whole request and return it only
|
||||
// once the handler chain has unwound. Releasing inside Inspect would let the
|
||||
// budget admit buffering that is still resident.
|
||||
func TestProtect_AppSecBudgetHeldForRequestAndReleasedAfter(t *testing.T) {
|
||||
srv := appsecEngine(t, http.StatusOK, `{"action":"allow"}`)
|
||||
budget := &countingBudget{total: 1 << 20}
|
||||
client, err := appsec.New(appsec.Config{
|
||||
URL: srv.URL,
|
||||
APIKey: "k",
|
||||
MaxBodyBytes: 4096,
|
||||
Budget: budget,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
mw.SetAppSec(client)
|
||||
require.NoError(t, mw.AddDomain("svc.example.com", DomainSettings{
|
||||
AccountID: types.AccountID("acct-1"),
|
||||
ServiceID: types.ServiceID("svc-1"),
|
||||
AppSecMode: restrict.AppSecEnforce,
|
||||
}))
|
||||
|
||||
var inHandler int64
|
||||
handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
// The backend reads the buffered body here, so the reservation must
|
||||
// still be held at this point.
|
||||
inHandler = budget.inUse()
|
||||
_, _ = io.ReadAll(r.Body)
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
|
||||
r := httptest.NewRequest(http.MethodPost, "http://svc.example.com/", strings.NewReader("payload"))
|
||||
r.Host = "svc.example.com"
|
||||
cd := proxy.NewCapturedData("req-1")
|
||||
r = r.WithContext(proxy.WithCapturedData(r.Context(), cd))
|
||||
|
||||
handler.ServeHTTP(httptest.NewRecorder(), r)
|
||||
|
||||
assert.Equal(t, int64(4096), inHandler, "the reservation must be held while the backend reads the body")
|
||||
assert.Equal(t, int64(0), budget.inUse(), "Protect must release the reservation once the request is served")
|
||||
}
|
||||
|
||||
// A denied request never reaches the backend, but Protect still has to hand the
|
||||
// reservation back or the pool leaks one cap per blocked request.
|
||||
func TestProtect_AppSecBudgetReleasedOnDeny(t *testing.T) {
|
||||
srv := appsecEngine(t, http.StatusForbidden, `{"action":"ban"}`)
|
||||
budget := &countingBudget{total: 1 << 20}
|
||||
client, err := appsec.New(appsec.Config{
|
||||
URL: srv.URL,
|
||||
APIKey: "k",
|
||||
MaxBodyBytes: 4096,
|
||||
Budget: budget,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
|
||||
r := httptest.NewRequest(http.MethodPost, "http://svc.example.com/", strings.NewReader("payload"))
|
||||
r.Host = "svc.example.com"
|
||||
rec, _, reached := serveWithAppSec(t, restrict.AppSecEnforce, client, r)
|
||||
|
||||
assert.False(t, reached, "a banned request must not reach the backend")
|
||||
assert.Equal(t, http.StatusForbidden, rec.Code)
|
||||
assert.Equal(t, int64(0), budget.inUse(), "a blocked request must not leak its reservation")
|
||||
}
|
||||
@@ -39,11 +39,6 @@ func (Header) Type() auth.Method {
|
||||
return auth.MethodHeader
|
||||
}
|
||||
|
||||
// HeaderName returns the request header this scheme reads its credential from.
|
||||
func (h Header) HeaderName() string {
|
||||
return h.headerName
|
||||
}
|
||||
|
||||
// Authenticate checks for the configured header in the request. If absent,
|
||||
// returns empty (unauthenticated). If present, validates via gRPC.
|
||||
func (h Header) Authenticate(r *http.Request) (string, string, error) {
|
||||
|
||||
@@ -18,7 +18,6 @@ import (
|
||||
"google.golang.org/grpc"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/auth"
|
||||
"github.com/netbirdio/netbird/proxy/internal/appsec"
|
||||
"github.com/netbirdio/netbird/proxy/internal/proxy"
|
||||
"github.com/netbirdio/netbird/proxy/internal/restrict"
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
@@ -26,11 +25,6 @@ import (
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
// sessionCookieNames is the cookie set AppSec redacts before mirroring: the
|
||||
// proxy's session cookie is a bearer credential for the service, and the
|
||||
// reverse proxy already strips it before forwarding upstream.
|
||||
var sessionCookieNames = []string{auth.SessionCookieName}
|
||||
|
||||
// errValidationUnavailable indicates that session validation failed due to
|
||||
// an infrastructure error (e.g. gRPC unavailable), not an invalid token.
|
||||
var errValidationUnavailable = errors.New("session validation unavailable")
|
||||
@@ -65,14 +59,6 @@ type DomainConfig struct {
|
||||
IPRestrictions *restrict.Filter
|
||||
// Private routes the domain through ValidateTunnelPeer; failure → 403.
|
||||
Private bool
|
||||
// AppSecMode enables CrowdSec AppSec request inspection for this domain.
|
||||
AppSecMode restrict.AppSecMode
|
||||
// redact* name the credentials this domain's schemes accept, resolved once
|
||||
// at registration. AppSec replaces their values before mirroring a request,
|
||||
// matching what the reverse proxy strips before forwarding upstream.
|
||||
redactBodyFields []string
|
||||
redactHeaders []string
|
||||
redactQueryParams []string
|
||||
}
|
||||
|
||||
type validationResult struct {
|
||||
@@ -96,9 +82,6 @@ type Middleware struct {
|
||||
sessionValidator SessionValidator
|
||||
geo restrict.GeoResolver
|
||||
tunnelCache *tunnelValidationCache
|
||||
// appsec is the shared CrowdSec AppSec client, nil when the proxy has no
|
||||
// AppSec endpoint configured. Set once during startup, before serving.
|
||||
appsec *appsec.Client
|
||||
}
|
||||
|
||||
// NewMiddleware creates a new authentication middleware. The sessionValidator is
|
||||
@@ -116,12 +99,6 @@ func NewMiddleware(logger *log.Logger, sessionValidator SessionValidator, geo re
|
||||
}
|
||||
}
|
||||
|
||||
// SetAppSec installs the shared CrowdSec AppSec client. Must be called during
|
||||
// startup, before the middleware serves any request.
|
||||
func (mw *Middleware) SetAppSec(client *appsec.Client) {
|
||||
mw.appsec = client
|
||||
}
|
||||
|
||||
// Protect wraps next with per-domain authentication and IP restriction checks.
|
||||
// Requests whose Host is not registered pass through unchanged.
|
||||
func (mw *Middleware) Protect(next http.Handler) http.Handler {
|
||||
@@ -146,14 +123,6 @@ func (mw *Middleware) Protect(next http.Handler) http.Handler {
|
||||
return
|
||||
}
|
||||
|
||||
// Deferred, not released here: the inspected body stays alive as r.Body
|
||||
// until the backend has read it, which happens inside next.ServeHTTP.
|
||||
appSecAllowed, releaseAppSec := mw.checkAppSec(w, r, config)
|
||||
defer releaseAppSec()
|
||||
if !appSecAllowed {
|
||||
return
|
||||
}
|
||||
|
||||
// Private services bypass operator schemes and gate on tunnel peer.
|
||||
if config.Private {
|
||||
if mw.forwardWithTunnelPeer(w, r, host, config, next) {
|
||||
@@ -293,134 +262,6 @@ func (mw *Middleware) checkIPRestrictions(w http.ResponseWriter, r *http.Request
|
||||
return false
|
||||
}
|
||||
|
||||
// checkAppSec mirrors the request to the CrowdSec AppSec engine when the domain
|
||||
// enables inspection. Returns false when the request was blocked and a response
|
||||
// has been written.
|
||||
//
|
||||
// The returned release frees the body-buffering budget the inspection reserved
|
||||
// and is never nil. It must run only after the request has been served, since
|
||||
// the buffered body stays alive as r.Body for the backend to read.
|
||||
//
|
||||
// Every non-allow remediation blocks with 403, captcha included: the proxy has
|
||||
// no challenge flow to serve. The distinct verdict is still recorded so the
|
||||
// access log shows which remediation the engine actually chose.
|
||||
//
|
||||
// Unlike the geo and IP-reputation checks, this runs for overlay traffic too:
|
||||
// AppSec inspects request content, which is just as meaningful when the caller
|
||||
// reached the proxy through the WireGuard tunnel.
|
||||
func (mw *Middleware) checkAppSec(w http.ResponseWriter, r *http.Request, config DomainConfig) (bool, func()) {
|
||||
if !config.AppSecMode.Enabled() {
|
||||
return true, func() {}
|
||||
}
|
||||
|
||||
verdict, release := mw.inspectAppSec(r, config)
|
||||
if verdict == restrict.Allow {
|
||||
return true, release
|
||||
}
|
||||
|
||||
observe := config.AppSecMode == restrict.AppSecObserve
|
||||
if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil {
|
||||
cd.SetMetadata("appsec_verdict", verdict.String())
|
||||
if observe {
|
||||
cd.SetMetadata("appsec_mode", "observe")
|
||||
}
|
||||
}
|
||||
|
||||
if observe {
|
||||
mw.logger.Debugf("AppSec observe: would block %s for %s (%s)", r.RemoteAddr, r.Host, verdict)
|
||||
return true, release
|
||||
}
|
||||
|
||||
mw.markDenied(r, verdict.String())
|
||||
mw.logger.Debugf("AppSec: %s for %s %s", verdict, r.Host, r.RemoteAddr)
|
||||
http.Error(w, "Forbidden", http.StatusForbidden)
|
||||
return false, release
|
||||
}
|
||||
|
||||
// inspectAppSec runs the AppSec call and returns its verdict. Failures come
|
||||
// back as DenyAppSecUnavailable regardless of mode so observe mode still
|
||||
// records that inspection did not happen; the caller decides what blocks. The
|
||||
// returned release is never nil.
|
||||
func (mw *Middleware) inspectAppSec(r *http.Request, config DomainConfig) (restrict.Verdict, func()) {
|
||||
// Mode requested but the proxy has no AppSec endpoint configured. Management
|
||||
// gates this on the cluster capability; a stale mapping can still arrive.
|
||||
if mw.appsec == nil {
|
||||
mw.logger.Debugf("AppSec mode %q requested for %s but no AppSec endpoint is configured", config.AppSecMode, r.Host)
|
||||
return restrict.DenyAppSecUnavailable, func() {}
|
||||
}
|
||||
|
||||
clientIP := mw.resolveClientIP(r)
|
||||
if !clientIP.IsValid() {
|
||||
// The engine requires a client address, and a request whose source we
|
||||
// cannot establish is exactly the kind we must not wave through.
|
||||
mw.logger.Debugf("AppSec: cannot resolve client address for %q", r.RemoteAddr)
|
||||
return restrict.DenyAppSecUnavailable, func() {}
|
||||
}
|
||||
|
||||
req := appsec.Request{
|
||||
HTTP: r,
|
||||
ClientIP: clientIP,
|
||||
RedactBodyFields: config.redactBodyFields,
|
||||
RedactHeaders: config.redactHeaders,
|
||||
RedactCookies: sessionCookieNames,
|
||||
RedactQueryParams: config.redactQueryParams,
|
||||
}
|
||||
if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil {
|
||||
req.TransactionID = cd.GetRequestID()
|
||||
}
|
||||
|
||||
result, err := mw.appsec.Inspect(r.Context(), req)
|
||||
if err != nil {
|
||||
mw.logger.Debugf("AppSec inspection failed for %s: %v", r.Host, err)
|
||||
}
|
||||
// Record when the body went uninspected: headers and URI were still
|
||||
// checked, but an operator reading the log should not read a clean verdict
|
||||
// as "the payload was examined". Oversize is reachable by padding, so its
|
||||
// absence from the log would hide a deliberate opt-out.
|
||||
if result.BodyBypass != "" {
|
||||
if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil {
|
||||
cd.SetMetadata("appsec_body_bypass", result.BodyBypass)
|
||||
}
|
||||
}
|
||||
return result.Verdict, result.Release
|
||||
}
|
||||
|
||||
// credentialFormFields lists the login form fields whose values are redacted
|
||||
// from the mirrored body, so a password or PIN submitted to the proxy's own
|
||||
// login form never reaches the Security Engine.
|
||||
func credentialFormFields(schemes []Scheme) []string {
|
||||
var fields []string
|
||||
for _, s := range schemes {
|
||||
switch s.Type() {
|
||||
case auth.MethodPassword:
|
||||
fields = append(fields, passwordFormId)
|
||||
case auth.MethodPIN:
|
||||
fields = append(fields, pinFormId)
|
||||
}
|
||||
}
|
||||
return fields
|
||||
}
|
||||
|
||||
// credentialHeaders lists the request headers whose values are redacted from
|
||||
// the mirrored request. A header-auth scheme carries a session token the proxy
|
||||
// validates and never forwards upstream, so the engine must not see it either.
|
||||
func credentialHeaders(schemes []Scheme) []string {
|
||||
var names []string
|
||||
for _, s := range schemes {
|
||||
// Structural, not a concrete Header assertion: if the scheme is ever
|
||||
// registered as a pointer, a type assertion would quietly stop matching
|
||||
// and the header would start reaching the engine again.
|
||||
named, ok := s.(interface{ HeaderName() string })
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
if name := named.HeaderName(); name != "" {
|
||||
names = append(names, name)
|
||||
}
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
// resolveClientIP extracts the real client IP from CapturedData, falling back to r.RemoteAddr.
|
||||
func (mw *Middleware) resolveClientIP(r *http.Request) netip.Addr {
|
||||
if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil {
|
||||
@@ -440,18 +281,12 @@ func (mw *Middleware) resolveClientIP(r *http.Request) netip.Addr {
|
||||
return addr.Unmap()
|
||||
}
|
||||
|
||||
// markDenied records the deny reason on the captured data so the access log
|
||||
// attributes the response to the proxy rather than the backend.
|
||||
func (mw *Middleware) markDenied(r *http.Request, reason string) {
|
||||
// blockIPRestriction sets captured data fields for an IP-restriction block event.
|
||||
func (mw *Middleware) blockIPRestriction(r *http.Request, reason string) {
|
||||
if cd := proxy.CapturedDataFromContext(r.Context()); cd != nil {
|
||||
cd.SetOrigin(proxy.OriginAuth)
|
||||
cd.SetAuthMethod(reason)
|
||||
}
|
||||
}
|
||||
|
||||
// blockIPRestriction sets captured data fields for an IP-restriction block event.
|
||||
func (mw *Middleware) blockIPRestriction(r *http.Request, reason string) {
|
||||
mw.markDenied(r, reason)
|
||||
mw.logger.Debugf("IP restriction: %s for %s", reason, r.RemoteAddr)
|
||||
}
|
||||
|
||||
@@ -802,61 +637,45 @@ func wasCredentialSubmitted(r *http.Request, method auth.Method) bool {
|
||||
case auth.MethodPassword:
|
||||
return r.FormValue("password") != ""
|
||||
case auth.MethodOIDC:
|
||||
return r.URL.Query().Get(sessionTokenParam) != ""
|
||||
return r.URL.Query().Get("session_token") != ""
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// DomainSettings is the per-domain configuration AddDomain applies.
|
||||
type DomainSettings struct {
|
||||
Schemes []Scheme
|
||||
// SessionPublicKey is the base64-encoded ed25519 key used to verify session
|
||||
// cookies. Required when Schemes is non-empty.
|
||||
SessionPublicKey string
|
||||
SessionExpiration time.Duration
|
||||
AccountID types.AccountID
|
||||
ServiceID types.ServiceID
|
||||
IPRestrictions *restrict.Filter
|
||||
// Private forces ValidateTunnelPeer enforcement (403 on failure) regardless
|
||||
// of the schemes list.
|
||||
Private bool
|
||||
AppSecMode restrict.AppSecMode
|
||||
}
|
||||
|
||||
// AddDomain registers authentication schemes for the given domain. With schemes
|
||||
// a valid session public key is required.
|
||||
func (mw *Middleware) AddDomain(domain string, settings DomainSettings) error {
|
||||
credentialFields := credentialFormFields(settings.Schemes)
|
||||
config := DomainConfig{
|
||||
AccountID: settings.AccountID,
|
||||
ServiceID: settings.ServiceID,
|
||||
IPRestrictions: settings.IPRestrictions,
|
||||
Private: settings.Private,
|
||||
AppSecMode: settings.AppSecMode,
|
||||
redactBodyFields: credentialFields,
|
||||
redactHeaders: credentialHeaders(settings.Schemes),
|
||||
// A credential can arrive in the query too: r.FormValue merges the URL
|
||||
// query into the form, so "?password=..." authenticates just as a form
|
||||
// post does and must not be mirrored in the clear either.
|
||||
redactQueryParams: append([]string{sessionTokenParam}, credentialFields...),
|
||||
// AddDomain registers authentication schemes for the given domain. With schemes a valid session public key is required.
|
||||
// private=true forces ValidateTunnelPeer enforcement (403 on failure) regardless of the schemes list.
|
||||
func (mw *Middleware) AddDomain(domain string, schemes []Scheme, publicKeyB64 string, expiration time.Duration, accountID types.AccountID, serviceID types.ServiceID, ipRestrictions *restrict.Filter, private bool) error {
|
||||
if len(schemes) == 0 {
|
||||
mw.domainsMux.Lock()
|
||||
defer mw.domainsMux.Unlock()
|
||||
mw.domains[domain] = DomainConfig{
|
||||
AccountID: accountID,
|
||||
ServiceID: serviceID,
|
||||
IPRestrictions: ipRestrictions,
|
||||
Private: private,
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
if len(settings.Schemes) > 0 {
|
||||
pubKeyBytes, err := base64.StdEncoding.DecodeString(settings.SessionPublicKey)
|
||||
if err != nil {
|
||||
return fmt.Errorf("decode session public key for domain %s: %w", domain, err)
|
||||
}
|
||||
if len(pubKeyBytes) != ed25519.PublicKeySize {
|
||||
return fmt.Errorf("invalid session public key size for domain %s: got %d, want %d", domain, len(pubKeyBytes), ed25519.PublicKeySize)
|
||||
}
|
||||
config.Schemes = settings.Schemes
|
||||
config.SessionPublicKey = pubKeyBytes
|
||||
config.SessionExpiration = settings.SessionExpiration
|
||||
pubKeyBytes, err := base64.StdEncoding.DecodeString(publicKeyB64)
|
||||
if err != nil {
|
||||
return fmt.Errorf("decode session public key for domain %s: %w", domain, err)
|
||||
}
|
||||
if len(pubKeyBytes) != ed25519.PublicKeySize {
|
||||
return fmt.Errorf("invalid session public key size for domain %s: got %d, want %d", domain, len(pubKeyBytes), ed25519.PublicKeySize)
|
||||
}
|
||||
|
||||
mw.domainsMux.Lock()
|
||||
defer mw.domainsMux.Unlock()
|
||||
mw.domains[domain] = config
|
||||
mw.domains[domain] = DomainConfig{
|
||||
Schemes: schemes,
|
||||
SessionPublicKey: pubKeyBytes,
|
||||
SessionExpiration: expiration,
|
||||
AccountID: accountID,
|
||||
ServiceID: serviceID,
|
||||
IPRestrictions: ipRestrictions,
|
||||
Private: private,
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -911,10 +730,10 @@ func (mw *Middleware) validateSessionToken(ctx context.Context, host, token stri
|
||||
// parameter removed so it doesn't linger in the browser's address bar or history.
|
||||
func stripSessionTokenParam(u *url.URL) string {
|
||||
q := u.Query()
|
||||
if !q.Has(sessionTokenParam) {
|
||||
if !q.Has("session_token") {
|
||||
return u.RequestURI()
|
||||
}
|
||||
q.Del(sessionTokenParam)
|
||||
q.Del("session_token")
|
||||
clean := *u
|
||||
clean.RawQuery = q.Encode()
|
||||
return clean.RequestURI()
|
||||
|
||||
@@ -64,7 +64,7 @@ func TestAddDomain_ValidKey(t *testing.T) {
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
err := mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour})
|
||||
err := mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false)
|
||||
require.NoError(t, err)
|
||||
|
||||
mw.domainsMux.RLock()
|
||||
@@ -81,7 +81,7 @@ func TestAddDomain_EmptyKey(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
err := mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionExpiration: time.Hour})
|
||||
err := mw.AddDomain("example.com", []Scheme{scheme}, "", time.Hour, "", "", nil, false)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "invalid session public key size")
|
||||
|
||||
@@ -95,7 +95,7 @@ func TestAddDomain_InvalidBase64(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
err := mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: "not-valid-base64!!!", SessionExpiration: time.Hour})
|
||||
err := mw.AddDomain("example.com", []Scheme{scheme}, "not-valid-base64!!!", time.Hour, "", "", nil, false)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "decode session public key")
|
||||
|
||||
@@ -110,7 +110,7 @@ func TestAddDomain_WrongKeySize(t *testing.T) {
|
||||
|
||||
shortKey := base64.StdEncoding.EncodeToString([]byte("tooshort"))
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
err := mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: shortKey, SessionExpiration: time.Hour})
|
||||
err := mw.AddDomain("example.com", []Scheme{scheme}, shortKey, time.Hour, "", "", nil, false)
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "invalid session public key size")
|
||||
|
||||
@@ -123,7 +123,7 @@ func TestAddDomain_WrongKeySize(t *testing.T) {
|
||||
func TestAddDomain_NoSchemes_NoKeyRequired(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
|
||||
err := mw.AddDomain("example.com", DomainSettings{SessionExpiration: time.Hour})
|
||||
err := mw.AddDomain("example.com", nil, "", time.Hour, "", "", nil, false)
|
||||
require.NoError(t, err, "domains with no auth schemes should not require a key")
|
||||
|
||||
mw.domainsMux.RLock()
|
||||
@@ -139,8 +139,8 @@ func TestAddDomain_OverwritesPreviousConfig(t *testing.T) {
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp1.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp2.PublicKey, SessionExpiration: 2 * time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp1.PublicKey, time.Hour, "", "", nil, false))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp2.PublicKey, 2*time.Hour, "", "", nil, false))
|
||||
|
||||
mw.domainsMux.RLock()
|
||||
config := mw.domains["example.com"]
|
||||
@@ -156,7 +156,7 @@ func TestRemoveDomain(t *testing.T) {
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
mw.RemoveDomain("example.com")
|
||||
|
||||
@@ -180,7 +180,7 @@ func TestProtect_UnknownDomainPassesThrough(t *testing.T) {
|
||||
|
||||
func TestProtect_DomainWithNoSchemesPassesThrough(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", nil, "", time.Hour, "", "", nil, false))
|
||||
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
@@ -197,7 +197,7 @@ func TestProtect_UnauthenticatedRequestIsBlocked(t *testing.T) {
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
var backendCalled bool
|
||||
backend := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
@@ -218,7 +218,7 @@ func TestProtect_HostWithPortIsMatched(t *testing.T) {
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
var backendCalled bool
|
||||
backend := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
@@ -239,7 +239,7 @@ func TestProtect_ValidSessionCookiePassesThrough(t *testing.T) {
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
token, err := sessionkey.SignToken(kp.PrivateKey, "test-user", "", "example.com", auth.MethodPIN, nil, nil, time.Hour)
|
||||
require.NoError(t, err)
|
||||
@@ -272,7 +272,7 @@ func TestProtect_SessionCookieGroupsPropagate(t *testing.T) {
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
groups := []string{"engineering", "sre"}
|
||||
token, err := sessionkey.SignToken(kp.PrivateKey, "test-user", "", "example.com", auth.MethodPIN, groups, nil, time.Hour)
|
||||
@@ -337,7 +337,7 @@ func TestProtect_PrivateService_TunnelPeerGroupsPropagate(t *testing.T) {
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
// Private service: no operator schemes — auth gates solely on the tunnel peer.
|
||||
require.NoError(t, mw.AddDomain("agent.example.com", DomainSettings{SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour, AccountID: "acct-1", ServiceID: "svc-1", Private: true}))
|
||||
require.NoError(t, mw.AddDomain("agent.example.com", nil, kp.PublicKey, time.Hour, "acct-1", "svc-1", nil, true))
|
||||
|
||||
cd := proxy.NewCapturedData("")
|
||||
cd.SetClientIP(netip.MustParseAddr("100.90.1.14")) // CGNAT tunnel source
|
||||
@@ -377,7 +377,7 @@ func TestProtect_PrivateService_TunnelPeerDenied(t *testing.T) {
|
||||
}}
|
||||
mw := NewMiddleware(log.StandardLogger(), validator, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
require.NoError(t, mw.AddDomain("agent.example.com", DomainSettings{SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour, AccountID: "acct-1", ServiceID: "svc-1", Private: true}))
|
||||
require.NoError(t, mw.AddDomain("agent.example.com", nil, kp.PublicKey, time.Hour, "acct-1", "svc-1", nil, true))
|
||||
|
||||
cd := proxy.NewCapturedData("")
|
||||
cd.SetClientIP(netip.MustParseAddr("100.90.1.14"))
|
||||
@@ -405,7 +405,7 @@ func TestProtect_ExpiredSessionCookieIsRejected(t *testing.T) {
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
// Sign a token that expired 1 second ago.
|
||||
token, err := sessionkey.SignToken(kp.PrivateKey, "test-user", "", "example.com", auth.MethodPIN, nil, nil, -time.Second)
|
||||
@@ -431,7 +431,7 @@ func TestProtect_WrongDomainCookieIsRejected(t *testing.T) {
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
// Token signed for a different domain audience.
|
||||
token, err := sessionkey.SignToken(kp.PrivateKey, "test-user", "", "other.com", auth.MethodPIN, nil, nil, time.Hour)
|
||||
@@ -458,7 +458,7 @@ func TestProtect_WrongKeyCookieIsRejected(t *testing.T) {
|
||||
kp2 := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp1.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp1.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
// Token signed with a different private key.
|
||||
token, err := sessionkey.SignToken(kp2.PrivateKey, "test-user", "", "example.com", auth.MethodPIN, nil, nil, time.Hour)
|
||||
@@ -495,7 +495,7 @@ func TestProtect_SchemeAuthRedirectsWithCookie(t *testing.T) {
|
||||
return "", "pin", nil
|
||||
},
|
||||
}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
var backendCalled bool
|
||||
backend := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
@@ -548,7 +548,7 @@ func TestProtect_FailedAuthDoesNotSetCookie(t *testing.T) {
|
||||
return "", "pin", nil
|
||||
},
|
||||
}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
@@ -584,7 +584,7 @@ func TestProtect_MultipleSchemes(t *testing.T) {
|
||||
return "", "password", nil
|
||||
},
|
||||
}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{pinScheme, passwordScheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{pinScheme, passwordScheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
var backendCalled bool
|
||||
backend := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
@@ -614,7 +614,7 @@ func TestProtect_InvalidTokenFromSchemeReturns400(t *testing.T) {
|
||||
return "invalid-jwt-token", "", nil
|
||||
},
|
||||
}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
@@ -638,7 +638,7 @@ func TestAddDomain_RandomBytes32NotEd25519(t *testing.T) {
|
||||
key := base64.StdEncoding.EncodeToString(randomBytes)
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
|
||||
err = mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: key, SessionExpiration: time.Hour})
|
||||
err = mw.AddDomain("example.com", []Scheme{scheme}, key, time.Hour, "", "", nil, false)
|
||||
require.NoError(t, err, "any 32-byte key should be accepted at registration time")
|
||||
}
|
||||
|
||||
@@ -647,10 +647,10 @@ func TestAddDomain_InvalidKeyDoesNotCorruptExistingConfig(t *testing.T) {
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
// Attempt to overwrite with an invalid key.
|
||||
err := mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: "bad", SessionExpiration: time.Hour})
|
||||
err := mw.AddDomain("example.com", []Scheme{scheme}, "bad", time.Hour, "", "", nil, false)
|
||||
require.Error(t, err)
|
||||
|
||||
// The original valid config should still be intact.
|
||||
@@ -674,7 +674,7 @@ func TestProtect_FailedPinAuthCapturesAuthMethod(t *testing.T) {
|
||||
return "", "pin", nil
|
||||
},
|
||||
}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
capturedData := proxy.NewCapturedData("")
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
@@ -701,7 +701,7 @@ func TestProtect_FailedPasswordAuthCapturesAuthMethod(t *testing.T) {
|
||||
return "", "password", nil
|
||||
},
|
||||
}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
capturedData := proxy.NewCapturedData("")
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
@@ -728,7 +728,7 @@ func TestProtect_NoCredentialsDoesNotCaptureAuthMethod(t *testing.T) {
|
||||
return "", "pin", nil
|
||||
},
|
||||
}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
capturedData := proxy.NewCapturedData("")
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
@@ -815,7 +815,8 @@ func TestWasCredentialSubmitted(t *testing.T) {
|
||||
func TestCheckIPRestrictions_UnparseableAddress(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
|
||||
err := mw.AddDomain("example.com", DomainSettings{AccountID: "acc1", ServiceID: "svc1", IPRestrictions: restrict.ParseFilter(restrict.FilterConfig{AllowedCIDRs: []string{"10.0.0.0/8"}})})
|
||||
err := mw.AddDomain("example.com", nil, "", 0, "acc1", "svc1",
|
||||
restrict.ParseFilter(restrict.FilterConfig{AllowedCIDRs: []string{"10.0.0.0/8"}}), false)
|
||||
require.NoError(t, err)
|
||||
|
||||
handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
@@ -850,7 +851,8 @@ func TestCheckIPRestrictions_UsesCapturedDataClientIP(t *testing.T) {
|
||||
// trusted proxies), checkIPRestrictions should use that IP, not RemoteAddr.
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
|
||||
err := mw.AddDomain("example.com", DomainSettings{AccountID: "acc1", ServiceID: "svc1", IPRestrictions: restrict.ParseFilter(restrict.FilterConfig{AllowedCIDRs: []string{"203.0.113.0/24"}})})
|
||||
err := mw.AddDomain("example.com", nil, "", 0, "acc1", "svc1",
|
||||
restrict.ParseFilter(restrict.FilterConfig{AllowedCIDRs: []string{"203.0.113.0/24"}}), false)
|
||||
require.NoError(t, err)
|
||||
|
||||
handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
@@ -890,7 +892,8 @@ func TestCheckIPRestrictions_NilGeoWithCountryRules(t *testing.T) {
|
||||
// Geo is nil, country restrictions are configured: must deny (fail-close).
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
|
||||
err := mw.AddDomain("example.com", DomainSettings{AccountID: "acc1", ServiceID: "svc1", IPRestrictions: restrict.ParseFilter(restrict.FilterConfig{AllowedCountries: []string{"US"}})})
|
||||
err := mw.AddDomain("example.com", nil, "", 0, "acc1", "svc1",
|
||||
restrict.ParseFilter(restrict.FilterConfig{AllowedCountries: []string{"US"}}), false)
|
||||
require.NoError(t, err)
|
||||
|
||||
handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
@@ -913,10 +916,11 @@ func TestCheckIPRestrictions_NilGeoWithCountryRules(t *testing.T) {
|
||||
func TestCheckIPRestrictions_OverlayOriginSkipsCountryRules(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
|
||||
err := mw.AddDomain("example.com", DomainSettings{AccountID: "acc1", ServiceID: "svc1", IPRestrictions: restrict.ParseFilter(restrict.FilterConfig{
|
||||
AllowedCIDRs: []string{"100.64.0.0/10"},
|
||||
AllowedCountries: []string{"US"},
|
||||
})})
|
||||
err := mw.AddDomain("example.com", nil, "", 0, "acc1", "svc1",
|
||||
restrict.ParseFilter(restrict.FilterConfig{
|
||||
AllowedCIDRs: []string{"100.64.0.0/10"},
|
||||
AllowedCountries: []string{"US"},
|
||||
}), false)
|
||||
require.NoError(t, err)
|
||||
|
||||
handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
@@ -949,7 +953,8 @@ func TestCheckIPRestrictions_OverlayOriginSkipsCountryRules(t *testing.T) {
|
||||
func TestCheckIPRestrictions_OverlayOriginRespectsCIDR(t *testing.T) {
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
|
||||
err := mw.AddDomain("example.com", DomainSettings{AccountID: "acc1", ServiceID: "svc1", IPRestrictions: restrict.ParseFilter(restrict.FilterConfig{AllowedCIDRs: []string{"100.64.0.0/16"}})})
|
||||
err := mw.AddDomain("example.com", nil, "", 0, "acc1", "svc1",
|
||||
restrict.ParseFilter(restrict.FilterConfig{AllowedCIDRs: []string{"100.64.0.0/16"}}), false)
|
||||
require.NoError(t, err)
|
||||
|
||||
handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
@@ -977,7 +982,7 @@ func TestProtect_OIDCOnlyRedirectsDirectly(t *testing.T) {
|
||||
return "", oidcURL, nil
|
||||
},
|
||||
}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
@@ -1006,7 +1011,7 @@ func TestProtect_OIDCWithOtherMethodShowsLoginPage(t *testing.T) {
|
||||
return "", "pin", nil
|
||||
},
|
||||
}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{oidcScheme, pinScheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{oidcScheme, pinScheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
@@ -1050,7 +1055,7 @@ func TestProtect_HeaderAuth_ForwardsOnSuccess(t *testing.T) {
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
hdr := newHeaderSchemeWithToken(t, kp, "X-API-Key", "secret-key")
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{hdr}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour, AccountID: "acc1", ServiceID: "svc1"}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false))
|
||||
|
||||
var backendCalled bool
|
||||
capturedData := proxy.NewCapturedData("")
|
||||
@@ -1093,7 +1098,7 @@ func TestProtect_HeaderAuth_MissingHeaderFallsThrough(t *testing.T) {
|
||||
hdr := newHeaderSchemeWithToken(t, kp, "X-API-Key", "secret-key")
|
||||
// Also add a PIN scheme so we can verify fallthrough behavior.
|
||||
pinScheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{hdr, pinScheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour, AccountID: "acc1", ServiceID: "svc1"}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr, pinScheme}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false))
|
||||
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
@@ -1113,7 +1118,7 @@ func TestProtect_HeaderAuth_WrongValueReturns401(t *testing.T) {
|
||||
return &proto.AuthenticateResponse{Success: false}, nil
|
||||
}}
|
||||
hdr := NewHeader(mock, "svc1", "acc1", "X-API-Key")
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{hdr}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour, AccountID: "acc1", ServiceID: "svc1"}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false))
|
||||
|
||||
capturedData := proxy.NewCapturedData("")
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
@@ -1136,7 +1141,7 @@ func TestProtect_HeaderAuth_InfraErrorReturns502(t *testing.T) {
|
||||
return nil, errors.New("gRPC unavailable")
|
||||
}}
|
||||
hdr := NewHeader(mock, "svc1", "acc1", "X-API-Key")
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{hdr}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour, AccountID: "acc1", ServiceID: "svc1"}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false))
|
||||
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
@@ -1153,7 +1158,7 @@ func TestProtect_HeaderAuth_SubsequentRequestUsesSessionCookie(t *testing.T) {
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
hdr := newHeaderSchemeWithToken(t, kp, "X-API-Key", "secret-key")
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{hdr}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour, AccountID: "acc1", ServiceID: "svc1"}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false))
|
||||
|
||||
handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
@@ -1213,7 +1218,7 @@ func TestProtect_HeaderAuth_MultipleValuesSameHeader(t *testing.T) {
|
||||
|
||||
// Single Header scheme (as if one entry existed), but the mock checks both values.
|
||||
hdr := NewHeader(mock, "svc1", "acc1", "Authorization")
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{hdr}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour, AccountID: "acc1", ServiceID: "svc1"}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false))
|
||||
|
||||
var backendCalled bool
|
||||
handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
@@ -1271,7 +1276,7 @@ func TestProtect_OIDCOnPlainHTTP_BlockedWith400(t *testing.T) {
|
||||
return "", "https://idp.example.com/authorize", nil
|
||||
},
|
||||
}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
@@ -1295,7 +1300,7 @@ func TestProtect_OIDCOverTLS_NotBlocked(t *testing.T) {
|
||||
return "", "https://idp.example.com/authorize", nil
|
||||
},
|
||||
}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
@@ -1315,7 +1320,7 @@ func TestProtect_NonOIDCSchemes_PlainHTTP_NotBlocked(t *testing.T) {
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
@@ -1345,7 +1350,7 @@ func TestProtect_TunnelPeerFastPath_RequiresInboundMarker(t *testing.T) {
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
@@ -1380,7 +1385,7 @@ func TestProtect_TunnelPeerFastPath_TakesPathWithInboundMarker(t *testing.T) {
|
||||
kp := generateTestKeyPair(t)
|
||||
|
||||
scheme := &stubScheme{method: auth.MethodPIN, promptID: "pin"}
|
||||
require.NoError(t, mw.AddDomain("example.com", DomainSettings{Schemes: []Scheme{scheme}, SessionPublicKey: kp.PublicKey, SessionExpiration: time.Hour}))
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{scheme}, kp.PublicKey, time.Hour, "", "", nil, false))
|
||||
|
||||
handler := mw.Protect(newPassthroughHandler())
|
||||
|
||||
|
||||
@@ -13,10 +13,6 @@ import (
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
// sessionTokenParam is the query parameter the management server uses to hand
|
||||
// the minted session token back to the proxy after an OIDC login.
|
||||
const sessionTokenParam = "session_token"
|
||||
|
||||
type urlGenerator interface {
|
||||
GetOIDCURL(context.Context, *proto.GetOIDCURLRequest, ...grpc.CallOption) (*proto.GetOIDCURLResponse, error)
|
||||
}
|
||||
@@ -47,7 +43,7 @@ func (o OIDC) Authenticate(r *http.Request) (string, string, error) {
|
||||
// Check for the session_token query param (from OIDC redirects).
|
||||
// The management server passes the token in the URL because it cannot set
|
||||
// cookies for the proxy's domain (cookies are domain-scoped per RFC 6265).
|
||||
if token := r.URL.Query().Get(sessionTokenParam); token != "" {
|
||||
if token := r.URL.Query().Get("session_token"); token != "" {
|
||||
return token, "", nil
|
||||
}
|
||||
|
||||
|
||||
@@ -44,7 +44,7 @@ func (s *stubSessionValidator) ValidateTunnelPeer(_ context.Context, in *proto.V
|
||||
func newTunnelMiddleware(t *testing.T, validator SessionValidator) *Middleware {
|
||||
t.Helper()
|
||||
mw := NewMiddleware(log.New(), validator, nil)
|
||||
require.NoError(t, mw.AddDomain("svc.example", DomainSettings{AccountID: "acct-1", ServiceID: "svc-1"}))
|
||||
require.NoError(t, mw.AddDomain("svc.example", nil, "", 0, "acct-1", "svc-1", nil, false))
|
||||
return mw
|
||||
}
|
||||
|
||||
@@ -235,8 +235,8 @@ func TestForwardWithTunnelPeer_RoutesAccountIDIntoCacheKey(t *testing.T) {
|
||||
}
|
||||
mw := NewMiddleware(log.New(), validator, nil)
|
||||
|
||||
require.NoError(t, mw.AddDomain("svc-a.example", DomainSettings{AccountID: "acct-a", ServiceID: "svc-a"}))
|
||||
require.NoError(t, mw.AddDomain("svc-b.example", DomainSettings{AccountID: "acct-b", ServiceID: "svc-b"}))
|
||||
require.NoError(t, mw.AddDomain("svc-a.example", nil, "", 0, "acct-a", "svc-a", nil, false))
|
||||
require.NoError(t, mw.AddDomain("svc-b.example", nil, "", 0, "acct-b", "svc-b", nil, false))
|
||||
|
||||
// The fast-path requires the inbound-listener marker on the context.
|
||||
// The peerstore lookup itself is account-agnostic at this level
|
||||
@@ -299,7 +299,7 @@ func TestForwardWithTunnelPeer_LocalLookupShortCircuitDoesNotPopulateCache(t *te
|
||||
|
||||
func TestPrivateService_FailsClosedOnTunnelPeerFailure(t *testing.T) {
|
||||
mw := NewMiddleware(log.New(), nil, nil)
|
||||
require.NoError(t, mw.AddDomain("private.svc", DomainSettings{AccountID: "acct-1", ServiceID: "svc-1", Private: true}))
|
||||
require.NoError(t, mw.AddDomain("private.svc", nil, "", 0, "acct-1", "svc-1", nil, true))
|
||||
|
||||
called := false
|
||||
handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
@@ -328,7 +328,7 @@ func TestPrivateService_ForwardsOnTunnelPeerSuccess(t *testing.T) {
|
||||
},
|
||||
}
|
||||
mw := NewMiddleware(log.New(), validator, nil)
|
||||
require.NoError(t, mw.AddDomain("private.svc", DomainSettings{AccountID: "acct-1", ServiceID: "svc-1", Private: true}))
|
||||
require.NoError(t, mw.AddDomain("private.svc", nil, "", 0, "acct-1", "svc-1", nil, true))
|
||||
|
||||
called := false
|
||||
handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
|
||||
@@ -21,8 +21,6 @@ import (
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/netutil"
|
||||
)
|
||||
|
||||
// MaxRoutingScanBytes bounds how far ScanRoutingFields will read into a
|
||||
@@ -36,6 +34,7 @@ const MaxRoutingScanBytes int64 = 32 << 20
|
||||
// metadata key by the chain when a request body is not surfaced.
|
||||
const (
|
||||
BypassUpgradeHeader = "upgrade_header"
|
||||
BypassConnectionUpgrd = "connection_upgrade"
|
||||
BypassContentType = "content_type_not_allowed"
|
||||
BypassBudget = "capture_budget_exhausted"
|
||||
BypassNoConfig = "no_capture_config"
|
||||
@@ -126,13 +125,12 @@ func CaptureRequest(r *http.Request, cfg *Config, b Budget) (body []byte, trunca
|
||||
if cfg.MaxRequestBytes <= 0 {
|
||||
return nil, false, 0, BypassCapZero, release, nil
|
||||
}
|
||||
// The predicate has to be the forwarder's own: a looser one (either header
|
||||
// on its own) skips capture for requests the forwarder still delivers to
|
||||
// the upstream with their body intact, which hides them from every
|
||||
// deny-capable middleware in the chain.
|
||||
if netutil.IsUpgradeRequest(r.Header) {
|
||||
if r.Header.Get("Upgrade") != "" {
|
||||
return nil, false, 0, BypassUpgradeHeader, release, nil
|
||||
}
|
||||
if strings.EqualFold(r.Header.Get("Connection"), "upgrade") {
|
||||
return nil, false, 0, BypassConnectionUpgrd, release, nil
|
||||
}
|
||||
if !contentTypeAllowed(r.Header.Get("Content-Type"), cfg.ContentTypes) {
|
||||
return nil, false, 0, BypassContentType, release, nil
|
||||
}
|
||||
|
||||
@@ -1,23 +0,0 @@
|
||||
package netutil
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
|
||||
"golang.org/x/net/http/httpguts"
|
||||
)
|
||||
|
||||
// IsUpgradeRequest reports whether r is a protocol-upgrade request, using the
|
||||
// same predicate httputil.ReverseProxy applies when it decides to hand the
|
||||
// connection over instead of proxying normally.
|
||||
//
|
||||
// Matching the forwarder exactly matters for anything that inspects a request
|
||||
// before it is proxied: a looser test (an Upgrade header on its own, say) marks
|
||||
// a request as an upgrade and skips inspection, while the forwarder still
|
||||
// delivers it to the backend as an ordinary request with its body intact. That
|
||||
// gap is a body-inspection bypass reachable by adding one header.
|
||||
func IsUpgradeRequest(h http.Header) bool {
|
||||
if !httpguts.HeaderValuesContainsToken(h["Connection"], "Upgrade") {
|
||||
return false
|
||||
}
|
||||
return h.Get("Upgrade") != ""
|
||||
}
|
||||
@@ -50,59 +50,6 @@ const (
|
||||
CrowdSecObserve CrowdSecMode = "observe"
|
||||
)
|
||||
|
||||
// AllowMatch controls how the configured allowlists (CIDR, country) combine.
|
||||
// Blocklists are always a separate hard-deny gate and are unaffected by it.
|
||||
type AllowMatch string
|
||||
|
||||
const (
|
||||
// AllowMatchAll requires the address to match every configured allowlist
|
||||
// (AND). This is the default and preserves the historical behavior.
|
||||
AllowMatchAll AllowMatch = "all"
|
||||
// AllowMatchAny requires the address to match at least one configured
|
||||
// allowlist (OR), e.g. "allowed country OR allowed CIDR".
|
||||
AllowMatchAny AllowMatch = "any"
|
||||
)
|
||||
|
||||
// normalizeAllowMatch maps unknown or empty values to the restrictive default
|
||||
// (AllowMatchAll) so an unrecognized mode never loosens access.
|
||||
func normalizeAllowMatch(m AllowMatch) AllowMatch {
|
||||
if m == AllowMatchAny {
|
||||
return AllowMatchAny
|
||||
}
|
||||
return AllowMatchAll
|
||||
}
|
||||
|
||||
// AppSecMode is the per-service CrowdSec AppSec (WAF) enforcement mode.
|
||||
type AppSecMode string
|
||||
|
||||
const (
|
||||
// AppSecOff disables request inspection.
|
||||
AppSecOff AppSecMode = ""
|
||||
// AppSecEnforce blocks requests the engine flags, and fails closed when the
|
||||
// engine is unreachable.
|
||||
AppSecEnforce AppSecMode = "enforce"
|
||||
// AppSecObserve records the verdict without blocking.
|
||||
AppSecObserve AppSecMode = "observe"
|
||||
)
|
||||
|
||||
// ParseAppSecMode maps a wire value to an AppSecMode. Unrecognized values map
|
||||
// to AppSecOff so a typo never turns inspection into an unintended block.
|
||||
func ParseAppSecMode(s string) AppSecMode {
|
||||
switch AppSecMode(s) {
|
||||
case AppSecEnforce:
|
||||
return AppSecEnforce
|
||||
case AppSecObserve:
|
||||
return AppSecObserve
|
||||
default:
|
||||
return AppSecOff
|
||||
}
|
||||
}
|
||||
|
||||
// Enabled reports whether the mode asks for request inspection.
|
||||
func (m AppSecMode) Enabled() bool {
|
||||
return m == AppSecEnforce || m == AppSecObserve
|
||||
}
|
||||
|
||||
// Filter evaluates IP restrictions. CIDR checks are performed first
|
||||
// (cheap), followed by country lookups (more expensive) only when needed.
|
||||
type Filter struct {
|
||||
@@ -112,9 +59,6 @@ type Filter struct {
|
||||
BlockedCountries []string
|
||||
CrowdSec CrowdSecChecker
|
||||
CrowdSecMode CrowdSecMode
|
||||
// AllowMatch controls how the allowlists combine (AND vs OR). Empty means
|
||||
// AllowMatchAll.
|
||||
AllowMatch AllowMatch
|
||||
}
|
||||
|
||||
// FilterConfig holds the raw configuration for building a Filter.
|
||||
@@ -125,7 +69,6 @@ type FilterConfig struct {
|
||||
BlockedCountries []string
|
||||
CrowdSec CrowdSecChecker
|
||||
CrowdSecMode CrowdSecMode
|
||||
AllowMatch AllowMatch
|
||||
Logger *log.Entry
|
||||
}
|
||||
|
||||
@@ -146,7 +89,6 @@ func ParseFilter(cfg FilterConfig) *Filter {
|
||||
f := &Filter{
|
||||
AllowedCountries: normalizeCountryCodes(cfg.AllowedCountries),
|
||||
BlockedCountries: normalizeCountryCodes(cfg.BlockedCountries),
|
||||
AllowMatch: normalizeAllowMatch(cfg.AllowMatch),
|
||||
}
|
||||
if hasCS {
|
||||
f.CrowdSec = cfg.CrowdSec
|
||||
@@ -204,13 +146,6 @@ const (
|
||||
// DenyCrowdSecUnavailable indicates enforce mode but the bouncer has not
|
||||
// completed its initial sync.
|
||||
DenyCrowdSecUnavailable
|
||||
// DenyAppSecBan indicates a CrowdSec AppSec "ban" remediation.
|
||||
DenyAppSecBan
|
||||
// DenyAppSecCaptcha indicates a CrowdSec AppSec "captcha" remediation.
|
||||
DenyAppSecCaptcha
|
||||
// DenyAppSecUnavailable indicates enforce mode but the AppSec engine could
|
||||
// not produce a verdict (unreachable, timed out, or it rejected the call).
|
||||
DenyAppSecUnavailable
|
||||
)
|
||||
|
||||
// String returns the deny reason string matching the HTTP auth mechanism names.
|
||||
@@ -232,12 +167,6 @@ func (v Verdict) String() string {
|
||||
return "crowdsec_throttle"
|
||||
case DenyCrowdSecUnavailable:
|
||||
return "crowdsec_unavailable"
|
||||
case DenyAppSecBan:
|
||||
return "appsec_ban"
|
||||
case DenyAppSecCaptcha:
|
||||
return "appsec_captcha"
|
||||
case DenyAppSecUnavailable:
|
||||
return "appsec_unavailable"
|
||||
default:
|
||||
return "unknown"
|
||||
}
|
||||
@@ -253,16 +182,6 @@ func (v Verdict) IsCrowdSec() bool {
|
||||
}
|
||||
}
|
||||
|
||||
// IsAppSec returns true when the verdict originates from an AppSec inspection.
|
||||
func (v Verdict) IsAppSec() bool {
|
||||
switch v {
|
||||
case DenyAppSecBan, DenyAppSecCaptcha, DenyAppSecUnavailable:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
// IsObserveOnly returns true when v is a CrowdSec verdict and the filter is in
|
||||
// observe mode. Callers should log the verdict but not block the request.
|
||||
func (f *Filter) IsObserveOnly(v Verdict) bool {
|
||||
@@ -297,10 +216,6 @@ func (f *Filter) Check(addr netip.Addr, geo GeoResolver) Verdict {
|
||||
// IPv4 CIDR rules match regardless of how the address was received.
|
||||
addr = addr.Unmap()
|
||||
|
||||
if f.AllowMatch == AllowMatchAny {
|
||||
return f.checkAny(addr, geo)
|
||||
}
|
||||
|
||||
if v := f.checkCIDR(addr); v != Allow {
|
||||
return v
|
||||
}
|
||||
@@ -310,68 +225,6 @@ func (f *Filter) Check(addr netip.Addr, geo GeoResolver) Verdict {
|
||||
return f.checkCrowdSec(addr)
|
||||
}
|
||||
|
||||
// checkAny evaluates the filter with OR semantics across allowlists: the
|
||||
// address is admitted if it matches any configured allowlist (CIDR or country).
|
||||
// Blocklists remain a hard-deny gate evaluated first and are independent of the
|
||||
// allow-combine mode, so a blocklist match (or unverifiable country block) still
|
||||
// denies. CrowdSec runs last, as in the default path.
|
||||
//
|
||||
// The country is resolved at most once and shared by both the blocklist and the
|
||||
// allowlist, matching what the all-mode path does. Splitting the two checks into
|
||||
// separate helpers cost a second geo lookup per connection whenever both country
|
||||
// lists were configured.
|
||||
func (f *Filter) checkAny(addr netip.Addr, geo GeoResolver) Verdict {
|
||||
for _, prefix := range f.BlockedCIDRs {
|
||||
if prefix.Contains(addr) {
|
||||
return DenyCIDR
|
||||
}
|
||||
}
|
||||
|
||||
cidrActive := len(f.AllowedCIDRs) > 0
|
||||
cidrAllowed := false
|
||||
if cidrActive {
|
||||
for _, prefix := range f.AllowedCIDRs {
|
||||
if prefix.Contains(addr) {
|
||||
cidrAllowed = true
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
countryActive := len(f.AllowedCountries) > 0
|
||||
// The blocklist is a hard gate, so it needs the country even when a CIDR
|
||||
// allowlist already admitted the address. The allowlist needs it only when
|
||||
// the CIDR list did not admit it, which is why a matching allowed CIDR
|
||||
// still skips the lookup when no country blocklist is configured.
|
||||
needCountry := len(f.BlockedCountries) > 0 || (countryActive && !cidrAllowed)
|
||||
|
||||
country := ""
|
||||
if needCountry {
|
||||
if geo == nil || !geo.Available() {
|
||||
return DenyGeoUnavailable
|
||||
}
|
||||
country = geo.LookupAddr(addr).CountryCode
|
||||
}
|
||||
|
||||
if country != "" && slices.Contains(f.BlockedCountries, country) {
|
||||
return DenyCountry
|
||||
}
|
||||
|
||||
allowed := (!cidrActive && !countryActive) ||
|
||||
cidrAllowed ||
|
||||
(countryActive && country != "" && slices.Contains(f.AllowedCountries, country))
|
||||
if !allowed {
|
||||
// Both allowlists missing is reported against the CIDR list, the one
|
||||
// checked first, so the reason stays stable for existing access logs.
|
||||
if cidrActive {
|
||||
return DenyCIDR
|
||||
}
|
||||
return DenyCountry
|
||||
}
|
||||
|
||||
return f.checkCrowdSec(addr)
|
||||
}
|
||||
|
||||
func (f *Filter) checkCIDR(addr netip.Addr) Verdict {
|
||||
if len(f.AllowedCIDRs) > 0 {
|
||||
allowed := false
|
||||
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user