mirror of
https://github.com/netbirdio/netbird.git
synced 2026-08-12 18:51:28 +02:00
Compare commits
17 Commits
feature/pr
...
feature/io
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
adf62a8080 | ||
|
|
77e5ac776b | ||
|
|
12546e231c | ||
|
|
052cf5a748 | ||
|
|
95a458801c | ||
|
|
14f9f8ce22 | ||
|
|
f805c149d9 | ||
|
|
223153993c | ||
|
|
99048e2bf2 | ||
|
|
27b2d3f351 | ||
|
|
ebfdf7d7b8 | ||
|
|
e8671a811d | ||
|
|
1ca26d8faa | ||
|
|
081f2d153b | ||
|
|
90bb5e7c0f | ||
|
|
1381dcf919 | ||
|
|
7ea0882975 |
@@ -43,19 +43,17 @@ archives:
|
||||
- netbird-ui-gtk3
|
||||
|
||||
nfpms:
|
||||
# Same package_name as the GTK4 packages -- the two are mutually-exclusive
|
||||
# alternatives served from separate repo paths (see uploads below); a given
|
||||
# distro points at exactly one of them. The file names must still differ:
|
||||
# the Debian pool is shared storage keyed by file name, so a default-named
|
||||
# gtk3 .deb would overwrite the stable one.
|
||||
# Mutually-exclusive alternative to the GTK4 netbird-ui package -- both
|
||||
# ship the same /usr/bin/netbird-ui from the shared stable/yum repos, so
|
||||
# this one carries its own name and conflicts with the GTK4 package.
|
||||
- maintainer: Netbird <dev@netbird.io>
|
||||
description: Netbird client UI.
|
||||
homepage: https://netbird.io/
|
||||
license: BSD-3-Clause
|
||||
vendor: NetBird
|
||||
id: netbird_ui_deb_gtk3
|
||||
package_name: netbird-ui
|
||||
file_name_template: "{{ .PackageName }}-gtk3_{{ .Version }}_{{ .Os }}_{{ .Arch }}"
|
||||
package_name: netbird-ui-gtk3
|
||||
file_name_template: "{{ .PackageName }}_{{ .Version }}_{{ .Os }}_{{ .Arch }}"
|
||||
builds:
|
||||
- netbird-ui-gtk3
|
||||
formats:
|
||||
@@ -67,6 +65,10 @@ nfpms:
|
||||
dst: /usr/share/applications/org.wails.netbird.desktop
|
||||
- src: client/ui/build/appicon.png
|
||||
dst: /usr/share/pixmaps/netbird.png
|
||||
conflicts:
|
||||
- netbird-ui
|
||||
replaces:
|
||||
- netbird-ui
|
||||
dependencies:
|
||||
- netbird (>= 0.75.0)
|
||||
- libgtk-3-0
|
||||
@@ -79,8 +81,8 @@ nfpms:
|
||||
license: BSD-3-Clause
|
||||
vendor: NetBird
|
||||
id: netbird_ui_rpm_gtk3
|
||||
package_name: netbird-ui
|
||||
file_name_template: "{{ .PackageName }}-gtk3_{{ .Version }}_{{ .Os }}_{{ .Arch }}"
|
||||
package_name: netbird-ui-gtk3
|
||||
file_name_template: "{{ .PackageName }}_{{ .Version }}_{{ .Os }}_{{ .Arch }}"
|
||||
builds:
|
||||
- netbird-ui-gtk3
|
||||
formats:
|
||||
@@ -92,6 +94,10 @@ nfpms:
|
||||
dst: /usr/share/applications/org.wails.netbird.desktop
|
||||
- src: client/ui/build/appicon.png
|
||||
dst: /usr/share/pixmaps/netbird.png
|
||||
# No `replaces` here: nfpm maps it to rpm Obsoletes, which would make
|
||||
# dnf swap installed GTK4 netbird-ui packages for this one on upgrade.
|
||||
conflicts:
|
||||
- netbird-ui
|
||||
dependencies:
|
||||
- netbird >= 0.75.0
|
||||
- (gtk3 or libgtk-3-0)
|
||||
@@ -111,32 +117,20 @@ changelog:
|
||||
disable: true
|
||||
|
||||
uploads:
|
||||
# The gtk3 packages reuse the netbird-ui package name, so they live in
|
||||
# dedicated repo paths (deb distribution `gtk3`, yum path `yum-gtk3`) that
|
||||
# legacy distros point their repo config at.
|
||||
#
|
||||
# GoReleaser derives the credential env var from the upload name, so these
|
||||
# would look for UPLOAD_DEBIAN-GTK3_SECRET / UPLOAD_YUM-GTK3_SECRET. The
|
||||
# release workflow only exports UPLOAD_DEBIAN_SECRET / UPLOAD_YUM_SECRET, and
|
||||
# a missing secret is a silent skip rather than a failure -- the packages
|
||||
# reached the GitHub release but never the package repositories. Point
|
||||
# `password` at the exported vars so both uploads authenticate.
|
||||
- name: debian-gtk3
|
||||
- name: debian
|
||||
skip: "{{ .Env.SKIP_PUBLISH }}"
|
||||
ids:
|
||||
- netbird_ui_deb_gtk3
|
||||
mode: archive
|
||||
target: https://pkgs.wiretrustee.com/debian/pool/{{ .ArtifactName }};deb.distribution=gtk3;deb.component=main;deb.architecture={{ if .Arm }}armhf{{ else }}{{ .Arch }}{{ end }};deb.package=
|
||||
target: https://pkgs.wiretrustee.com/debian/pool/{{ .ArtifactName }};deb.distribution=stable;deb.component=main;deb.architecture={{ if .Arm }}armhf{{ else }}{{ .Arch }}{{ end }};deb.package=
|
||||
username: dev@wiretrustee.com
|
||||
password: "{{ .Env.UPLOAD_DEBIAN_SECRET }}"
|
||||
method: PUT
|
||||
|
||||
- name: yum-gtk3
|
||||
- name: yum
|
||||
skip: "{{ .Env.SKIP_PUBLISH }}"
|
||||
ids:
|
||||
- netbird_ui_rpm_gtk3
|
||||
mode: archive
|
||||
target: https://pkgs.wiretrustee.com/yum-gtk3/{{ .Arch }}{{ if .Arm }}{{ .Arm }}{{ end }}
|
||||
target: https://pkgs.wiretrustee.com/yum/{{ .Arch }}{{ if .Arm }}{{ .Arm }}{{ end }}
|
||||
username: dev@wiretrustee.com
|
||||
password: "{{ .Env.UPLOAD_YUM_SECRET }}"
|
||||
method: PUT
|
||||
|
||||
@@ -112,6 +112,7 @@ aligns with our security standards and design expectations.
|
||||
- [Test suite](#test-suite)
|
||||
- [Checklist before submitting a PR](#checklist-before-submitting-a-pr)
|
||||
- [When we close a PR](#when-we-close-a-pr)
|
||||
- [Translations](#translations)
|
||||
- [Other project repositories](#other-project-repositories)
|
||||
- [Contributor License Agreement](#contributor-license-agreement)
|
||||
|
||||
@@ -612,6 +613,17 @@ A closed PR is not a rejected idea. Take it back to the
|
||||
[discussion](https://github.com/netbirdio/netbird/discussions), settle the
|
||||
approach, and reopen the work from there.
|
||||
|
||||
## Translations
|
||||
|
||||
Desktop UI translations are not contributed through pull requests. Translate on
|
||||
[Crowdin](https://crowdin.com/project/netbird) instead: no ticket needed, just
|
||||
join the project and pick your language. Crowdin syncs with this repository and
|
||||
opens the service PRs itself, so hand-edited locale files would conflict with
|
||||
the next sync. Style, terminology, and review guidance live in
|
||||
[client/ui/i18n/TRANSLATING.md](client/ui/i18n/TRANSLATING.md). To request a
|
||||
language the project does not offer yet, ask on the Crowdin project page or in
|
||||
a [discussion](https://github.com/netbirdio/netbird/discussions).
|
||||
|
||||
## Other project repositories
|
||||
|
||||
NetBird project is composed of 3 main repositories:
|
||||
|
||||
@@ -22,6 +22,7 @@ import (
|
||||
"github.com/netbirdio/netbird/client/internal/listener"
|
||||
"github.com/netbirdio/netbird/client/internal/peer"
|
||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||
"github.com/netbirdio/netbird/client/ssh/jwtcache"
|
||||
"github.com/netbirdio/netbird/client/system"
|
||||
"github.com/netbirdio/netbird/formatter"
|
||||
"github.com/netbirdio/netbird/route"
|
||||
@@ -92,7 +93,14 @@ type Client struct {
|
||||
|
||||
stateMu sync.RWMutex
|
||||
connectClient *internal.ConnectClient
|
||||
config *profilemanager.Config
|
||||
// config holds the active configuration once Run has loaded it. Consumed by
|
||||
// the in-app SSH client for the NetBird SSH key and the OAuth flow.
|
||||
config *profilemanager.Config
|
||||
|
||||
// sshJWTCache keeps the SSH JWT token between reconnects so the user is not
|
||||
// forced through the browser OAuth flow on every session. Lives on Client
|
||||
// (not SSHClient) because the app creates a new SSHClient per session.
|
||||
sshJWTCache *jwtcache.Cache
|
||||
}
|
||||
|
||||
// NewClient instantiate a new Client
|
||||
@@ -109,6 +117,7 @@ func NewClient(cfgFile, stateFile, cacheDir, logFilePath, deviceName string, osV
|
||||
ctxCancelLock: &sync.Mutex{},
|
||||
networkChangeListener: networkChangeListener,
|
||||
dnsManager: dnsManager,
|
||||
sshJWTCache: jwtcache.New(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -183,6 +192,7 @@ func (c *Client) Run(fd int32, interfaceName string, envList *EnvList) error {
|
||||
ctx = internal.CtxInitState(ctx)
|
||||
c.onHostDnsFn = func([]string) {}
|
||||
cfg.WgIface = interfaceName
|
||||
c.config = cfg
|
||||
|
||||
connectClient := internal.NewConnectClient(ctx, cfg, c.recorder)
|
||||
c.setState(cfg, connectClient)
|
||||
@@ -708,6 +718,13 @@ func (c *Client) stateSnapshot() (*profilemanager.Config, *internal.ConnectClien
|
||||
return c.config, c.connectClient
|
||||
}
|
||||
|
||||
// sshState returns the active config and the running connect client for the
|
||||
// in-app SSH client. Both are nil until Run has loaded the config and started
|
||||
// the tunnel.
|
||||
func (c *Client) sshState() (*profilemanager.Config, *internal.ConnectClient) {
|
||||
return c.stateSnapshot()
|
||||
}
|
||||
|
||||
func formatDuration(d time.Duration) string {
|
||||
ds := d.String()
|
||||
dotIndex := strings.Index(ds, ".")
|
||||
|
||||
512
client/ios/NetBirdSDK/ssh_client.go
Normal file
512
client/ios/NetBirdSDK/ssh_client.go
Normal file
@@ -0,0 +1,512 @@
|
||||
//go:build ios
|
||||
|
||||
package NetBirdSDK
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"io"
|
||||
"net"
|
||||
"strconv"
|
||||
"strings"
|
||||
"sync"
|
||||
"syscall"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
gossh "golang.org/x/crypto/ssh"
|
||||
|
||||
"github.com/netbirdio/netbird/client/internal"
|
||||
"github.com/netbirdio/netbird/client/internal/auth"
|
||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||
nbssh "github.com/netbirdio/netbird/client/ssh"
|
||||
"github.com/netbirdio/netbird/client/ssh/detection"
|
||||
"github.com/netbirdio/netbird/client/ssh/jwtcache"
|
||||
)
|
||||
|
||||
const (
|
||||
sshDialTimeout = 30 * time.Second
|
||||
sshDetectionTimeout = 5 * time.Second
|
||||
)
|
||||
|
||||
// SSHTerminalListener receives SSH session events. It is implemented in Swift.
|
||||
//
|
||||
// All callbacks are invoked from goroutines and may run concurrently with each
|
||||
// other; the implementation must be safe to call from any thread.
|
||||
type SSHTerminalListener interface {
|
||||
OnConnected()
|
||||
OnData(data []byte)
|
||||
OnClose(reason string)
|
||||
OnError(message string)
|
||||
}
|
||||
|
||||
// SSHClient is a NetBird-aware SSH client exposed to Swift via gomobile.
|
||||
//
|
||||
// It dials through the running NetBird tunnel and runs a standard SSH session
|
||||
// on top with PTY enabled. Host-key verification uses the NetBird-provided
|
||||
// peer SSH host keys, identical to the desktop client.
|
||||
type SSHClient struct {
|
||||
nb *Client
|
||||
mu sync.Mutex
|
||||
listener SSHTerminalListener
|
||||
urlOpener URLOpener
|
||||
|
||||
sshClient *gossh.Client
|
||||
session *gossh.Session
|
||||
stdin io.WriteCloser
|
||||
closed bool
|
||||
}
|
||||
|
||||
// NewSSHClient creates a new SSH client bound to the running NetBird Client.
|
||||
func NewSSHClient(c *Client) *SSHClient {
|
||||
return &SSHClient{nb: c}
|
||||
}
|
||||
|
||||
// SetListener registers the Swift listener. Must be called before Connect to
|
||||
// receive any events.
|
||||
func (s *SSHClient) SetListener(l SSHTerminalListener) {
|
||||
s.mu.Lock()
|
||||
s.listener = l
|
||||
s.mu.Unlock()
|
||||
}
|
||||
|
||||
// SetURLOpener registers the Swift URL opener used to display the device-code
|
||||
// authorization page in an in-app browser when the target peer requires JWT
|
||||
// authentication. Must be set before Connect to be effective.
|
||||
func (s *SSHClient) SetURLOpener(opener URLOpener) {
|
||||
s.mu.Lock()
|
||||
s.urlOpener = opener
|
||||
s.mu.Unlock()
|
||||
}
|
||||
|
||||
// Connect dials the SSH server through the NetBird tunnel and performs the
|
||||
// SSH handshake. It auto-detects the server type via SSH banner inspection
|
||||
// and selects the appropriate authentication path:
|
||||
//
|
||||
// - NetBird-SSH server requiring JWT: launches the OAuth 2.0 device-code
|
||||
// flow, opens the verification URL through the registered URLOpener, and
|
||||
// uses the resulting token as the SSH password. Host-key verification
|
||||
// uses the NetBird peer registry.
|
||||
// - NetBird-SSH server without JWT: authenticates with the NetBird SSH
|
||||
// private key. Host-key verification uses the NetBird peer registry.
|
||||
// - Regular SSH server (e.g. OpenSSH): authenticates with the NetBird key
|
||||
// first (so a user-installed NetBird public key works), then falls back
|
||||
// to the supplied password if non-empty. Host-key verification is
|
||||
// disabled (TOFU pending).
|
||||
//
|
||||
// The password parameter is only consulted for regular SSH servers.
|
||||
//
|
||||
// This is the only way to open a session, deliberately so: the JWT is a bearer
|
||||
// token, and detection is what proves the listener is a NetBird SSH service
|
||||
// before the token is offered to it. A peer whose NetBird SSH is disabled is
|
||||
// served by plain sshd and authenticates by key or password like any other
|
||||
// host, so it stays reachable without an OAuth round trip.
|
||||
func (s *SSHClient) Connect(host string, port int, user, password string) error {
|
||||
cfg, cc := s.nb.sshState()
|
||||
if cc == nil {
|
||||
return errors.New("netbird client not running")
|
||||
}
|
||||
if cfg == nil {
|
||||
return errors.New("netbird config not loaded")
|
||||
}
|
||||
engine := cc.Engine()
|
||||
if engine == nil {
|
||||
return errors.New("netbird engine not available")
|
||||
}
|
||||
|
||||
wgDialer := makeWGDialer(cfg.WgIface, sshDialTimeout)
|
||||
|
||||
serverType := detectServerType(host, port, wgDialer)
|
||||
log.Infof("SSH server type for %s:%d: %s", host, port, serverType)
|
||||
|
||||
authMethods, hostKeyCallback, err := s.buildAuth(cfg, engine, serverType, password)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
clientConfig := &gossh.ClientConfig{
|
||||
User: user,
|
||||
Auth: authMethods,
|
||||
HostKeyCallback: hostKeyCallback,
|
||||
Timeout: sshDialTimeout,
|
||||
}
|
||||
if err := s.dialAndHandshake(host, port, clientConfig, wgDialer); err != nil {
|
||||
return annotateAuthError(err, serverType, user)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// annotateAuthError adds NetBird-specific guidance to an authentication
|
||||
// failure, but only for a server that detection identified as NetBird-SSH.
|
||||
// There a rejection nearly always means the peer's dashboard SSH access is not
|
||||
// configured for this account, which the raw gossh error does not convey.
|
||||
func annotateAuthError(err error, serverType detection.ServerType, user string) error {
|
||||
if !serverType.RequiresJWT() {
|
||||
return err
|
||||
}
|
||||
msg := err.Error()
|
||||
if !strings.Contains(msg, "no supported methods remain") &&
|
||||
!strings.Contains(msg, "unable to authenticate") {
|
||||
return err
|
||||
}
|
||||
return fmt.Errorf("NetBird SSH authentication rejected.\n\n"+
|
||||
"Checklist:\n"+
|
||||
" 1. SSH is enabled for this peer in the NetBird dashboard\n"+
|
||||
" 2. Your account is listed under SSH access for this peer\n"+
|
||||
" 3. The OS username (%q) is mapped to your account\n\n"+
|
||||
"If SSH access is not configured, connect with a password instead.\n\n"+
|
||||
"Original: %w", user, err)
|
||||
}
|
||||
|
||||
// StartSession requests a PTY and starts an interactive shell. Output from
|
||||
// the session is forwarded to the listener via OnData.
|
||||
func (s *SSHClient) StartSession(cols, rows int) error {
|
||||
log.Debugf("SSH: starting session %dx%d", cols, rows)
|
||||
s.mu.Lock()
|
||||
sshClient := s.sshClient
|
||||
s.mu.Unlock()
|
||||
|
||||
if sshClient == nil {
|
||||
return errors.New("ssh client not connected")
|
||||
}
|
||||
|
||||
session, err := sshClient.NewSession()
|
||||
if err != nil {
|
||||
return fmt.Errorf("new session: %w", err)
|
||||
}
|
||||
|
||||
modes := gossh.TerminalModes{
|
||||
gossh.ECHO: 1,
|
||||
gossh.TTY_OP_ISPEED: 14400,
|
||||
gossh.TTY_OP_OSPEED: 14400,
|
||||
gossh.VINTR: 3,
|
||||
gossh.VQUIT: 28,
|
||||
gossh.VERASE: 127,
|
||||
}
|
||||
if err := session.RequestPty("xterm-256color", rows, cols, modes); err != nil {
|
||||
closeQuiet(session, "session after pty error")
|
||||
return fmt.Errorf("request pty: %w", err)
|
||||
}
|
||||
|
||||
stdin, err := session.StdinPipe()
|
||||
if err != nil {
|
||||
closeQuiet(session, "session after stdin error")
|
||||
return fmt.Errorf("stdin pipe: %w", err)
|
||||
}
|
||||
stdout, err := session.StdoutPipe()
|
||||
if err != nil {
|
||||
closeQuiet(session, "session after stdout error")
|
||||
return fmt.Errorf("stdout pipe: %w", err)
|
||||
}
|
||||
stderr, err := session.StderrPipe()
|
||||
if err != nil {
|
||||
closeQuiet(session, "session after stderr error")
|
||||
return fmt.Errorf("stderr pipe: %w", err)
|
||||
}
|
||||
|
||||
if err := session.Shell(); err != nil {
|
||||
closeQuiet(session, "session after shell error")
|
||||
return fmt.Errorf("start shell: %w", err)
|
||||
}
|
||||
|
||||
s.mu.Lock()
|
||||
s.session = session
|
||||
s.stdin = stdin
|
||||
s.mu.Unlock()
|
||||
|
||||
go s.readLoop(stdout, "stdout")
|
||||
go s.readLoop(stderr, "stderr")
|
||||
log.Debug("SSH: session started, shell running")
|
||||
return nil
|
||||
}
|
||||
|
||||
// Write sends data to the SSH session stdin.
|
||||
func (s *SSHClient) Write(data []byte) error {
|
||||
s.mu.Lock()
|
||||
stdin := s.stdin
|
||||
s.mu.Unlock()
|
||||
if stdin == nil {
|
||||
return errors.New("ssh session not started")
|
||||
}
|
||||
if _, err := stdin.Write(data); err != nil {
|
||||
return fmt.Errorf("write stdin: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Resize updates the PTY window size.
|
||||
func (s *SSHClient) Resize(cols, rows int) error {
|
||||
s.mu.Lock()
|
||||
session := s.session
|
||||
s.mu.Unlock()
|
||||
if session == nil {
|
||||
return errors.New("ssh session not started")
|
||||
}
|
||||
return session.WindowChange(rows, cols)
|
||||
}
|
||||
|
||||
// Close terminates the SSH session and underlying connection. Safe to call
|
||||
// multiple times.
|
||||
func (s *SSHClient) Close() error {
|
||||
s.mu.Lock()
|
||||
sshClient := s.sshClient
|
||||
session := s.session
|
||||
stdin := s.stdin
|
||||
s.sshClient = nil
|
||||
s.session = nil
|
||||
s.stdin = nil
|
||||
s.mu.Unlock()
|
||||
|
||||
if stdin != nil {
|
||||
if err := stdin.Close(); err != nil {
|
||||
log.Debugf("ssh: stdin close: %v", err)
|
||||
}
|
||||
}
|
||||
if session != nil {
|
||||
if err := session.Close(); err != nil && !errors.Is(err, io.EOF) {
|
||||
log.Debugf("ssh: session close: %v", err)
|
||||
}
|
||||
}
|
||||
var firstErr error
|
||||
if sshClient != nil {
|
||||
if err := sshClient.Close(); err != nil {
|
||||
firstErr = err
|
||||
}
|
||||
}
|
||||
s.notifyClose("closed by client")
|
||||
return firstErr
|
||||
}
|
||||
|
||||
func (s *SSHClient) buildAuth(cfg *profilemanager.Config, engine *internal.Engine,
|
||||
serverType detection.ServerType, password string) ([]gossh.AuthMethod, gossh.HostKeyCallback, error) {
|
||||
|
||||
switch serverType {
|
||||
case detection.ServerTypeNetBirdJWT:
|
||||
token, err := s.requestJWTToken(cfg)
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("jwt: %w", err)
|
||||
}
|
||||
auths := []gossh.AuthMethod{gossh.Password(token)}
|
||||
return auths, nbssh.CreateHostKeyCallback(&engineHostKeyVerifier{engine: engine}), nil
|
||||
|
||||
case detection.ServerTypeNetBirdNoJWT:
|
||||
if cfg.SSHKey == "" {
|
||||
return nil, nil, errors.New("no NetBird SSH key available")
|
||||
}
|
||||
signer, err := gossh.ParsePrivateKey([]byte(cfg.SSHKey))
|
||||
if err != nil {
|
||||
return nil, nil, fmt.Errorf("parse netbird ssh key: %w", err)
|
||||
}
|
||||
auths := []gossh.AuthMethod{gossh.PublicKeys(signer)}
|
||||
return auths, nbssh.CreateHostKeyCallback(&engineHostKeyVerifier{engine: engine}), nil
|
||||
|
||||
default: // regular SSH
|
||||
var auths []gossh.AuthMethod
|
||||
if cfg.SSHKey != "" {
|
||||
if signer, err := gossh.ParsePrivateKey([]byte(cfg.SSHKey)); err == nil {
|
||||
auths = append(auths, gossh.PublicKeys(signer))
|
||||
} else {
|
||||
log.Debugf("ssh: parse netbird key for regular auth: %v", err)
|
||||
}
|
||||
}
|
||||
if password != "" {
|
||||
pw := password
|
||||
auths = append(auths, gossh.Password(pw))
|
||||
auths = append(auths, gossh.KeyboardInteractive(func(_, _ string, questions []string, _ []bool) ([]string, error) {
|
||||
answers := make([]string, len(questions))
|
||||
for i := range questions {
|
||||
answers[i] = pw
|
||||
}
|
||||
return answers, nil
|
||||
}))
|
||||
}
|
||||
if len(auths) == 0 {
|
||||
return nil, nil, errors.New("no auth method available: provide a password or configure NetBird SSH key")
|
||||
}
|
||||
return auths, gossh.InsecureIgnoreHostKey(), nil // nolint:gosec // TOFU not yet implemented
|
||||
}
|
||||
}
|
||||
|
||||
func (s *SSHClient) requestJWTToken(cfg *profilemanager.Config) (string, error) {
|
||||
// Reuse a cached token so the user is not forced through the browser OAuth
|
||||
// flow on every reconnect. TTL comes from cfg.SSHJWTCacheTTL, same as the
|
||||
// daemon's cache; unset/0 disables caching.
|
||||
if token, ok := s.nb.sshJWTCache.Get(); ok {
|
||||
log.Debug("SSH: reusing cached JWT token")
|
||||
return token, nil
|
||||
}
|
||||
|
||||
s.mu.Lock()
|
||||
urlOpener := s.urlOpener
|
||||
s.mu.Unlock()
|
||||
if urlOpener == nil {
|
||||
return "", errors.New("URL opener not configured for JWT auth")
|
||||
}
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), 5*time.Minute)
|
||||
defer cancel()
|
||||
|
||||
flow, err := auth.NewOAuthFlow(ctx, cfg, false, true, profilemanager.GetLoginHint())
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("create oauth flow: %w", err)
|
||||
}
|
||||
|
||||
flowInfo, err := flow.RequestAuthInfo(ctx)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("request auth info: %w", err)
|
||||
}
|
||||
|
||||
go urlOpener.Open(flowInfo.VerificationURIComplete, flowInfo.UserCode)
|
||||
|
||||
tokenInfo, err := flow.WaitToken(ctx, flowInfo)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("wait for token: %w", err)
|
||||
}
|
||||
|
||||
token := tokenInfo.GetTokenToUse()
|
||||
if token == "" {
|
||||
return "", errors.New("empty token returned by IdP")
|
||||
}
|
||||
|
||||
if ttl := jwtcache.ResolveTTL(cfg.SSHJWTCacheTTL); ttl > 0 {
|
||||
s.nb.sshJWTCache.Store(token, ttl)
|
||||
}
|
||||
return token, nil
|
||||
}
|
||||
|
||||
func (s *SSHClient) dialAndHandshake(host string, port int, clientConfig *gossh.ClientConfig, dialer *net.Dialer) error {
|
||||
addr := net.JoinHostPort(host, strconv.Itoa(port))
|
||||
log.Infof("SSH: connecting to %s as %s", addr, clientConfig.User)
|
||||
|
||||
ctx, cancel := context.WithTimeout(context.Background(), sshDialTimeout)
|
||||
defer cancel()
|
||||
|
||||
conn, err := dialer.DialContext(ctx, "tcp", addr)
|
||||
if err != nil {
|
||||
return fmt.Errorf("dial %s: %w", addr, err)
|
||||
}
|
||||
|
||||
sshConn, chans, reqs, err := gossh.NewClientConn(conn, addr, clientConfig)
|
||||
if err != nil {
|
||||
if cerr := conn.Close(); cerr != nil {
|
||||
log.Debugf("ssh: close after handshake error: %v", cerr)
|
||||
}
|
||||
return fmt.Errorf("ssh handshake: %w", err)
|
||||
}
|
||||
|
||||
s.mu.Lock()
|
||||
s.sshClient = gossh.NewClient(sshConn, chans, reqs)
|
||||
listener := s.listener
|
||||
s.mu.Unlock()
|
||||
|
||||
log.Infof("SSH: connected to %s", addr)
|
||||
if listener != nil {
|
||||
listener.OnConnected()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *SSHClient) readLoop(r io.Reader, name string) {
|
||||
buf := make([]byte, 4096)
|
||||
for {
|
||||
n, err := r.Read(buf)
|
||||
if n > 0 {
|
||||
s.mu.Lock()
|
||||
listener := s.listener
|
||||
s.mu.Unlock()
|
||||
if listener != nil {
|
||||
chunk := make([]byte, n)
|
||||
copy(chunk, buf[:n])
|
||||
listener.OnData(chunk)
|
||||
}
|
||||
}
|
||||
if err != nil {
|
||||
if !errors.Is(err, io.EOF) {
|
||||
log.Debugf("ssh %s read: %v", name, err)
|
||||
}
|
||||
s.notifyClose(err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *SSHClient) notifyClose(reason string) {
|
||||
s.mu.Lock()
|
||||
if s.closed {
|
||||
s.mu.Unlock()
|
||||
return
|
||||
}
|
||||
s.closed = true
|
||||
listener := s.listener
|
||||
s.mu.Unlock()
|
||||
if listener != nil {
|
||||
listener.OnClose(reason)
|
||||
}
|
||||
}
|
||||
|
||||
// engineHostKeyVerifier adapts *internal.Engine to nbssh.HostKeyVerifier.
|
||||
type engineHostKeyVerifier struct {
|
||||
engine *internal.Engine
|
||||
}
|
||||
|
||||
func (v *engineHostKeyVerifier) VerifySSHHostKey(peerAddress string, presented []byte) error {
|
||||
storedKey, found := v.engine.GetPeerSSHKey(peerAddress)
|
||||
if !found {
|
||||
return nbssh.ErrPeerNotFound
|
||||
}
|
||||
return nbssh.VerifyHostKey(storedKey, presented, peerAddress)
|
||||
}
|
||||
|
||||
func detectServerType(host string, port int, dialer *net.Dialer) detection.ServerType {
|
||||
ctx, cancel := context.WithTimeout(context.Background(), sshDetectionTimeout)
|
||||
defer cancel()
|
||||
|
||||
serverType, err := detection.DetectSSHServerType(ctx, dialer, host, port)
|
||||
if err != nil {
|
||||
log.Debugf("ssh: server detection for %s:%d failed: %v (assuming regular SSH)", host, port, err)
|
||||
return detection.ServerTypeRegular
|
||||
}
|
||||
return serverType
|
||||
}
|
||||
|
||||
// makeWGDialer returns a net.Dialer whose sockets are bound to the WireGuard
|
||||
// interface (wgIface, e.g. "utun100"). This is required in the iOS Network
|
||||
// Extension process, where the OS deliberately excludes the provider's own
|
||||
// traffic from the VPN tunnel to prevent routing loops. Without binding to
|
||||
// the WireGuard interface, TCP connections to NetBird peer IPs (100.x.x.x
|
||||
// CGNAT space) would be sent over the physical network and fail with
|
||||
// "network is unreachable". Falls back to an unbound dialer if the interface
|
||||
// cannot be found (e.g. tunnel not yet up).
|
||||
func makeWGDialer(wgIface string, timeout time.Duration) *net.Dialer {
|
||||
return &net.Dialer{
|
||||
Timeout: timeout,
|
||||
Control: func(network, address string, c syscall.RawConn) error {
|
||||
iface, err := net.InterfaceByName(wgIface)
|
||||
if err != nil {
|
||||
log.Debugf("ssh: WG interface %q not found, dialing without bind: %v", wgIface, err)
|
||||
return nil
|
||||
}
|
||||
var innerErr error
|
||||
if ctrlErr := c.Control(func(fd uintptr) {
|
||||
// IP_BOUND_IF (Darwin) = 25: binds the socket to a specific interface index.
|
||||
innerErr = syscall.SetsockoptInt(int(fd), syscall.IPPROTO_IP, 25, iface.Index)
|
||||
}); ctrlErr != nil {
|
||||
return ctrlErr
|
||||
}
|
||||
if innerErr != nil {
|
||||
log.Debugf("ssh: IP_BOUND_IF bind to %q failed: %v", wgIface, innerErr)
|
||||
}
|
||||
return innerErr
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func closeQuiet(c io.Closer, label string) {
|
||||
if c == nil {
|
||||
return
|
||||
}
|
||||
if err := c.Close(); err != nil && !errors.Is(err, io.EOF) {
|
||||
log.Debugf("ssh: close %s: %v", label, err)
|
||||
}
|
||||
}
|
||||
@@ -1,79 +0,0 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/awnumar/memguard"
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
type jwtCache struct {
|
||||
mu sync.RWMutex
|
||||
enclave *memguard.Enclave
|
||||
expiresAt time.Time
|
||||
timer *time.Timer
|
||||
maxTokenSize int
|
||||
}
|
||||
|
||||
func newJWTCache() *jwtCache {
|
||||
return &jwtCache{
|
||||
maxTokenSize: 8192,
|
||||
}
|
||||
}
|
||||
|
||||
func (c *jwtCache) store(token string, maxAge time.Duration) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
c.cleanup()
|
||||
|
||||
if c.timer != nil {
|
||||
c.timer.Stop()
|
||||
}
|
||||
|
||||
tokenBytes := []byte(token)
|
||||
c.enclave = memguard.NewEnclave(tokenBytes)
|
||||
|
||||
c.expiresAt = time.Now().Add(maxAge)
|
||||
|
||||
var timer *time.Timer
|
||||
timer = time.AfterFunc(maxAge, func() {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.timer != timer {
|
||||
return
|
||||
}
|
||||
c.cleanup()
|
||||
c.timer = nil
|
||||
log.Debugf("JWT token cache expired after %v, securely wiped from memory", maxAge)
|
||||
})
|
||||
c.timer = timer
|
||||
}
|
||||
|
||||
func (c *jwtCache) get() (string, bool) {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
|
||||
if c.enclave == nil || time.Now().After(c.expiresAt) {
|
||||
return "", false
|
||||
}
|
||||
|
||||
buffer, err := c.enclave.Open()
|
||||
if err != nil {
|
||||
log.Debugf("Failed to open JWT token enclave: %v", err)
|
||||
return "", false
|
||||
}
|
||||
defer buffer.Destroy()
|
||||
|
||||
token := string(buffer.Bytes())
|
||||
return token, true
|
||||
}
|
||||
|
||||
// cleanup destroys the secure enclave, must be called with lock held
|
||||
func (c *jwtCache) cleanup() {
|
||||
if c.enclave != nil {
|
||||
c.enclave = nil
|
||||
}
|
||||
c.expiresAt = time.Time{}
|
||||
}
|
||||
@@ -26,6 +26,7 @@ import (
|
||||
"github.com/netbirdio/netbird/client/internal/profilemanager"
|
||||
sleephandler "github.com/netbirdio/netbird/client/internal/sleep/handler"
|
||||
"github.com/netbirdio/netbird/client/mdm"
|
||||
"github.com/netbirdio/netbird/client/ssh/jwtcache"
|
||||
"github.com/netbirdio/netbird/client/system"
|
||||
mgm "github.com/netbirdio/netbird/shared/management/client"
|
||||
"github.com/netbirdio/netbird/shared/management/domain"
|
||||
@@ -50,9 +51,6 @@ const (
|
||||
defaultMaxRetryTime = 14 * 24 * time.Hour
|
||||
defaultRetryMultiplier = 1.7
|
||||
|
||||
// JWT token cache TTL for the client daemon (disabled by default)
|
||||
defaultJWTCacheTTL = 0
|
||||
|
||||
errRestoreResidualState = "failed to restore residual state: %v"
|
||||
errProfilesDisabled = "profiles are disabled, you cannot use this feature without profiles enabled"
|
||||
errUpdateSettingsDisabled = "update settings are disabled, you cannot use this feature without update settings enabled"
|
||||
@@ -134,7 +132,7 @@ type Server struct {
|
||||
|
||||
updateManager *updater.Manager
|
||||
|
||||
jwtCache *jwtCache
|
||||
jwtCache *jwtcache.Cache
|
||||
|
||||
// loginAttemptFn stands in for the Management login round trip. Tests set
|
||||
// it to drive the login outcomes that need a server on the other end;
|
||||
@@ -163,7 +161,7 @@ func New(ctx context.Context, logFile string, configFile string, profilesDisable
|
||||
updateSettingsDisabled: updateSettingsDisabled,
|
||||
captureEnabled: captureEnabled,
|
||||
networksDisabled: networksDisabled,
|
||||
jwtCache: newJWTCache(),
|
||||
jwtCache: jwtcache.New(),
|
||||
extendAuthSessionFlow: auth.NewPendingFlow(),
|
||||
probeThrottle: newProbeThrottle(probeThreshold),
|
||||
}
|
||||
@@ -1670,19 +1668,11 @@ func (s *Server) getJWTCacheTTL() time.Duration {
|
||||
config := s.config
|
||||
s.mutex.Unlock()
|
||||
|
||||
if config == nil || config.SSHJWTCacheTTL == nil {
|
||||
return defaultJWTCacheTTL
|
||||
if config == nil {
|
||||
return jwtcache.DefaultTTL
|
||||
}
|
||||
|
||||
seconds := *config.SSHJWTCacheTTL
|
||||
if seconds == 0 {
|
||||
log.Debug("SSH JWT cache disabled (configured to 0)")
|
||||
return 0
|
||||
}
|
||||
|
||||
ttl := time.Duration(seconds) * time.Second
|
||||
log.Debugf("SSH JWT cache TTL set to %v from config", ttl)
|
||||
return ttl
|
||||
return jwtcache.ResolveTTL(config.SSHJWTCacheTTL)
|
||||
}
|
||||
|
||||
// RequestJWTAuth initiates JWT authentication flow for SSH
|
||||
@@ -1704,7 +1694,7 @@ func (s *Server) RequestJWTAuth(
|
||||
|
||||
jwtCacheTTL := s.getJWTCacheTTL()
|
||||
if jwtCacheTTL > 0 {
|
||||
if cachedToken, found := s.jwtCache.get(); found {
|
||||
if cachedToken, found := s.jwtCache.Get(); found {
|
||||
log.Debugf("JWT token found in cache, returning cached token for SSH authentication")
|
||||
|
||||
return &proto.RequestJWTAuthResponse{
|
||||
@@ -1777,7 +1767,7 @@ func (s *Server) WaitJWTToken(
|
||||
|
||||
jwtCacheTTL := s.getJWTCacheTTL()
|
||||
if jwtCacheTTL > 0 {
|
||||
s.jwtCache.store(token, jwtCacheTTL)
|
||||
s.jwtCache.Store(token, jwtCacheTTL)
|
||||
log.Debugf("JWT token cached for SSH authentication, TTL: %v", jwtCacheTTL)
|
||||
} else {
|
||||
log.Debug("JWT caching disabled, not storing token")
|
||||
|
||||
109
client/ssh/jwtcache/cache.go
Normal file
109
client/ssh/jwtcache/cache.go
Normal file
@@ -0,0 +1,109 @@
|
||||
// Package jwtcache provides an in-memory, TTL-bound cache for SSH JWT tokens.
|
||||
// The token is kept in a secure memguard enclave and wiped from memory when it
|
||||
// expires. It is shared by the daemon gRPC server and the mobile SDKs, which
|
||||
// have no daemon process to delegate caching to.
|
||||
package jwtcache
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/awnumar/memguard"
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
// DefaultTTL is used when no TTL is configured: caching disabled.
|
||||
const DefaultTTL = 0
|
||||
|
||||
// Cache stores a single JWT token in a secure enclave until it expires.
|
||||
type Cache struct {
|
||||
mu sync.RWMutex
|
||||
enclave *memguard.Enclave
|
||||
expiresAt time.Time
|
||||
timer *time.Timer
|
||||
maxTokenSize int
|
||||
}
|
||||
|
||||
// New creates an empty Cache.
|
||||
func New() *Cache {
|
||||
return &Cache{
|
||||
maxTokenSize: 8192,
|
||||
}
|
||||
}
|
||||
|
||||
// Store caches the token for maxAge. A previously stored token is wiped.
|
||||
func (c *Cache) Store(token string, maxAge time.Duration) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
c.cleanup()
|
||||
|
||||
if c.timer != nil {
|
||||
c.timer.Stop()
|
||||
}
|
||||
|
||||
tokenBytes := []byte(token)
|
||||
c.enclave = memguard.NewEnclave(tokenBytes)
|
||||
|
||||
c.expiresAt = time.Now().Add(maxAge)
|
||||
|
||||
var timer *time.Timer
|
||||
timer = time.AfterFunc(maxAge, func() {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.timer != timer {
|
||||
return
|
||||
}
|
||||
c.cleanup()
|
||||
c.timer = nil
|
||||
log.Debugf("JWT token cache expired after %v, securely wiped from memory", maxAge)
|
||||
})
|
||||
c.timer = timer
|
||||
}
|
||||
|
||||
// Get returns the cached token, or false if none is stored or it has expired.
|
||||
func (c *Cache) Get() (string, bool) {
|
||||
c.mu.RLock()
|
||||
defer c.mu.RUnlock()
|
||||
|
||||
if c.enclave == nil || time.Now().After(c.expiresAt) {
|
||||
return "", false
|
||||
}
|
||||
|
||||
buffer, err := c.enclave.Open()
|
||||
if err != nil {
|
||||
log.Debugf("Failed to open JWT token enclave: %v", err)
|
||||
return "", false
|
||||
}
|
||||
defer buffer.Destroy()
|
||||
|
||||
token := string(buffer.Bytes())
|
||||
return token, true
|
||||
}
|
||||
|
||||
// cleanup destroys the secure enclave, must be called with lock held
|
||||
func (c *Cache) cleanup() {
|
||||
if c.enclave != nil {
|
||||
c.enclave = nil
|
||||
}
|
||||
c.expiresAt = time.Time{}
|
||||
}
|
||||
|
||||
// ResolveTTL converts the configured TTL (seconds, from
|
||||
// profilemanager.Config.SSHJWTCacheTTL) into a duration. Returns DefaultTTL
|
||||
// when unset; 0 means caching is disabled.
|
||||
func ResolveTTL(configuredSeconds *int) time.Duration {
|
||||
if configuredSeconds == nil {
|
||||
return DefaultTTL
|
||||
}
|
||||
|
||||
seconds := *configuredSeconds
|
||||
if seconds == 0 {
|
||||
log.Debug("SSH JWT cache disabled (configured to 0)")
|
||||
return 0
|
||||
}
|
||||
|
||||
ttl := time.Duration(seconds) * time.Second
|
||||
log.Debugf("SSH JWT cache TTL set to %v from config", ttl)
|
||||
return ttl
|
||||
}
|
||||
@@ -243,7 +243,7 @@ func (s *Server) setUserEnvironmentVariables(envMap map[string]string, userProfi
|
||||
|
||||
// prepareCommandEnv prepares environment variables for command execution on Windows
|
||||
func (s *Server) prepareCommandEnv(logger *log.Entry, localUser *user.User, session ssh.Session) []string {
|
||||
username, domain := s.parseUsername(localUser.Username)
|
||||
username, domain := parseUsername(localUser.Username)
|
||||
userEnv, err := s.getUserEnvironment(logger, username, domain)
|
||||
if err != nil {
|
||||
log.Debugf("failed to get user environment for %s\\%s, using fallback: %v", domain, username, err)
|
||||
@@ -383,7 +383,7 @@ func (s *Server) executeCommandWithPty(logger *log.Entry, session ssh.Session, _
|
||||
return false
|
||||
}
|
||||
|
||||
username, domain := s.parseUsername(localUser.Username)
|
||||
username, domain := parseUsername(localUser.Username)
|
||||
shell := getUserShell(localUser.Uid)
|
||||
|
||||
req := PtyExecutionRequest{
|
||||
|
||||
@@ -133,7 +133,12 @@ func (s *Server) checkPrivilegedPortAccess(forwardType string, port uint32, resu
|
||||
return nil
|
||||
}
|
||||
|
||||
if result.User != nil && isPrivilegedUsername(result.User.Username) {
|
||||
// Only uid 0 may bind below the threshold, which is the kernel's own rule and
|
||||
// is asked directly rather than through isPrivilegedOrUnknown: that helper
|
||||
// reports an account it cannot evaluate as privileged, which is safe for a
|
||||
// refusal and unsafe for a grant such as this one. Windows has returned
|
||||
// above, so Uid here is a Unix uid and never a SID.
|
||||
if result.User != nil && result.User.Uid == "0" {
|
||||
return nil
|
||||
}
|
||||
|
||||
|
||||
16
client/ssh/server/privileges_other.go
Normal file
16
client/ssh/server/privileges_other.go
Normal file
@@ -0,0 +1,16 @@
|
||||
//go:build !windows
|
||||
|
||||
package server
|
||||
|
||||
// isProcessElevated is only meaningful on Windows; other platforms use the
|
||||
// effective UID check in isCurrentProcessPrivileged.
|
||||
func isProcessElevated() bool {
|
||||
return false
|
||||
}
|
||||
|
||||
// isWindowsAccountPrivilegedOrUnknown is only reachable on Windows. Report
|
||||
// privileged on other platforms so a caller refusing privileged accounts fails
|
||||
// closed.
|
||||
func isWindowsAccountPrivilegedOrUnknown(string) bool {
|
||||
return true
|
||||
}
|
||||
228
client/ssh/server/privileges_windows.go
Normal file
228
client/ssh/server/privileges_windows.go
Normal file
@@ -0,0 +1,228 @@
|
||||
//go:build windows
|
||||
|
||||
package server
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
"unsafe"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"golang.org/x/sys/windows"
|
||||
)
|
||||
|
||||
var (
|
||||
netapi32 = windows.NewLazySystemDLL("netapi32.dll")
|
||||
procNetUserGetLocalGroups = netapi32.NewProc("NetUserGetLocalGroups")
|
||||
)
|
||||
|
||||
const (
|
||||
// lgIncludeIndirect makes NetUserGetLocalGroups also return local groups
|
||||
// the user belongs to through a global group.
|
||||
lgIncludeIndirect = 0x1
|
||||
maxPreferredLength = 0xFFFFFFFF
|
||||
)
|
||||
|
||||
// localGroupUsersInfo0 mirrors LOCALGROUP_USERS_INFO_0.
|
||||
type localGroupUsersInfo0 struct {
|
||||
name *uint16
|
||||
}
|
||||
|
||||
// isProcessElevated reports whether the current process token is elevated
|
||||
// (TokenElevation): true for elevated administrators, the built-in
|
||||
// Administrator, administrators with UAC disabled, and SYSTEM; false for
|
||||
// standard users and administrators running with a UAC-filtered token.
|
||||
func isProcessElevated() bool {
|
||||
return windows.GetCurrentProcessToken().IsElevated()
|
||||
}
|
||||
|
||||
// isWindowsAccountPrivilegedOrUnknown reports whether the account is privileged
|
||||
// on this machine: a well-known service account, a built-in Administrator
|
||||
// (RID 500), or a member of the local Administrators group, directly or through
|
||||
// nested groups.
|
||||
//
|
||||
// An account whose privilege cannot be determined counts as privileged, which
|
||||
// is why the name says "or unknown". That is fail-closed for a caller that
|
||||
// refuses privileged accounts, and fail-open for a caller that grants something
|
||||
// to them, so only the former may use this.
|
||||
func isWindowsAccountPrivilegedOrUnknown(username string) bool {
|
||||
sid, _, _, err := windows.LookupSID("", username)
|
||||
if err != nil {
|
||||
log.Warnf("privilege check: SID lookup for %q failed, treating as privileged: %v", username, err)
|
||||
return true
|
||||
}
|
||||
|
||||
if isPrivilegedUserSID(sid) {
|
||||
return true
|
||||
}
|
||||
|
||||
member, err := isLocalAdminsMember(username)
|
||||
if err != nil {
|
||||
log.Warnf("privilege check: cannot determine Administrators membership for %q, treating as privileged: %v", username, err)
|
||||
return true
|
||||
}
|
||||
return member
|
||||
}
|
||||
|
||||
// isPrivilegedUserSID reports whether the SID itself identifies a privileged
|
||||
// principal, without consulting group membership.
|
||||
func isPrivilegedUserSID(sid *windows.SID) bool {
|
||||
wellKnown := []windows.WELL_KNOWN_SID_TYPE{
|
||||
windows.WinLocalSystemSid,
|
||||
windows.WinLocalServiceSid,
|
||||
windows.WinNetworkServiceSid,
|
||||
windows.WinBuiltinAdministratorsSid,
|
||||
}
|
||||
for _, sidType := range wellKnown {
|
||||
if sid.IsWellKnown(sidType) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return isBuiltinAdministratorSID(sid)
|
||||
}
|
||||
|
||||
// isBuiltinAdministratorSID reports whether the SID is a machine or domain
|
||||
// built-in Administrator account (S-1-5-21-...-500). RID 500 is reserved for
|
||||
// that account; it can be renamed but cannot be removed from the
|
||||
// Administrators group.
|
||||
func isBuiltinAdministratorSID(sid *windows.SID) bool {
|
||||
if sid.IdentifierAuthority() != windows.SECURITY_NT_AUTHORITY {
|
||||
return false
|
||||
}
|
||||
count := sid.SubAuthorityCount()
|
||||
if count < 2 || sid.SubAuthority(0) != 21 {
|
||||
return false
|
||||
}
|
||||
return sid.SubAuthority(uint32(count-1)) == 500
|
||||
}
|
||||
|
||||
// isLocalAdminsMember reports whether the account is a member of the local
|
||||
// Administrators group.
|
||||
//
|
||||
// Local accounts are checked against the local SAM, which is authoritative for
|
||||
// them and, unlike a token, cannot under-report: UAC filters the tokens of
|
||||
// local administrators, and a filtered token carries Administrators as
|
||||
// deny-only, which a membership check on the token would read as "not a
|
||||
// member". Domain accounts are exempt from that filtering, so for them an S4U
|
||||
// token is preferred because its group list is LSA's transitive expansion and
|
||||
// therefore covers nested and universal groups plus the machine's own local
|
||||
// groups. NetUserGetLocalGroups expands only one global-group hop but needs no
|
||||
// logon, so it serves as the fallback when no token can be obtained.
|
||||
func isLocalAdminsMember(username string) (bool, error) {
|
||||
adminSid, err := windows.CreateWellKnownSid(windows.WinBuiltinAdministratorsSid)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("create Administrators SID: %w", err)
|
||||
}
|
||||
|
||||
account, domain := parseUsername(username)
|
||||
if NewPrivilegeDropper().isLocalUser(domain) {
|
||||
return localGroupsContainSID(account, adminSid)
|
||||
}
|
||||
|
||||
member, s4uErr := s4uTokenIsMember(account, domain, adminSid)
|
||||
if s4uErr == nil {
|
||||
return member, nil
|
||||
}
|
||||
log.Debugf("privilege check: S4U membership check for %q failed, falling back to local group enumeration: %v", username, s4uErr)
|
||||
|
||||
member, err = localGroupsContainSID(buildUserCpn(account, domain), adminSid)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("S4U check: %w; local group enumeration: %w", s4uErr, err)
|
||||
}
|
||||
return member, nil
|
||||
}
|
||||
|
||||
// s4uTokenIsMember obtains an S4U token for the account and checks whether the
|
||||
// given SID is enabled in it.
|
||||
func s4uTokenIsMember(account, domain string, sid *windows.SID) (bool, error) {
|
||||
token, err := generateS4UUserToken(log.NewEntry(log.StandardLogger()), account, domain)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
defer func() {
|
||||
if err := windows.CloseHandle(token); err != nil {
|
||||
log.Debugf("close S4U token: %v", err)
|
||||
}
|
||||
}()
|
||||
return windows.Token(token).IsMember(sid)
|
||||
}
|
||||
|
||||
// localGroupsContainSID reports whether the wanted group is among the local
|
||||
// groups the account belongs to, directly or through a global group.
|
||||
//
|
||||
// The wanted SID is resolved to its group name once and compared against the
|
||||
// enumerated names. Well-known SIDs resolve from a static table, so that lookup
|
||||
// needs no domain controller, and it keeps the comparison correct for a renamed
|
||||
// or localized group because both sides then carry the new name. Resolving each
|
||||
// enumerated name back to a SID instead would add a lookup per group that can
|
||||
// block until it times out while a domain controller is unreachable, and cannot
|
||||
// change the outcome: the names enumerated here are local groups of this
|
||||
// machine, whose names are unique, so a name match identifies the group.
|
||||
//
|
||||
// A failure to resolve the wanted SID is returned rather than reported as
|
||||
// "not a member", so a privilege check built on this fails closed.
|
||||
func localGroupsContainSID(username string, want *windows.SID) (bool, error) {
|
||||
wantName, _, _, err := want.LookupAccount("")
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("resolve group SID %s to a name: %w", want, err)
|
||||
}
|
||||
|
||||
groups, err := netUserGetLocalGroups(username)
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
|
||||
for _, group := range groups {
|
||||
if strings.EqualFold(group, wantName) {
|
||||
return true, nil
|
||||
}
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// netUserGetLocalGroups returns the names of the local groups the account is a
|
||||
// member of, including indirect membership through global groups.
|
||||
func netUserGetLocalGroups(username string) ([]string, error) {
|
||||
name16, err := windows.UTF16PtrFromString(username)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("convert username: %w", err)
|
||||
}
|
||||
|
||||
var buf *byte
|
||||
var entriesRead, totalEntries uint32
|
||||
status, _, _ := procNetUserGetLocalGroups.Call(
|
||||
0, // local server
|
||||
uintptr(unsafe.Pointer(name16)),
|
||||
0, // level 0: LOCALGROUP_USERS_INFO_0
|
||||
lgIncludeIndirect,
|
||||
uintptr(unsafe.Pointer(&buf)),
|
||||
maxPreferredLength,
|
||||
uintptr(unsafe.Pointer(&entriesRead)),
|
||||
uintptr(unsafe.Pointer(&totalEntries)),
|
||||
)
|
||||
if status != 0 {
|
||||
return nil, fmt.Errorf("NetUserGetLocalGroups for %q: status %d", username, status)
|
||||
}
|
||||
if buf == nil {
|
||||
return nil, nil
|
||||
}
|
||||
defer func() {
|
||||
if err := windows.NetApiBufferFree(buf); err != nil {
|
||||
log.Debugf("free NetApi buffer: %v", err)
|
||||
}
|
||||
}()
|
||||
|
||||
// MAX_PREFERRED_LENGTH makes the API allocate as much as it needs, so a
|
||||
// short read is not expected. Report it rather than silently returning a
|
||||
// subset of the account's groups.
|
||||
if entriesRead != totalEntries {
|
||||
return nil, fmt.Errorf("NetUserGetLocalGroups for %q returned %d of %d groups", username, entriesRead, totalEntries)
|
||||
}
|
||||
|
||||
entries := unsafe.Slice((*localGroupUsersInfo0)(unsafe.Pointer(buf)), entriesRead)
|
||||
groups := make([]string, 0, entriesRead)
|
||||
for _, entry := range entries {
|
||||
groups = append(groups, windows.UTF16PtrToString(entry.name))
|
||||
}
|
||||
return groups, nil
|
||||
}
|
||||
293
client/ssh/server/privileges_windows_test.go
Normal file
293
client/ssh/server/privileges_windows_test.go
Normal file
@@ -0,0 +1,293 @@
|
||||
//go:build windows
|
||||
|
||||
package server
|
||||
|
||||
import (
|
||||
"os/user"
|
||||
"testing"
|
||||
"unsafe"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/sys/windows"
|
||||
)
|
||||
|
||||
// filterNormalAccount limits NetUserEnum to normal user accounts.
|
||||
const filterNormalAccount = 0x2
|
||||
|
||||
// TOKEN_ELEVATION_TYPE values.
|
||||
const (
|
||||
tokenElevationTypeDefault = 1
|
||||
tokenElevationTypeFull = 2
|
||||
tokenElevationTypeLimited = 3
|
||||
)
|
||||
|
||||
// tokenElevationType reads TokenElevationType from a token.
|
||||
func tokenElevationType(token windows.Token) (uint32, error) {
|
||||
var elevationType, returnedLen uint32
|
||||
err := windows.GetTokenInformation(token, windows.TokenElevationType,
|
||||
(*byte)(unsafe.Pointer(&elevationType)), uint32(unsafe.Sizeof(elevationType)), &returnedLen)
|
||||
if err != nil {
|
||||
return 0, err
|
||||
}
|
||||
return elevationType, nil
|
||||
}
|
||||
|
||||
// userInfo0 mirrors USER_INFO_0.
|
||||
type userInfo0 struct {
|
||||
name *uint16
|
||||
}
|
||||
|
||||
func mustParseSID(t *testing.T, s string) *windows.SID {
|
||||
t.Helper()
|
||||
sid, err := windows.StringToSid(s)
|
||||
require.NoError(t, err, "parse SID %s", s)
|
||||
return sid
|
||||
}
|
||||
|
||||
// localAccountNames returns the names of the local user accounts.
|
||||
func localAccountNames(t *testing.T) []string {
|
||||
t.Helper()
|
||||
|
||||
var buf *byte
|
||||
var entriesRead, totalEntries, resume uint32
|
||||
err := windows.NetUserEnum(nil, 0, filterNormalAccount, &buf, maxPreferredLength,
|
||||
&entriesRead, &totalEntries, &resume)
|
||||
require.NoError(t, err, "enumerate local users")
|
||||
t.Cleanup(func() {
|
||||
require.NoError(t, windows.NetApiBufferFree(buf), "free NetApi buffer")
|
||||
})
|
||||
|
||||
entries := unsafe.Slice((*userInfo0)(unsafe.Pointer(buf)), entriesRead)
|
||||
names := make([]string, 0, entriesRead)
|
||||
for _, entry := range entries {
|
||||
names = append(names, windows.UTF16PtrToString(entry.name))
|
||||
}
|
||||
return names
|
||||
}
|
||||
|
||||
// localAccountNameByRID returns the name of the local account carrying the
|
||||
// given RID. Accounts such as Administrator and Guest can be renamed and are
|
||||
// localized, so tests must not name them literally.
|
||||
func localAccountNameByRID(t *testing.T, rid uint32) string {
|
||||
t.Helper()
|
||||
|
||||
for _, name := range localAccountNames(t) {
|
||||
sid, _, _, err := windows.LookupSID("", name)
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
if sid.IdentifierAuthority() != windows.SECURITY_NT_AUTHORITY {
|
||||
continue
|
||||
}
|
||||
count := sid.SubAuthorityCount()
|
||||
if count < 2 || sid.SubAuthority(0) != 21 {
|
||||
continue
|
||||
}
|
||||
if sid.SubAuthority(uint32(count-1)) == rid {
|
||||
return name
|
||||
}
|
||||
}
|
||||
|
||||
t.Fatalf("no local account with RID %d", rid)
|
||||
return ""
|
||||
}
|
||||
|
||||
// wellKnownAccountName resolves a well-known SID to the qualified account name
|
||||
// the local system uses for it, which is localized.
|
||||
func wellKnownAccountName(t *testing.T, sidType windows.WELL_KNOWN_SID_TYPE) string {
|
||||
t.Helper()
|
||||
|
||||
sid, err := windows.CreateWellKnownSid(sidType)
|
||||
require.NoError(t, err, "create well-known SID")
|
||||
name, domain, _, err := sid.LookupAccount("")
|
||||
require.NoError(t, err, "resolve %s to an account name", sid)
|
||||
if domain == "" {
|
||||
return name
|
||||
}
|
||||
return domain + `\` + name
|
||||
}
|
||||
|
||||
func TestIsBuiltinAdministratorSID(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
sid string
|
||||
want bool
|
||||
}{
|
||||
{"machine_administrator", "S-1-5-21-1111111111-2222222222-3333333333-500", true},
|
||||
{"domain_administrator", "S-1-5-21-3390233681-4087452608-412898826-500", true},
|
||||
{"regular_user", "S-1-5-21-1111111111-2222222222-3333333333-1001", false},
|
||||
{"guest_account", "S-1-5-21-1111111111-2222222222-3333333333-501", false},
|
||||
{"domain_admins_group", "S-1-5-21-1111111111-2222222222-3333333333-512", false},
|
||||
{"system", "S-1-5-18", false},
|
||||
{"administrators_group", "S-1-5-32-544", false},
|
||||
{"non_nt_authority", "S-1-1-0", false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := isBuiltinAdministratorSID(mustParseSID(t, tt.sid))
|
||||
assert.Equal(t, tt.want, result, "RID 500 detection for %s", tt.sid)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsPrivilegedUserSID(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
sid string
|
||||
want bool
|
||||
}{
|
||||
{"local_system", "S-1-5-18", true},
|
||||
{"local_service", "S-1-5-19", true},
|
||||
{"network_service", "S-1-5-20", true},
|
||||
{"administrators_group", "S-1-5-32-544", true},
|
||||
{"builtin_administrator", "S-1-5-21-1111111111-2222222222-3333333333-500", true},
|
||||
{"regular_user", "S-1-5-21-1111111111-2222222222-3333333333-1001", false},
|
||||
{"users_group", "S-1-5-32-545", false},
|
||||
{"everyone", "S-1-1-0", false},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := isPrivilegedUserSID(mustParseSID(t, tt.sid))
|
||||
assert.Equal(t, tt.want, result, "SID privilege classification for %s", tt.sid)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsWindowsAccountPrivilegedOrUnknown(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
username string
|
||||
want bool
|
||||
}{
|
||||
{"system", wellKnownAccountName(t, windows.WinLocalSystemSid), true},
|
||||
{"local_service", wellKnownAccountName(t, windows.WinLocalServiceSid), true},
|
||||
{"network_service", wellKnownAccountName(t, windows.WinNetworkServiceSid), true},
|
||||
{"administrators_group", wellKnownAccountName(t, windows.WinBuiltinAdministratorsSid), true},
|
||||
// The built-in Administrator (RID 500) and Guest (RID 501) accounts
|
||||
// exist on every Windows installation, though they may be disabled.
|
||||
{"builtin_administrator", localAccountNameByRID(t, 500), true},
|
||||
{"guest", localAccountNameByRID(t, 501), false},
|
||||
// Unresolvable accounts fail closed.
|
||||
{"nonexistent_user", "netbird-no-such-user", true},
|
||||
{"empty_username", "", true},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := isWindowsAccountPrivilegedOrUnknown(tt.username)
|
||||
assert.Equal(t, tt.want, result, "account privilege classification for %q", tt.username)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestIsProcessElevated(t *testing.T) {
|
||||
elevated := isProcessElevated()
|
||||
|
||||
// TokenElevationType is a second, independent view of the same token:
|
||||
// Full means elevated and Limited means a filtered administrator, while
|
||||
// Default covers both a standard user and an administrator with no linked
|
||||
// token (UAC off, the built-in Administrator, SYSTEM), so it implies nothing.
|
||||
elevationType, err := tokenElevationType(windows.GetCurrentProcessToken())
|
||||
require.NoError(t, err, "read token elevation type")
|
||||
|
||||
adminSid, err := windows.CreateWellKnownSid(windows.WinBuiltinAdministratorsSid)
|
||||
require.NoError(t, err, "create Administrators SID")
|
||||
|
||||
// Token(0) makes CheckTokenMembership evaluate the caller's own token. It
|
||||
// counts only enabled SIDs, so a filtered administrator reports false here.
|
||||
member, err := windows.Token(0).IsMember(adminSid)
|
||||
require.NoError(t, err, "check own Administrators membership")
|
||||
|
||||
t.Logf("elevated=%v elevationType=%d memberOfAdministrators=%v", elevated, elevationType, member)
|
||||
|
||||
switch elevationType {
|
||||
case tokenElevationTypeFull:
|
||||
assert.True(t, elevated, "a token of elevation type Full must report elevated")
|
||||
case tokenElevationTypeLimited:
|
||||
assert.False(t, elevated, "a filtered administrator token must not report elevated")
|
||||
}
|
||||
|
||||
// Administrators enabled in the token means the token wields administrative
|
||||
// rights, which is what elevation reports.
|
||||
if member {
|
||||
assert.True(t, elevated, "token with enabled Administrators membership must report elevated")
|
||||
}
|
||||
}
|
||||
|
||||
// TestS4UMembershipAgreesWithLocalGroups exercises the S4U token path used
|
||||
// for domain accounts. S4U logons need the TCB privilege, so the test runs
|
||||
// only as SYSTEM (which is how CI executes the suite). For local accounts the
|
||||
// token's Administrators membership must agree with the SAM enumeration.
|
||||
func TestS4UMembershipAgreesWithLocalGroups(t *testing.T) {
|
||||
system, err := windows.CreateWellKnownSid(windows.WinLocalSystemSid)
|
||||
require.NoError(t, err, "create SYSTEM SID")
|
||||
current, err := user.Current()
|
||||
require.NoError(t, err, "get current user")
|
||||
if current.Uid != system.String() {
|
||||
t.Skipf("S4U logon requires SYSTEM (running as %s)", current.Username)
|
||||
}
|
||||
|
||||
adminSid, err := windows.CreateWellKnownSid(windows.WinBuiltinAdministratorsSid)
|
||||
require.NoError(t, err, "create Administrators SID")
|
||||
|
||||
checked := 0
|
||||
for _, name := range localAccountNames(t) {
|
||||
viaToken, err := s4uTokenIsMember(name, ".", adminSid)
|
||||
if err != nil {
|
||||
// Disabled or logon-restricted accounts cannot get an S4U logon.
|
||||
t.Logf("skipping %s: %v", name, err)
|
||||
continue
|
||||
}
|
||||
viaSAM, err := localGroupsContainSID(name, adminSid)
|
||||
require.NoError(t, err, "enumerate local groups for %s", name)
|
||||
|
||||
assert.Equal(t, viaSAM, viaToken, "S4U token and SAM enumeration must agree on Administrators membership for %s", name)
|
||||
checked++
|
||||
}
|
||||
// Ineligible accounts are skipped, so without this the test could report
|
||||
// success while comparing nothing at all.
|
||||
require.Positive(t, checked, "no local account completed an S4U logon, so nothing was compared")
|
||||
t.Logf("checked %d local accounts via S4U", checked)
|
||||
}
|
||||
|
||||
// TestLocalGroupsContainSID_Administrator checks the positive case against the
|
||||
// built-in Administrator, a member of Administrators on every installation.
|
||||
func TestLocalGroupsContainSID_Administrator(t *testing.T) {
|
||||
adminSid, err := windows.CreateWellKnownSid(windows.WinBuiltinAdministratorsSid)
|
||||
require.NoError(t, err, "create Administrators SID")
|
||||
|
||||
administrator := localAccountNameByRID(t, 500)
|
||||
member, err := localGroupsContainSID(administrator, adminSid)
|
||||
require.NoError(t, err, "enumerate local groups for %s", administrator)
|
||||
assert.True(t, member, "%s is a member of the Administrators group", administrator)
|
||||
}
|
||||
|
||||
// TestLocalGroupsContainSID_UnresolvableGroupFailsClosed covers a wanted SID
|
||||
// that resolves to no group: the error must surface rather than being reported
|
||||
// as "not a member", so the privilege check treats the account as privileged.
|
||||
func TestLocalGroupsContainSID_UnresolvableGroupFailsClosed(t *testing.T) {
|
||||
unknown := mustParseSID(t, "S-1-5-21-1111111111-2222222222-3333333333-4444")
|
||||
|
||||
_, err := localGroupsContainSID(localAccountNameByRID(t, 500), unknown)
|
||||
require.Error(t, err, "must report an error when the wanted group cannot be identified")
|
||||
}
|
||||
|
||||
func TestLocalGroupsContainSID_Guest(t *testing.T) {
|
||||
guestsSid, err := windows.CreateWellKnownSid(windows.WinBuiltinGuestsSid)
|
||||
require.NoError(t, err, "create Guests SID")
|
||||
adminsSid, err := windows.CreateWellKnownSid(windows.WinBuiltinAdministratorsSid)
|
||||
require.NoError(t, err, "create Administrators SID")
|
||||
|
||||
guest := localAccountNameByRID(t, 501)
|
||||
|
||||
inGuests, err := localGroupsContainSID(guest, guestsSid)
|
||||
require.NoError(t, err, "enumerate local groups for %s", guest)
|
||||
assert.True(t, inGuests, "%s is a member of the Guests group", guest)
|
||||
|
||||
inAdmins, err := localGroupsContainSID(guest, adminsSid)
|
||||
require.NoError(t, err, "enumerate local groups for %s", guest)
|
||||
assert.False(t, inAdmins, "%s is not a member of the Administrators group", guest)
|
||||
}
|
||||
@@ -239,6 +239,7 @@ func TestServer_PrivilegedPortAccess(t *testing.T) {
|
||||
forwardType string
|
||||
port uint32
|
||||
username string
|
||||
uid string
|
||||
expectError bool
|
||||
errorMsg string
|
||||
skipOnWindows bool
|
||||
@@ -248,6 +249,7 @@ func TestServer_PrivilegedPortAccess(t *testing.T) {
|
||||
forwardType: "remote",
|
||||
port: 80,
|
||||
username: "testuser",
|
||||
uid: "1000",
|
||||
expectError: true,
|
||||
errorMsg: "cannot bind to privileged port",
|
||||
skipOnWindows: true,
|
||||
@@ -257,6 +259,7 @@ func TestServer_PrivilegedPortAccess(t *testing.T) {
|
||||
forwardType: "tcpip-forward",
|
||||
port: 443,
|
||||
username: "testuser",
|
||||
uid: "1000",
|
||||
expectError: true,
|
||||
errorMsg: "cannot bind to privileged port",
|
||||
skipOnWindows: true,
|
||||
@@ -266,6 +269,7 @@ func TestServer_PrivilegedPortAccess(t *testing.T) {
|
||||
forwardType: "remote",
|
||||
port: 8080,
|
||||
username: "testuser",
|
||||
uid: "1000",
|
||||
expectError: false,
|
||||
},
|
||||
{
|
||||
@@ -273,6 +277,7 @@ func TestServer_PrivilegedPortAccess(t *testing.T) {
|
||||
forwardType: "remote",
|
||||
port: 0,
|
||||
username: "testuser",
|
||||
uid: "1000",
|
||||
expectError: false,
|
||||
},
|
||||
{
|
||||
@@ -280,13 +285,35 @@ func TestServer_PrivilegedPortAccess(t *testing.T) {
|
||||
forwardType: "remote",
|
||||
port: 22,
|
||||
username: "root",
|
||||
uid: "0",
|
||||
expectError: false,
|
||||
},
|
||||
{
|
||||
// Only uid 0 is privileged, whatever the account is called.
|
||||
name: "uid 0 under another name may bind a privileged port",
|
||||
forwardType: "remote",
|
||||
port: 22,
|
||||
username: "toor",
|
||||
uid: "0",
|
||||
expectError: false,
|
||||
skipOnWindows: true,
|
||||
},
|
||||
{
|
||||
name: "account named root without uid 0 may not",
|
||||
forwardType: "remote",
|
||||
port: 22,
|
||||
username: "root",
|
||||
uid: "1000",
|
||||
expectError: true,
|
||||
errorMsg: "cannot bind to privileged port",
|
||||
skipOnWindows: true,
|
||||
},
|
||||
{
|
||||
name: "local forward privileged port allowed for non-root",
|
||||
forwardType: "local",
|
||||
port: 80,
|
||||
username: "testuser",
|
||||
uid: "1000",
|
||||
expectError: false,
|
||||
},
|
||||
}
|
||||
@@ -299,7 +326,7 @@ func TestServer_PrivilegedPortAccess(t *testing.T) {
|
||||
|
||||
result := PrivilegeCheckResult{
|
||||
Allowed: true,
|
||||
User: &user.User{Username: tt.username},
|
||||
User: &user.User{Username: tt.username, Uid: tt.uid},
|
||||
}
|
||||
|
||||
err := server.checkPrivilegedPortAccess(tt.forwardType, tt.port, result)
|
||||
@@ -420,6 +447,13 @@ func TestServer_PortConflictHandling(t *testing.T) {
|
||||
|
||||
func TestServer_IsPrivilegedUser(t *testing.T) {
|
||||
|
||||
// Windows classification depends on account SIDs and group membership, and
|
||||
// the accounts involved carry localized, renameable names. It is covered by
|
||||
// TestIsWindowsAccountPrivileged, which resolves them from well-known SIDs.
|
||||
if runtime.GOOS == "windows" {
|
||||
t.Skip("covered by TestIsWindowsAccountPrivileged")
|
||||
}
|
||||
|
||||
tests := []struct {
|
||||
username string
|
||||
expected bool
|
||||
@@ -440,44 +474,16 @@ func TestServer_IsPrivilegedUser(t *testing.T) {
|
||||
expected: false,
|
||||
description: "empty username should not be privileged",
|
||||
},
|
||||
}
|
||||
|
||||
// Add Windows-specific tests
|
||||
if runtime.GOOS == "windows" {
|
||||
tests = append(tests, []struct {
|
||||
username string
|
||||
expected bool
|
||||
description string
|
||||
}{
|
||||
{
|
||||
username: "Administrator",
|
||||
expected: true,
|
||||
description: "Administrator should be considered privileged on Windows",
|
||||
},
|
||||
{
|
||||
username: "administrator",
|
||||
expected: true,
|
||||
description: "administrator should be considered privileged on Windows (case insensitive)",
|
||||
},
|
||||
}...)
|
||||
} else {
|
||||
// On non-Windows systems, Administrator should not be privileged
|
||||
tests = append(tests, []struct {
|
||||
username string
|
||||
expected bool
|
||||
description string
|
||||
}{
|
||||
{
|
||||
username: "Administrator",
|
||||
expected: false,
|
||||
description: "Administrator should not be privileged on non-Windows systems",
|
||||
},
|
||||
}...)
|
||||
{
|
||||
username: "Administrator",
|
||||
expected: false,
|
||||
description: "Administrator should not be privileged on non-Windows systems",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.description, func(t *testing.T) {
|
||||
result := isPrivilegedUsername(tt.username)
|
||||
result := isPrivilegedOrUnknown(tt.username)
|
||||
assert.Equal(t, tt.expected, result, tt.description)
|
||||
})
|
||||
}
|
||||
|
||||
@@ -17,7 +17,7 @@ import (
|
||||
// createSftpCommand creates a Windows SFTP command with user switching.
|
||||
// The caller must close the returned token handle after starting the process.
|
||||
func (s *Server) createSftpCommand(targetUser *user.User, sess ssh.Session) (*exec.Cmd, windows.Token, error) {
|
||||
username, domain := s.parseUsername(targetUser.Username)
|
||||
username, domain := parseUsername(targetUser.Username)
|
||||
|
||||
netbirdPath, err := os.Executable()
|
||||
if err != nil {
|
||||
|
||||
@@ -16,11 +16,6 @@ var (
|
||||
ErrPrivilegedUserSwitch = errors.New("cannot switch to privileged user - current user lacks required privileges")
|
||||
)
|
||||
|
||||
// isPlatformUnix returns true for Unix-like platforms (Linux, macOS, etc.)
|
||||
func isPlatformUnix() bool {
|
||||
return getCurrentOS() != "windows"
|
||||
}
|
||||
|
||||
// Dependency injection variables for testing - allows mocking dynamic runtime checks
|
||||
var (
|
||||
getCurrentUser = currentUserWithGetent
|
||||
@@ -29,6 +24,9 @@ var (
|
||||
getIsProcessPrivileged = isCurrentProcessPrivileged
|
||||
|
||||
getEuid = os.Geteuid
|
||||
|
||||
getProcessElevated = isProcessElevated
|
||||
getWindowsAccountPrivilegedOrUnknown = isWindowsAccountPrivilegedOrUnknown
|
||||
)
|
||||
|
||||
const (
|
||||
@@ -65,6 +63,13 @@ type PrivilegeCheckResult struct {
|
||||
RequiresUserSwitching bool
|
||||
}
|
||||
|
||||
// privilegeCheckContext holds all context needed for privilege checking
|
||||
type privilegeCheckContext struct {
|
||||
currentUser *user.User
|
||||
currentUserPrivileged bool
|
||||
allowRoot bool
|
||||
}
|
||||
|
||||
// CheckPrivileges performs comprehensive privilege checking for all SSH features.
|
||||
// This is the single source of truth for privilege decisions across the SSH server.
|
||||
func (s *Server) CheckPrivileges(req PrivilegeCheckRequest) PrivilegeCheckResult {
|
||||
@@ -75,7 +80,7 @@ func (s *Server) CheckPrivileges(req PrivilegeCheckRequest) PrivilegeCheckResult
|
||||
|
||||
// Handle empty username case - but still check root access controls
|
||||
if req.RequestedUsername == "" {
|
||||
if isPrivilegedUsername(context.currentUser.Username) && !context.allowRoot {
|
||||
if isPrivilegedOrUnknown(context.currentUser.Username) && !context.allowRoot {
|
||||
return PrivilegeCheckResult{
|
||||
Allowed: false,
|
||||
Error: &PrivilegedUserError{Username: context.currentUser.Username},
|
||||
@@ -135,7 +140,7 @@ func (s *Server) checkUserRequest(ctx *privilegeCheckContext, req PrivilegeCheck
|
||||
|
||||
needsUserSwitching := !isSameResolvedUser(resolvedUser, ctx.currentUser)
|
||||
|
||||
if isPrivilegedUsername(resolvedUser.Username) && !ctx.allowRoot {
|
||||
if isPrivilegedOrUnknown(resolvedUser.Username) && !ctx.allowRoot {
|
||||
return PrivilegeCheckResult{
|
||||
Allowed: false,
|
||||
Error: &PrivilegedUserError{Username: resolvedUser.Username},
|
||||
@@ -175,6 +180,42 @@ func (s *Server) resolveRequestedUser(requestedUsername string) (*user.User, err
|
||||
return u, nil
|
||||
}
|
||||
|
||||
// SetAllowRootLogin configures root login access
|
||||
func (s *Server) SetAllowRootLogin(allow bool) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.allowRootLogin = allow
|
||||
}
|
||||
|
||||
// userNameLookup performs user lookup with root login permission check
|
||||
func (s *Server) userNameLookup(username string) (*user.User, error) {
|
||||
result, err := s.userPrivilegeCheck(username)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return result.User, nil
|
||||
}
|
||||
|
||||
// userPrivilegeCheck performs user lookup with full privilege check result
|
||||
func (s *Server) userPrivilegeCheck(username string) (PrivilegeCheckResult, error) {
|
||||
result := s.CheckPrivileges(PrivilegeCheckRequest{
|
||||
RequestedUsername: username,
|
||||
FeatureSupportsUserSwitch: true,
|
||||
FeatureName: FeatureSSHLogin,
|
||||
})
|
||||
|
||||
if !result.Allowed {
|
||||
return result, result.Error
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// isPlatformUnix returns true for Unix-like platforms (Linux, macOS, etc.)
|
||||
func isPlatformUnix() bool {
|
||||
return getCurrentOS() != "windows"
|
||||
}
|
||||
|
||||
// isSameResolvedUser compares two resolved user identities
|
||||
func isSameResolvedUser(user1, user2 *user.User) bool {
|
||||
if user1 == nil || user2 == nil {
|
||||
@@ -183,13 +224,6 @@ func isSameResolvedUser(user1, user2 *user.User) bool {
|
||||
return user1.Uid == user2.Uid
|
||||
}
|
||||
|
||||
// privilegeCheckContext holds all context needed for privilege checking
|
||||
type privilegeCheckContext struct {
|
||||
currentUser *user.User
|
||||
currentUserPrivileged bool
|
||||
allowRoot bool
|
||||
}
|
||||
|
||||
// isSameUser checks if two usernames refer to the same user
|
||||
// SECURITY: This function must be conservative - it should only return true
|
||||
// when we're certain both usernames refer to the exact same user identity
|
||||
@@ -253,159 +287,30 @@ func isWindowsSameUser(requestedUsername, currentUsername string) bool {
|
||||
return strings.EqualFold(reqDomain, curDomain)
|
||||
}
|
||||
|
||||
// SetAllowRootLogin configures root login access
|
||||
func (s *Server) SetAllowRootLogin(allow bool) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.allowRootLogin = allow
|
||||
}
|
||||
|
||||
// userNameLookup performs user lookup with root login permission check
|
||||
func (s *Server) userNameLookup(username string) (*user.User, error) {
|
||||
result := s.CheckPrivileges(PrivilegeCheckRequest{
|
||||
RequestedUsername: username,
|
||||
FeatureSupportsUserSwitch: true,
|
||||
FeatureName: FeatureSSHLogin,
|
||||
})
|
||||
|
||||
if !result.Allowed {
|
||||
return nil, result.Error
|
||||
}
|
||||
|
||||
return result.User, nil
|
||||
}
|
||||
|
||||
// userPrivilegeCheck performs user lookup with full privilege check result
|
||||
func (s *Server) userPrivilegeCheck(username string) (PrivilegeCheckResult, error) {
|
||||
result := s.CheckPrivileges(PrivilegeCheckRequest{
|
||||
RequestedUsername: username,
|
||||
FeatureSupportsUserSwitch: true,
|
||||
FeatureName: FeatureSSHLogin,
|
||||
})
|
||||
|
||||
if !result.Allowed {
|
||||
return result, result.Error
|
||||
}
|
||||
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// isPrivilegedUsername checks if the given username represents a privileged user across platforms.
|
||||
// On Unix: root
|
||||
// On Windows: Administrator, SYSTEM (case-insensitive)
|
||||
// Handles domain-qualified usernames like "DOMAIN\Administrator" or "user@domain.com"
|
||||
func isPrivilegedUsername(username string) bool {
|
||||
// isPrivilegedOrUnknown reports whether the given username represents a
|
||||
// privileged user, or on Windows an account whose privilege could not be
|
||||
// determined.
|
||||
// On Unix: root.
|
||||
// On Windows: well-known service accounts, built-in Administrator accounts,
|
||||
// and members of the local Administrators group; handles domain-qualified
|
||||
// usernames like "DOMAIN\user" or "user@domain.com". An account that cannot be
|
||||
// resolved or evaluated is reported as privileged.
|
||||
//
|
||||
// Use this to refuse privileged accounts, never to grant them anything: the
|
||||
// undetermined case is safe for a refusal and unsafe for a grant.
|
||||
func isPrivilegedOrUnknown(username string) bool {
|
||||
if getCurrentOS() != "windows" {
|
||||
return username == "root"
|
||||
}
|
||||
|
||||
bareUsername := username
|
||||
// Handle Windows domain format: DOMAIN\username
|
||||
if idx := strings.LastIndex(username, `\`); idx != -1 {
|
||||
bareUsername = username[idx+1:]
|
||||
}
|
||||
// Handle email-style format: username@domain.com
|
||||
if idx := strings.Index(bareUsername, "@"); idx != -1 {
|
||||
bareUsername = bareUsername[:idx]
|
||||
}
|
||||
|
||||
return isWindowsPrivilegedUser(bareUsername)
|
||||
}
|
||||
|
||||
// isWindowsPrivilegedUser checks if a bare username (domain already stripped) represents a Windows privileged account
|
||||
func isWindowsPrivilegedUser(bareUsername string) bool {
|
||||
// common privileged usernames (case insensitive)
|
||||
privilegedNames := []string{
|
||||
"administrator",
|
||||
"admin",
|
||||
"root",
|
||||
"system",
|
||||
"localsystem",
|
||||
"networkservice",
|
||||
"localservice",
|
||||
}
|
||||
|
||||
usernameLower := strings.ToLower(bareUsername)
|
||||
for _, privilegedName := range privilegedNames {
|
||||
if usernameLower == privilegedName {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
// computer accounts (ending with $) are not privileged by themselves
|
||||
// They only gain privileges through group membership or specific SIDs
|
||||
|
||||
if targetUser, err := lookupUser(bareUsername); err == nil {
|
||||
return isWindowsPrivilegedSID(targetUser.Uid)
|
||||
}
|
||||
|
||||
return false
|
||||
}
|
||||
|
||||
// isWindowsPrivilegedSID checks if a Windows SID represents a privileged account
|
||||
func isWindowsPrivilegedSID(sid string) bool {
|
||||
privilegedSIDs := []string{
|
||||
"S-1-5-18", // Local System (SYSTEM)
|
||||
"S-1-5-19", // Local Service (NT AUTHORITY\LOCAL SERVICE)
|
||||
"S-1-5-20", // Network Service (NT AUTHORITY\NETWORK SERVICE)
|
||||
"S-1-5-32-544", // Administrators group (BUILTIN\Administrators)
|
||||
"S-1-5-500", // Built-in Administrator account (local machine RID 500)
|
||||
}
|
||||
|
||||
for _, privilegedSID := range privilegedSIDs {
|
||||
if sid == privilegedSID {
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
// Check for domain administrator accounts (RID 500 in any domain)
|
||||
// Format: S-1-5-21-domain-domain-domain-500
|
||||
// This is reliable as RID 500 is reserved for the domain Administrator account
|
||||
if strings.HasPrefix(sid, "S-1-5-21-") && strings.HasSuffix(sid, "-500") {
|
||||
return true
|
||||
}
|
||||
|
||||
// Check for other well-known privileged RIDs in domain contexts
|
||||
// RID 512 = Domain Admins group, RID 516 = Domain Controllers group
|
||||
if strings.HasPrefix(sid, "S-1-5-21-") {
|
||||
if strings.HasSuffix(sid, "-512") || // Domain Admins group
|
||||
strings.HasSuffix(sid, "-516") || // Domain Controllers group
|
||||
strings.HasSuffix(sid, "-519") { // Enterprise Admins group
|
||||
return true
|
||||
}
|
||||
}
|
||||
|
||||
return false
|
||||
return getWindowsAccountPrivilegedOrUnknown(username)
|
||||
}
|
||||
|
||||
// isCurrentProcessPrivileged checks if the current process is running with elevated privileges.
|
||||
// On Unix systems, this means running as root (UID 0).
|
||||
// On Windows, this means running as Administrator or SYSTEM.
|
||||
// On Windows, this means the process token is elevated (administrators, SYSTEM).
|
||||
func isCurrentProcessPrivileged() bool {
|
||||
if getCurrentOS() == "windows" {
|
||||
return isWindowsElevated()
|
||||
return getProcessElevated()
|
||||
}
|
||||
return getEuid() == 0
|
||||
}
|
||||
|
||||
// isWindowsElevated checks if the current process is running with elevated privileges on Windows
|
||||
func isWindowsElevated() bool {
|
||||
currentUser, err := getCurrentUser()
|
||||
if err != nil {
|
||||
log.Errorf("failed to get current user for privilege check, assuming non-privileged: %v", err)
|
||||
return false
|
||||
}
|
||||
|
||||
if isWindowsPrivilegedSID(currentUser.Uid) {
|
||||
log.Debugf("Windows user switching supported: running as privileged SID %s", currentUser.Uid)
|
||||
return true
|
||||
}
|
||||
|
||||
if isPrivilegedUsername(currentUser.Username) {
|
||||
log.Debugf("Windows user switching supported: running as privileged username %s", currentUser.Username)
|
||||
return true
|
||||
}
|
||||
|
||||
log.Debugf("Windows user switching not supported: not running as privileged user (current: %s)", currentUser.Uid)
|
||||
return false
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ import (
|
||||
"errors"
|
||||
"os/user"
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
@@ -27,8 +28,8 @@ func setupTestDependencies(currentUser *user.User, currentUserErr error, os stri
|
||||
originalLookupUser := lookupUser
|
||||
originalGetCurrentOS := getCurrentOS
|
||||
originalGetEuid := getEuid
|
||||
|
||||
// Reset caches to ensure clean test state
|
||||
originalGetProcessElevated := getProcessElevated
|
||||
originalGetWindowsAccountPrivilegedOrUnknown := getWindowsAccountPrivilegedOrUnknown
|
||||
|
||||
// Set test values - inject platform dependencies
|
||||
getCurrentUser = func() (*user.User, error) {
|
||||
@@ -53,16 +54,31 @@ func setupTestDependencies(currentUser *user.User, currentUserErr error, os stri
|
||||
return euid
|
||||
}
|
||||
|
||||
// Mock privilege detection based on the test user
|
||||
getIsProcessPrivileged = func() bool {
|
||||
// Simulate the Windows token elevation check based on the fixture user:
|
||||
// the built-in Administrator (RID 500) and SYSTEM run elevated.
|
||||
getProcessElevated = func() bool {
|
||||
if currentUser == nil {
|
||||
return false
|
||||
}
|
||||
// Check both username and SID for Windows systems
|
||||
if os == "windows" && isWindowsPrivilegedSID(currentUser.Uid) {
|
||||
return currentUser.Uid == "S-1-5-18" || strings.HasSuffix(currentUser.Uid, "-500")
|
||||
}
|
||||
|
||||
// Simulate the Windows account classifier for the fixture accounts.
|
||||
// "root" does not exist on Windows; the real classifier fails closed on
|
||||
// unresolvable accounts, so it counts as privileged here too.
|
||||
getWindowsAccountPrivilegedOrUnknown = func(username string) bool {
|
||||
bare := username
|
||||
if idx := strings.LastIndex(bare, `\`); idx != -1 {
|
||||
bare = bare[idx+1:]
|
||||
}
|
||||
if idx := strings.Index(bare, "@"); idx != -1 {
|
||||
bare = bare[:idx]
|
||||
}
|
||||
switch strings.ToLower(bare) {
|
||||
case "administrator", "system", "root":
|
||||
return true
|
||||
}
|
||||
return isPrivilegedUsername(currentUser.Username)
|
||||
return false
|
||||
}
|
||||
|
||||
// Return cleanup function
|
||||
@@ -71,10 +87,8 @@ func setupTestDependencies(currentUser *user.User, currentUserErr error, os stri
|
||||
lookupUser = originalLookupUser
|
||||
getCurrentOS = originalGetCurrentOS
|
||||
getEuid = originalGetEuid
|
||||
|
||||
getIsProcessPrivileged = isCurrentProcessPrivileged
|
||||
|
||||
// Reset caches after test
|
||||
getProcessElevated = originalGetProcessElevated
|
||||
getWindowsAccountPrivilegedOrUnknown = originalGetWindowsAccountPrivilegedOrUnknown
|
||||
}
|
||||
}
|
||||
|
||||
@@ -421,6 +435,9 @@ func TestUsedFallback_MeansNoPrivilegeDropping(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestPrivilegedUsernameDetection(t *testing.T) {
|
||||
// Windows classification is syscall-backed (SID resolution, group
|
||||
// membership) and is covered by privileges_windows_test.go; here only the
|
||||
// Unix logic and the platform dispatch are exercised.
|
||||
tests := []struct {
|
||||
name string
|
||||
username string
|
||||
@@ -432,25 +449,9 @@ func TestPrivilegedUsernameDetection(t *testing.T) {
|
||||
{"unix_regular_user", "alice", "linux", false},
|
||||
{"unix_root_capital", "Root", "linux", false}, // Case-sensitive
|
||||
|
||||
// Windows tests
|
||||
// Windows dispatch to the (mocked) account classifier
|
||||
{"windows_administrator", "Administrator", "windows", true},
|
||||
{"windows_system", "SYSTEM", "windows", true},
|
||||
{"windows_admin", "admin", "windows", true},
|
||||
{"windows_admin_lowercase", "administrator", "windows", true}, // Case-insensitive
|
||||
{"windows_domain_admin", "DOMAIN\\Administrator", "windows", true},
|
||||
{"windows_email_admin", "admin@domain.com", "windows", true},
|
||||
{"windows_regular_user", "alice", "windows", false},
|
||||
{"windows_domain_user", "DOMAIN\\alice", "windows", false},
|
||||
{"windows_localsystem", "localsystem", "windows", true},
|
||||
{"windows_networkservice", "networkservice", "windows", true},
|
||||
{"windows_localservice", "localservice", "windows", true},
|
||||
|
||||
// Computer accounts (these depend on current user context in real implementation)
|
||||
{"windows_computer_account", "WIN2K19-C2$", "windows", false}, // Computer account by itself not privileged
|
||||
{"windows_domain_computer", "DOMAIN\\COMPUTER$", "windows", false}, // Domain computer account
|
||||
|
||||
// Cross-platform
|
||||
{"root_on_windows", "root", "windows", true}, // Root should be privileged everywhere
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
@@ -459,50 +460,8 @@ func TestPrivilegedUsernameDetection(t *testing.T) {
|
||||
cleanup := setupTestDependencies(nil, nil, tt.platform, 1000, nil, nil)
|
||||
defer cleanup()
|
||||
|
||||
result := isPrivilegedUsername(tt.username)
|
||||
assert.Equal(t, tt.privileged, result)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWindowsPrivilegedSIDDetection(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
sid string
|
||||
privileged bool
|
||||
description string
|
||||
}{
|
||||
// Well-known system accounts
|
||||
{"system_account", "S-1-5-18", true, "Local System (SYSTEM)"},
|
||||
{"local_service", "S-1-5-19", true, "Local Service"},
|
||||
{"network_service", "S-1-5-20", true, "Network Service"},
|
||||
{"administrators_group", "S-1-5-32-544", true, "Administrators group"},
|
||||
{"builtin_administrator", "S-1-5-500", true, "Built-in Administrator"},
|
||||
|
||||
// Domain accounts
|
||||
{"domain_administrator", "S-1-5-21-1234567890-1234567890-1234567890-500", true, "Domain Administrator (RID 500)"},
|
||||
{"domain_admins_group", "S-1-5-21-1234567890-1234567890-1234567890-512", true, "Domain Admins group"},
|
||||
{"domain_controllers_group", "S-1-5-21-1234567890-1234567890-1234567890-516", true, "Domain Controllers group"},
|
||||
{"enterprise_admins_group", "S-1-5-21-1234567890-1234567890-1234567890-519", true, "Enterprise Admins group"},
|
||||
|
||||
// Regular users
|
||||
{"regular_user", "S-1-5-21-1234567890-1234567890-1234567890-1001", false, "Regular domain user"},
|
||||
{"another_regular_user", "S-1-5-21-1234567890-1234567890-1234567890-1234", false, "Another regular user"},
|
||||
{"local_user", "S-1-5-21-1234567890-1234567890-1234567890-1000", false, "Local regular user"},
|
||||
|
||||
// Groups that are not privileged
|
||||
{"domain_users", "S-1-5-21-1234567890-1234567890-1234567890-513", false, "Domain Users group"},
|
||||
{"power_users", "S-1-5-32-547", false, "Power Users group"},
|
||||
|
||||
// Invalid SIDs
|
||||
{"malformed_sid", "S-1-5-invalid", false, "Malformed SID"},
|
||||
{"empty_sid", "", false, "Empty SID"},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
result := isWindowsPrivilegedSID(tt.sid)
|
||||
assert.Equal(t, tt.privileged, result, "Failed for %s: %s", tt.description, tt.sid)
|
||||
result := isPrivilegedOrUnknown(tt.username)
|
||||
assert.Equal(t, tt.privileged, result, "privilege classification for %s on %s", tt.username, tt.platform)
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
@@ -91,7 +91,7 @@ func validateUsernameFormat(username string) error {
|
||||
func (s *Server) createExecutorCommand(logger *log.Entry, session ssh.Session, localUser *user.User, hasPty bool) (*exec.Cmd, func(), error) {
|
||||
logger.Debugf("creating Windows executor command for user %s (Pty: %v)", localUser.Username, hasPty)
|
||||
|
||||
username, _ := s.parseUsername(localUser.Username)
|
||||
username, _ := parseUsername(localUser.Username)
|
||||
if err := validateUsername(username); err != nil {
|
||||
return nil, nil, fmt.Errorf("invalid username %q: %w", username, err)
|
||||
}
|
||||
@@ -102,7 +102,7 @@ func (s *Server) createExecutorCommand(logger *log.Entry, session ssh.Session, l
|
||||
// createUserSwitchCommand creates a command with Windows user switching.
|
||||
// Returns the command and a cleanup function that must be called after starting the process.
|
||||
func (s *Server) createUserSwitchCommand(logger *log.Entry, session ssh.Session, localUser *user.User) (*exec.Cmd, func(), error) {
|
||||
username, domain := s.parseUsername(localUser.Username)
|
||||
username, domain := parseUsername(localUser.Username)
|
||||
|
||||
shell := getUserShell(localUser.Uid)
|
||||
|
||||
@@ -138,7 +138,7 @@ func (s *Server) createUserSwitchCommand(logger *log.Entry, session ssh.Session,
|
||||
}
|
||||
|
||||
// parseUsername extracts username and domain from a Windows username
|
||||
func (s *Server) parseUsername(fullUsername string) (username, domain string) {
|
||||
func parseUsername(fullUsername string) (username, domain string) {
|
||||
// Handle DOMAIN\username format
|
||||
if idx := strings.LastIndex(fullUsername, `\`); idx != -1 {
|
||||
domain = fullUsername[:idx]
|
||||
|
||||
@@ -2,9 +2,24 @@
|
||||
|
||||
A short brief for translating the desktop UI — for any translator, human or AI agent (*"you"* = whoever's translating).
|
||||
|
||||
**Drive an agent with:** *"Read `i18n/TRANSLATING.md` and translate the UI to Russian"* — or *"…and review the existing German translation."*
|
||||
**Translations are managed on Crowdin: <https://crowdin.com/project/netbird>.** Join the project, pick your language, and translate in the editor. Each string carries a context note (the `description` from the source file) telling you what it is and where it shows up, and the project's glossary, style guide, and QA checks mirror this document.
|
||||
|
||||
> 💡 **The one habit that matters most:** read each key's `description` before translating it. Labels are terse and ambiguous on their own; the `description` tells you what the string is, where it shows up, what to keep verbatim, and what it actually means.
|
||||
> 💡 **The one habit that matters most:** read each string's context before translating it. Labels are terse and ambiguous on their own; the context tells you what the string is, where it shows up, what to keep verbatim, and what it actually means.
|
||||
|
||||
---
|
||||
|
||||
## How contributions flow
|
||||
|
||||
```text
|
||||
i18n/locales/en/common.json ──sync──▶ Crowdin ──service PR──▶ i18n/locales/<code>/common.json
|
||||
```
|
||||
|
||||
- `i18n/locales/en/common.json` is the source of truth. New and changed strings sync to Crowdin automatically (see `crowdin.yml` in the repository root).
|
||||
- Crowdin opens and updates a service pull request with the translated bundles, keeping the source's file shape and key order. Keys nobody has translated yet are left out of the export; the app falls back to English for them at runtime. Maintainers review and merge that PR.
|
||||
- Don't hand-edit `i18n/locales/<code>/common.json` in your own PRs: the next sync would conflict with or overwrite your changes. Translate on Crowdin instead.
|
||||
- Missing your language? Request it on the Crowdin project page or in a [GitHub discussion](https://github.com/netbirdio/netbird/discussions). When a language first ships, a maintainer adds its row to `i18n/locales/_index.json` with `code`, `displayName` (the native name), and `englishName`, which puts it in the app's language picker.
|
||||
|
||||
**Prefer translating with an AI agent?** That still works: drive it with *"Read `i18n/TRANSLATING.md` and translate the UI to Russian"* as before, but deliver the result to Crowdin instead of a pull request. Download your language's file from the Crowdin editor, let the agent translate it, and upload it back (the editor's offline translation flow). Crowdin runs its QA checks on upload, and the next service PR carries the strings into the repo.
|
||||
|
||||
---
|
||||
|
||||
@@ -30,25 +45,6 @@ A **business zero-trust VPN** — an encrypted **overlay mesh** between a compan
|
||||
|
||||
---
|
||||
|
||||
## The files
|
||||
|
||||
```
|
||||
i18n/locales/_index.json shipped-language list
|
||||
i18n/locales/en/common.json source of truth — message + description
|
||||
i18n/locales/<code>/common.json a target — message only
|
||||
```
|
||||
|
||||
Chrome-extension JSON, each key → `{ "message", "description" }`. You translate the **`message`**.
|
||||
|
||||
| ✅ Do | ❌ Don't |
|
||||
|---|---|
|
||||
| Keep **every key** from `en`, in the same order | Translate, rename, reorder, drop, or add keys (they're identifiers; the set grows over time) |
|
||||
| Put **only `message`** in target bundles | Copy `description` into a target bundle |
|
||||
| Give every key a non-empty `message` | Leave keys missing or empty |
|
||||
| Save valid UTF-8 JSON, no BOM | Add trailing commas or break the JSON |
|
||||
|
||||
---
|
||||
|
||||
## Hard rules — get these exactly right
|
||||
|
||||
These are the usual ways a translation *breaks the app*, not just reads oddly.
|
||||
@@ -58,7 +54,7 @@ These are the usual ways a translation *breaks the app*, not just reads oddly.
|
||||
| Copy `{placeholders}` verbatim — `{version}`, `{count}`, `{name}`… | Translate the word inside the braces (`{verbleibend}` breaks it) |
|
||||
| Reposition a placeholder so the sentence flows | Drop or duplicate a placeholder |
|
||||
| Preserve every `\n`, leading/trailing space, and trailing `...` | Trim "invisible" spaces or the `...` (they're load-bearing) |
|
||||
| Keep `®` in WireGuard® and quotes around `{name}` | Strip punctuation the description flags |
|
||||
| Keep `®` in WireGuard® and quotes around `{name}` | Strip punctuation the context flags |
|
||||
|
||||
**Plurals:** the app has only a *one / other* split — the singular key fires only when `count == 1`; the `{count}` key covers everything else (0, 2, 5, 100…). Languages with more than two forms (ru, pl, uk) can't be fully correct here — use the form that fits the widest range (Russian genitive plural: `минут` / `часов` / `дней`). Don't invent extra keys or cram multiple forms into one string. When no single form fits every value — a unit label after a number field, say — reach for a number-agnostic form (an abbreviation, or wording that reads the same for 1 and 100) instead of forcing a plural the *one / other* split can't supply.
|
||||
|
||||
@@ -78,13 +74,15 @@ When a brand sits beside a common noun, keep its exact spelling but join them th
|
||||
|
||||
> **Use the word that language's IT users actually say.** Translate when a natural, common term exists; keep the English term *only* when the literal translation would be awkward or no one in that field really uses it.
|
||||
|
||||
Apply each term **consistently** — same English term → same translation everywhere — and keep a term once you've settled it. Whether a term stays English or takes a native word is **language-dependent**: a technical loanword (e.g. *Daemon*, *Handshake*) often stays, an everyday word (e.g. *Latency*, *Public key*) usually localizes, and some (*Exit Node*, *Peer*) go either way depending on the language. Decide per term with the rule above — a foreign origin alone is no reason to keep English. **Your main reference is the existing bundles:** match how a term was already rendered for your language rather than re-deciding it.
|
||||
Apply each term **consistently** — same English term → same translation everywhere — and keep a term once you've settled it. Whether a term stays English or takes a native word is **language-dependent**: a technical loanword (e.g. *Daemon*, *Handshake*) often stays, an everyday word (e.g. *Latency*, *Public key*) usually localizes, and some (*Exit Node*, *Peer*) go either way depending on the language. Decide per term with the rule above — a foreign origin alone is no reason to keep English. **Your main reference is the existing translation:** match how a term was already rendered for your language rather than re-deciding it.
|
||||
|
||||
Two checks before you commit a term:
|
||||
|
||||
- **Prefer established localized wording.** If a widely used tool in this space (for example WireGuard) ships your language, its wording for a shared term such as *handshake* is what users already expect — look at the translated app, not just English docs. For generic UI verbs and formal address, follow your OS vendor's style guide (Microsoft / Apple / Google).
|
||||
- **Watch for false friends.** A literal translation can collide with a *different* established term in your field — confirm your word doesn't already mean something else in this domain before using it.
|
||||
|
||||
These tiers are mirrored in the Crowdin project glossary, so the editor highlights them inline. When you settle a new Tier C term for your language, add its translation to the glossary entry so it sticks for everyone who comes after you.
|
||||
|
||||
---
|
||||
|
||||
## Style
|
||||
@@ -98,7 +96,7 @@ Two checks before you commit a term:
|
||||
|
||||
Where it reads naturally, aim to keep each string **roughly the same length** as the English — the UI is tight and over-long strings can wrap or truncate. It's a soft preference, not a rule: if your language simply needs more words, use them.
|
||||
|
||||
A few habits that keep a bundle reading like one product rather than a word-for-word port:
|
||||
A few habits that keep a translation reading like one product rather than a word-for-word port:
|
||||
|
||||
- **Translate meaning, not words.** Render what a string *does*. An idiom or an awkward source phrase should become natural in your language, not a literal calque.
|
||||
- **Keep one voice within a family.** Sibling strings — the connection states, every settings *help* caption, every "… Failed" title — should share a grammatical form. If one member sounds wrong in that form, re-voice the whole family rather than leave one odd sibling.
|
||||
@@ -107,27 +105,26 @@ A few habits that keep a bundle reading like one product rather than a word-for-
|
||||
|
||||
---
|
||||
|
||||
## Procedure
|
||||
## Reviewing a language
|
||||
|
||||
**New language** — read `en/common.json` *with* descriptions → settle your Tier C terms → write `i18n/locales/<code>/common.json` (same keys and order as `en`, `message` only, placeholders & brands preserved) → add a row to `_index.json` (`{"code","displayName"` = native name`,"englishName"}`) → run the QA list. Use the locale-code style the existing entries use (e.g. `fr`, `pt`, `zh-CN`).
|
||||
**On Crowdin:** proofread in the editor — context, glossary highlights, and QA flags sit inline next to each string.
|
||||
|
||||
**Review (de / hu / …)** — read source and target side by side; for each key check glossary conformance (e.g. de `Exit-Node` → `Exit Node`, hu `Kilépő csomópont` → `Exit Node`), placeholder/`\n` integrity, consistency, tone, and that the meaning matches the English `description`. Fix in place, then report what you changed (especially term standardizations) so a native speaker can sanity-check.
|
||||
**In the repo** — e.g. driving an AI agent with *"Read `i18n/TRANSLATING.md` and review the existing German translation"* — read source and target side by side; for each key check glossary conformance (e.g. de `Exit-Node` → `Exit Node`, hu `Kilépő csomópont` → `Exit Node`), placeholder/`\n` integrity, consistency, tone, and that the meaning matches the English `description`. Report what you found, and apply the fixes **on Crowdin** — direct edits to the locale files are overwritten by the next sync.
|
||||
|
||||
---
|
||||
|
||||
## QA before you finish
|
||||
|
||||
- [ ] Valid JSON · **every `en` key** present, same order · **no `description`** fields
|
||||
- [ ] Every `{placeholder}`, `\n`, and intentional space preserved · `...` / `… Failed` / `{name}` quotes kept
|
||||
- [ ] Tier A/B left intact · Tier C applied consistently (and matching the existing bundle for your language)
|
||||
- [ ] Tier A/B left intact · Tier C applied consistently (and matching the existing translation for your language)
|
||||
- [ ] Buttons & tray short · locale punctuation and capitalization applied
|
||||
- [ ] New language added to `_index.json`
|
||||
- [ ] Crowdin QA flags resolved (variables, glossary terms, punctuation)
|
||||
- [ ] **Tested in the running app** ↓
|
||||
|
||||
---
|
||||
|
||||
## Test it in the app
|
||||
|
||||
A bundle can pass every check above and still read wrong on screen. **Run the app, switch to your language, and click through the real surfaces** — tray menu, main window, every Settings tab, the dialogs. Watch for text overflow or truncation, labels that are technically right but wrong *for what the control does*, leaked placeholders, and terms that drift between screens.
|
||||
A translation can pass every check above and still read wrong on screen. **Run the app, switch to your language, and click through the real surfaces** — tray menu, main window, every Settings tab, the dialogs. Watch for text overflow or truncation, labels that are technically right but wrong *for what the control does*, leaked placeholders, and terms that drift between screens.
|
||||
|
||||
How to run the app and switch language: see the project README. Can't run it (e.g. a headless agent)? Say so in your summary — don't silently skip this step.
|
||||
|
||||
11
crowdin.yml
Normal file
11
crowdin.yml
Normal file
@@ -0,0 +1,11 @@
|
||||
skip_untranslated_strings: true
|
||||
skip_untranslated_files: true
|
||||
import_eq_suggestions: true
|
||||
|
||||
files:
|
||||
- source: /client/ui/i18n/locales/en/common.json
|
||||
translation: /client/ui/i18n/locales/%two_letters_code%/common.json
|
||||
type: chrome
|
||||
languages_mapping:
|
||||
two_letters_code:
|
||||
zh-CN: zh-CN
|
||||
@@ -438,14 +438,10 @@ func TestProvidersMatrix(t *testing.T) {
|
||||
// Create every provider, all enabled, each with a unique model string so the
|
||||
// proxy's connect-time snapshot carries them all and model→provider routing
|
||||
// is unambiguous (provider toggles after connect don't reconcile to the
|
||||
// proxy, so we enable everything up front). The first create bootstraps the
|
||||
// cluster.
|
||||
// proxy, so we enable everything up front).
|
||||
ids := make([]string, 0, len(matrix))
|
||||
for i, pc := range matrix {
|
||||
for _, pc := range matrix {
|
||||
req := providerRequest(pc)
|
||||
if i == 0 {
|
||||
req.BootstrapCluster = ptr(harness.AgentNetworkCluster)
|
||||
}
|
||||
prov, perr := srv.CreateProvider(ctx, req)
|
||||
require.NoError(t, perr, "create provider %s", pc.name)
|
||||
ids = append(ids, prov.Id)
|
||||
|
||||
@@ -82,13 +82,12 @@ func provisionPricedProvider(t *testing.T, ctx context.Context, name string, mod
|
||||
// need NOT be in the catalog — the operator names it and prices it here.
|
||||
dummyKey := "sk-price-e2e"
|
||||
prov, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{
|
||||
Name: name,
|
||||
ProviderId: "openai_api",
|
||||
UpstreamUrl: vllm.URL,
|
||||
ApiKey: &dummyKey,
|
||||
Enabled: ptr(true),
|
||||
BootstrapCluster: ptr(harness.AgentNetworkCluster),
|
||||
Models: &models,
|
||||
Name: name,
|
||||
ProviderId: "openai_api",
|
||||
UpstreamUrl: vllm.URL,
|
||||
ApiKey: &dummyKey,
|
||||
Enabled: ptr(true),
|
||||
Models: &models,
|
||||
})
|
||||
require.NoError(t, err, "create provider")
|
||||
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), prov.Id) })
|
||||
|
||||
@@ -113,15 +113,14 @@ func runPathRoutedGuardrailCase(t *testing.T, tc pathRoutedGuardrailCase) {
|
||||
|
||||
// Catch-all provider (no models) so the router forwards any model; a static
|
||||
// bearer key means the router injects a static auth header instead of minting
|
||||
// a GCP token. Bootstraps the cluster if it isn't already.
|
||||
// a GCP token.
|
||||
staticKey := "static-e2e-token"
|
||||
prov, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{
|
||||
Name: tc.name,
|
||||
ProviderId: tc.catalogID,
|
||||
UpstreamUrl: vllm.URL,
|
||||
ApiKey: &staticKey,
|
||||
Enabled: ptr(true),
|
||||
BootstrapCluster: ptr(harness.AgentNetworkCluster),
|
||||
Name: tc.name,
|
||||
ProviderId: tc.catalogID,
|
||||
UpstreamUrl: vllm.URL,
|
||||
ApiKey: &staticKey,
|
||||
Enabled: ptr(true),
|
||||
})
|
||||
require.NoError(t, err, "create %s provider", tc.name)
|
||||
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), prov.Id) })
|
||||
|
||||
@@ -73,7 +73,6 @@ func TestGuardrailGroupSwitchTakesEffectAfterTTL(t *testing.T) {
|
||||
{Id: modelA, InputPer1k: 0.001, OutputPer1k: 0.001},
|
||||
{Id: modelB, InputPer1k: 0.001, OutputPer1k: 0.001},
|
||||
},
|
||||
BootstrapCluster: ptr(harness.AgentNetworkCluster),
|
||||
})
|
||||
require.NoError(t, err, "create provider")
|
||||
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), prov.Id) })
|
||||
|
||||
@@ -61,15 +61,14 @@ func TestGuardrailMultiPolicyModelAllowlist(t *testing.T) {
|
||||
}
|
||||
|
||||
// pRestricted declares the two guardrailed models so routing is deterministic
|
||||
// (model -> provider). Created first, so it carries the bootstrap cluster.
|
||||
// (model -> provider).
|
||||
pRestricted, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{
|
||||
Name: "restricted",
|
||||
ProviderId: "openai_api",
|
||||
UpstreamUrl: vllm.URL,
|
||||
ApiKey: &staticKey,
|
||||
Enabled: ptr(true),
|
||||
Models: models(modelSelected, modelOther),
|
||||
BootstrapCluster: ptr(harness.AgentNetworkCluster),
|
||||
Name: "restricted",
|
||||
ProviderId: "openai_api",
|
||||
UpstreamUrl: vllm.URL,
|
||||
ApiKey: &staticKey,
|
||||
Enabled: ptr(true),
|
||||
Models: models(modelSelected, modelOther),
|
||||
})
|
||||
require.NoError(t, err, "create restricted provider")
|
||||
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), pRestricted.Id) })
|
||||
|
||||
@@ -115,7 +115,7 @@ func TestGuardrailPerGroupAllowlist_AllProviders(t *testing.T) {
|
||||
staticKey := "static-e2e-token"
|
||||
enabled := true
|
||||
|
||||
for i, c := range cases {
|
||||
for _, c := range cases {
|
||||
req := api.AgentNetworkProviderRequest{
|
||||
Name: "e2e-pergroup-" + c.name,
|
||||
ProviderId: c.catalogID,
|
||||
@@ -124,9 +124,6 @@ func TestGuardrailPerGroupAllowlist_AllProviders(t *testing.T) {
|
||||
Enabled: ptr(true),
|
||||
Models: c.models,
|
||||
}
|
||||
if i == 0 {
|
||||
req.BootstrapCluster = ptr(harness.AgentNetworkCluster)
|
||||
}
|
||||
prov, perr := srv.CreateProvider(ctx, req)
|
||||
require.NoError(t, perr, "create provider %s", c.name)
|
||||
c.providerID = prov.Id
|
||||
@@ -283,13 +280,12 @@ func TestGuardrailMultiGroupUser(t *testing.T) {
|
||||
|
||||
// P1 — union scenario: two restricting policies, one per group.
|
||||
p1, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{
|
||||
Name: "e2e-mg-union",
|
||||
ProviderId: "openai_api",
|
||||
UpstreamUrl: vllm.URL,
|
||||
ApiKey: &staticKey,
|
||||
Enabled: ptr(true),
|
||||
Models: priced(unionA, unionB, unionC),
|
||||
BootstrapCluster: ptr(harness.AgentNetworkCluster),
|
||||
Name: "e2e-mg-union",
|
||||
ProviderId: "openai_api",
|
||||
UpstreamUrl: vllm.URL,
|
||||
ApiKey: &staticKey,
|
||||
Enabled: ptr(true),
|
||||
Models: priced(unionA, unionB, unionC),
|
||||
})
|
||||
require.NoError(t, err, "create union provider")
|
||||
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), p1.Id) })
|
||||
|
||||
@@ -115,14 +115,11 @@ func TestModelAllowlistEnforced(t *testing.T) {
|
||||
})
|
||||
require.NoError(t, err, "mint setup key")
|
||||
|
||||
// Providers with their configured (allowed) models; the first bootstraps the cluster.
|
||||
// Providers with their configured (allowed) models
|
||||
ids := make([]string, 0, len(providers))
|
||||
allowed := make([]string, 0, len(providers))
|
||||
for i, pc := range providers {
|
||||
for _, pc := range providers {
|
||||
req := providerRequest(pc)
|
||||
if i == 0 {
|
||||
req.BootstrapCluster = ptr(harness.AgentNetworkCluster)
|
||||
}
|
||||
prov, perr := srv.CreateProvider(ctx, req)
|
||||
require.NoError(t, perr, "create provider %s", pc.name)
|
||||
id := prov.Id
|
||||
|
||||
@@ -14,6 +14,7 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/netbirdio/netbird/e2e/harness"
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
)
|
||||
|
||||
// srv is the shared combined server for the package, ready (PAT-authenticated)
|
||||
@@ -42,5 +43,14 @@ func run(m *testing.M) int {
|
||||
return 1
|
||||
}
|
||||
|
||||
// Bootstrap the account's agent-network endpoint once for the package:
|
||||
// providers no longer have settings side effects, and every data-plane
|
||||
// test expects the shared account pinned to the combined proxy cluster.
|
||||
cluster := harness.AgentNetworkCluster
|
||||
if _, err := srv.CreateSettings(ctx, api.AgentNetworkSettingsCreateRequest{ProxyAddress: &cluster}); err != nil {
|
||||
fmt.Fprintf(os.Stderr, "e2e: bootstrap agent-network endpoint: %v\n", err)
|
||||
return 1
|
||||
}
|
||||
|
||||
return m.Run()
|
||||
}
|
||||
|
||||
@@ -21,11 +21,10 @@ func ptr[T any](v T) *T { return &v }
|
||||
func newProvider(t *testing.T, ctx context.Context, name string) api.AgentNetworkProvider {
|
||||
t.Helper()
|
||||
prov, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{
|
||||
Name: name,
|
||||
ProviderId: "openai_api",
|
||||
UpstreamUrl: "https://api.openai.com",
|
||||
ApiKey: ptr("sk-dummy-e2e-key"),
|
||||
BootstrapCluster: ptr("eu.proxy.netbird.test"),
|
||||
Name: name,
|
||||
ProviderId: "openai_api",
|
||||
UpstreamUrl: "https://api.openai.com",
|
||||
ApiKey: ptr("sk-dummy-e2e-key"),
|
||||
})
|
||||
require.NoError(t, err, "create provider %q", name)
|
||||
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), prov.Id) })
|
||||
@@ -57,17 +56,11 @@ func TestProviderLifecycle(t *testing.T) {
|
||||
}}
|
||||
}
|
||||
|
||||
for i, pc := range cases {
|
||||
i, pc := i, pc
|
||||
for _, pc := range cases {
|
||||
pc := pc
|
||||
t.Run(pc.name, func(t *testing.T) {
|
||||
req := providerRequest(pc)
|
||||
req.Name = "lc-" + pc.name
|
||||
// Bootstrap the cluster on the first create in case the matrix has
|
||||
// not run (e.g. no provider keys → settings not yet bootstrapped).
|
||||
if i == 0 {
|
||||
req.BootstrapCluster = ptr(harness.AgentNetworkCluster)
|
||||
}
|
||||
|
||||
prov, err := srv.CreateProvider(ctx, req)
|
||||
require.NoError(t, err, "create %s provider", pc.name)
|
||||
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), prov.Id) })
|
||||
@@ -137,45 +130,65 @@ func TestProviderValidation(t *testing.T) {
|
||||
requireClientError(t, err)
|
||||
}
|
||||
|
||||
// TestSettingsRoundTrip flips the collection toggles and confirms cluster /
|
||||
// subdomain stay immutable, then restores the original state.
|
||||
// TestSettingsRoundTrip flips the collection toggles and confirms the
|
||||
// endpoint and proxy address stay immutable, then restores the original
|
||||
// state. A second bootstrap attempt must be rejected as a conflict.
|
||||
func TestSettingsRoundTrip(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
// Settings are bootstrapped on first provider create.
|
||||
newProvider(t, ctx, "Settings Bootstrap")
|
||||
|
||||
// The package's TestMain bootstrapped the shared account's endpoint.
|
||||
before, err := srv.GetSettings(ctx)
|
||||
require.NoError(t, err, "get settings")
|
||||
require.NotEmpty(t, before.Cluster, "settings must carry an assigned cluster")
|
||||
require.NotEmpty(t, before.Endpoint, "settings must carry the bootstrapped endpoint")
|
||||
require.NotEmpty(t, before.ProxyAddress, "settings must carry the bootstrapped proxy address")
|
||||
|
||||
require.NotNil(t, before.AccessLogRetentionDays, "bootstrapped settings must carry a retention")
|
||||
beforeRetention := *before.AccessLogRetentionDays
|
||||
|
||||
flipped, err := srv.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{
|
||||
Endpoint: before.Endpoint,
|
||||
ProxyAddress: before.ProxyAddress,
|
||||
EnableLogCollection: !before.EnableLogCollection,
|
||||
EnablePromptCollection: !before.EnablePromptCollection,
|
||||
RedactPii: !before.RedactPii,
|
||||
AccessLogRetentionDays: beforeRetention,
|
||||
})
|
||||
require.NoError(t, err, "update settings")
|
||||
assert.Equal(t, !before.EnableLogCollection, flipped.EnableLogCollection, "log collection toggle must flip")
|
||||
assert.Equal(t, !before.EnablePromptCollection, flipped.EnablePromptCollection, "prompt collection toggle must flip")
|
||||
assert.Equal(t, before.Cluster, flipped.Cluster, "cluster must be immutable across updates")
|
||||
assert.Equal(t, before.Subdomain, flipped.Subdomain, "subdomain must be immutable across updates")
|
||||
require.NotNil(t, flipped.AccessLogRetentionDays)
|
||||
assert.Equal(t, beforeRetention, *flipped.AccessLogRetentionDays,
|
||||
"retention sent unchanged must round-trip, not reset to the zero value")
|
||||
assert.Equal(t, before.Endpoint, flipped.Endpoint, "endpoint must be immutable across updates")
|
||||
assert.Equal(t, before.ProxyAddress, flipped.ProxyAddress, "proxy address must be immutable across updates")
|
||||
|
||||
// A cluster different from the pinned one must be rejected; echoing the
|
||||
// pinned one back is valid.
|
||||
// The account is already bootstrapped: a second bootstrap is a conflict,
|
||||
// whatever shape it asks for.
|
||||
_, err = srv.CreateSettings(ctx, api.AgentNetworkSettingsCreateRequest{
|
||||
Endpoint: ptr("attacker.cluster.invalid"),
|
||||
})
|
||||
requireClientError(t, err)
|
||||
|
||||
// The identity fields ride along on the PUT as a required echo: a request
|
||||
// carrying a different endpoint is rejected without applying anything.
|
||||
_, err = srv.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{
|
||||
Cluster: ptr("attacker.cluster.invalid"),
|
||||
Endpoint: "other.cluster.invalid",
|
||||
ProxyAddress: before.ProxyAddress,
|
||||
EnableLogCollection: before.EnableLogCollection,
|
||||
EnablePromptCollection: before.EnablePromptCollection,
|
||||
RedactPii: before.RedactPii,
|
||||
AccessLogRetentionDays: beforeRetention,
|
||||
})
|
||||
requireClientError(t, err)
|
||||
|
||||
// Restore the original toggles.
|
||||
_, err = srv.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{
|
||||
Cluster: ptr(before.Cluster),
|
||||
Endpoint: before.Endpoint,
|
||||
ProxyAddress: before.ProxyAddress,
|
||||
EnableLogCollection: before.EnableLogCollection,
|
||||
EnablePromptCollection: before.EnablePromptCollection,
|
||||
RedactPii: before.RedactPii,
|
||||
AccessLogRetentionDays: beforeRetention,
|
||||
})
|
||||
require.NoError(t, err, "restore settings")
|
||||
}
|
||||
|
||||
@@ -4,6 +4,7 @@ package agentnetwork
|
||||
|
||||
import (
|
||||
"context"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
@@ -14,7 +15,8 @@ import (
|
||||
)
|
||||
|
||||
// harnessStartFresh boots a dedicated combined server with its own fresh
|
||||
// account and registers its teardown on t.
|
||||
// account and registers its teardown on t. Unlike the shared srv, the fresh
|
||||
// account has NOT had its agent-network endpoint bootstrapped.
|
||||
func harnessStartFresh(ctx context.Context, t *testing.T) (*harness.Combined, error) {
|
||||
t.Helper()
|
||||
fresh, err := harness.StartCombined(ctx)
|
||||
@@ -28,16 +30,16 @@ func harnessStartFresh(ctx context.Context, t *testing.T) (*harness.Combined, er
|
||||
return fresh, nil
|
||||
}
|
||||
|
||||
// TestSettingsBootstrapViaPut covers the settings-first bootstrap path on an
|
||||
// TestSettingsBootstrapViaPost covers the explicit bootstrap contract on an
|
||||
// account that has never been bootstrapped: the GET reads as the defaults
|
||||
// with an empty cluster/subdomain/endpoint, a PUT without a cluster has
|
||||
// nothing to pin and fails, and a PUT carrying a cluster creates the row and
|
||||
// pins it immutably. The shared srv cannot provide that starting state (any
|
||||
// provider-creating test bootstraps it, and test order is deliberately not
|
||||
// relied on), so this boots a dedicated combined server — the image is
|
||||
// already built and cached by TestMain's StartCombined, so the extra cost is
|
||||
// one container start.
|
||||
func TestSettingsBootstrapViaPut(t *testing.T) {
|
||||
// with an empty endpoint/proxy_address, a PUT has no row to update and fails,
|
||||
// and a POST creates the row and assigns the immutable endpoint — labeled
|
||||
// beneath a proxy address here, with the toggle overrides from the same
|
||||
// request applied. The shared srv cannot provide that starting state
|
||||
// (TestMain bootstraps it), so this boots a dedicated combined server — the
|
||||
// image is already built and cached by TestMain's StartCombined, so the extra
|
||||
// cost is one container start.
|
||||
func TestSettingsBootstrapViaPost(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
fresh, err := harnessStartFresh(ctx, t)
|
||||
@@ -47,32 +49,35 @@ func TestSettingsBootstrapViaPut(t *testing.T) {
|
||||
// as an error and not as a null body.
|
||||
before, err := fresh.GetSettings(ctx)
|
||||
require.NoError(t, err, "get settings on a fresh account must succeed")
|
||||
assert.Empty(t, before.Cluster, "cluster must be empty before bootstrap")
|
||||
assert.Empty(t, before.Subdomain, "subdomain must be empty before bootstrap")
|
||||
assert.Empty(t, before.Endpoint, "endpoint must be empty before bootstrap, not a bare dot")
|
||||
assert.Empty(t, before.Endpoint, "endpoint must be empty before bootstrap")
|
||||
assert.Empty(t, before.ProxyAddress, "proxy address must be empty before bootstrap")
|
||||
assert.False(t, before.Dedicated, "an unbootstrapped account has no serving shape")
|
||||
assert.True(t, before.EnableLogCollection, "defaults must show log collection on, matching bootstrap")
|
||||
assert.False(t, before.EnablePromptCollection, "defaults must show prompt collection off")
|
||||
|
||||
// A PUT without a cluster has nothing to pin the account to.
|
||||
// A PUT has no row to update yet — bootstrap is the explicit POST.
|
||||
_, err = fresh.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{
|
||||
EnableLogCollection: true,
|
||||
EnableLogCollection: true,
|
||||
AccessLogRetentionDays: 30,
|
||||
})
|
||||
requireClientError(t, err)
|
||||
|
||||
// A PUT carrying a cluster bootstraps the account and applies the
|
||||
// mutable fields from the same request. Every toggle is set away from
|
||||
// its bootstrap default so each assertion can actually fail.
|
||||
// A POST with a proxy address bootstraps a labeled endpoint and applies
|
||||
// the toggles from the same request. Every toggle is set away from its
|
||||
// bootstrap default so each assertion can actually fail.
|
||||
const cluster = "e2e.bootstrap.netbird.selfhosted"
|
||||
bootstrapped, err := fresh.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{
|
||||
Cluster: ptr(cluster),
|
||||
EnableLogCollection: false,
|
||||
EnablePromptCollection: true,
|
||||
RedactPii: true,
|
||||
bootstrapped, err := fresh.CreateSettings(ctx, api.AgentNetworkSettingsCreateRequest{
|
||||
ProxyAddress: ptr(cluster),
|
||||
EnableLogCollection: ptr(false),
|
||||
EnablePromptCollection: ptr(true),
|
||||
RedactPii: ptr(true),
|
||||
})
|
||||
require.NoError(t, err, "bootstrap settings via PUT must succeed")
|
||||
assert.Equal(t, cluster, bootstrapped.Cluster, "cluster must be pinned from the request")
|
||||
require.NotEmpty(t, bootstrapped.Subdomain, "subdomain must be assigned at bootstrap")
|
||||
assert.Equal(t, bootstrapped.Subdomain+"."+cluster, bootstrapped.Endpoint, "endpoint must combine subdomain and cluster")
|
||||
require.NoError(t, err, "bootstrap settings via POST must succeed")
|
||||
assert.Equal(t, cluster, bootstrapped.ProxyAddress, "proxy address must be pinned from the request")
|
||||
require.NotEmpty(t, bootstrapped.Endpoint, "endpoint must be assigned at bootstrap")
|
||||
assert.True(t, strings.HasSuffix(bootstrapped.Endpoint, "."+cluster),
|
||||
"labeled endpoint must hang one label beneath the proxy address: %s", bootstrapped.Endpoint)
|
||||
assert.False(t, bootstrapped.Dedicated, "a labeled pin is not dedicated")
|
||||
assert.False(t, bootstrapped.EnableLogCollection, "log collection from the bootstrap request must override the default")
|
||||
assert.True(t, bootstrapped.EnablePromptCollection, "prompt collection from the bootstrap request must apply")
|
||||
assert.True(t, bootstrapped.RedactPii, "redact toggle from the bootstrap request must apply")
|
||||
@@ -85,30 +90,90 @@ func TestSettingsBootstrapViaPut(t *testing.T) {
|
||||
assert.Equal(t, bootstrapped.EnablePromptCollection, after.EnablePromptCollection, "prompt collection must persist")
|
||||
assert.Equal(t, bootstrapped.RedactPii, after.RedactPii, "redact toggle must persist")
|
||||
|
||||
// Once bootstrapped, later updates may omit the cluster entirely.
|
||||
// Once bootstrapped, PUT updates the toggles. The identity fields ride
|
||||
// along as a required echo of the assigned values; a matching echo is
|
||||
// accepted and never written.
|
||||
persisted, err := fresh.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{
|
||||
Endpoint: bootstrapped.Endpoint,
|
||||
ProxyAddress: bootstrapped.ProxyAddress,
|
||||
EnableLogCollection: true,
|
||||
EnablePromptCollection: false,
|
||||
RedactPii: true,
|
||||
AccessLogRetentionDays: 21,
|
||||
})
|
||||
require.NoError(t, err, "post-bootstrap update without cluster must succeed")
|
||||
assert.Equal(t, cluster, persisted.Cluster, "omitted cluster must keep the pinned value")
|
||||
require.NoError(t, err, "post-bootstrap update must succeed")
|
||||
require.NotNil(t, persisted.AccessLogRetentionDays)
|
||||
assert.Equal(t, 21, *persisted.AccessLogRetentionDays, "retention from the update must apply")
|
||||
assert.Equal(t, bootstrapped.Endpoint, persisted.Endpoint, "endpoint must survive updates untouched")
|
||||
assert.Equal(t, cluster, persisted.ProxyAddress, "proxy address must survive updates untouched")
|
||||
assert.True(t, persisted.EnableLogCollection, "post-bootstrap toggle must apply")
|
||||
assert.False(t, persisted.EnablePromptCollection, "post-bootstrap toggle must apply")
|
||||
|
||||
// The cluster is immutable: a different value is rejected rather than
|
||||
// silently ignored, and the rejected update must not disturb anything.
|
||||
// The endpoint is immutable: a PUT carrying a different endpoint is
|
||||
// rejected, and a second bootstrap is rejected as a conflict. Neither
|
||||
// rejected write may disturb anything.
|
||||
_, err = fresh.UpdateSettings(ctx, api.AgentNetworkSettingsRequest{
|
||||
Cluster: ptr("other.cluster.invalid"),
|
||||
EnableLogCollection: false,
|
||||
Endpoint: "other.cluster.invalid",
|
||||
ProxyAddress: persisted.ProxyAddress,
|
||||
EnableLogCollection: persisted.EnableLogCollection,
|
||||
EnablePromptCollection: persisted.EnablePromptCollection,
|
||||
RedactPii: persisted.RedactPii,
|
||||
AccessLogRetentionDays: 21,
|
||||
})
|
||||
requireClientError(t, err)
|
||||
|
||||
_, err = fresh.CreateSettings(ctx, api.AgentNetworkSettingsCreateRequest{
|
||||
Endpoint: ptr("other.cluster.invalid"),
|
||||
})
|
||||
requireClientError(t, err)
|
||||
|
||||
final, err := fresh.GetSettings(ctx)
|
||||
require.NoError(t, err, "get settings after the rejected cluster change must succeed")
|
||||
assert.Equal(t, persisted.Cluster, final.Cluster, "rejected update must not change the cluster")
|
||||
assert.Equal(t, persisted.Endpoint, final.Endpoint, "rejected update must not change the endpoint")
|
||||
assert.Equal(t, persisted.EnableLogCollection, final.EnableLogCollection, "rejected update must not apply its toggles")
|
||||
assert.Equal(t, persisted.EnablePromptCollection, final.EnablePromptCollection, "rejected update must not apply its toggles")
|
||||
assert.Equal(t, persisted.RedactPii, final.RedactPii, "rejected update must not apply its toggles")
|
||||
require.NoError(t, err, "get settings after the rejected bootstrap must succeed")
|
||||
assert.Equal(t, persisted.Endpoint, final.Endpoint, "rejected bootstrap must not change the endpoint")
|
||||
assert.Equal(t, persisted.ProxyAddress, final.ProxyAddress, "rejected bootstrap must not change the proxy address")
|
||||
assert.Equal(t, persisted.EnableLogCollection, final.EnableLogCollection, "rejected bootstrap must not apply its toggles")
|
||||
assert.Equal(t, persisted.EnablePromptCollection, final.EnablePromptCollection, "rejected bootstrap must not apply its toggles")
|
||||
assert.Equal(t, persisted.RedactPii, final.RedactPii, "rejected bootstrap must not apply its toggles")
|
||||
}
|
||||
|
||||
// TestSettingsBootstrapSelfAddressed covers the dedicated shape end to end:
|
||||
// a POST carrying an endpoint claims the hostname verbatim, the proxy address
|
||||
// equals it, and the pin reads as dedicated — the address-first flow a
|
||||
// self-hosted operator uses before deploying the proxy that will declare it.
|
||||
// The tail covers the recovery path the guarded DELETE exists for: with no
|
||||
// providers and no proxy at the address, the claim can be released and a
|
||||
// fresh bootstrap succeeds — the fix for a typo'd immutable endpoint.
|
||||
func TestSettingsBootstrapSelfAddressed(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
fresh, err := harnessStartFresh(ctx, t)
|
||||
require.NoError(t, err, "start dedicated combined server")
|
||||
|
||||
created, err := fresh.CreateSettings(ctx, api.AgentNetworkSettingsCreateRequest{
|
||||
Endpoint: ptr("gw.e2e.netbird.selfhosted"),
|
||||
})
|
||||
require.NoError(t, err, "self-addressed bootstrap must succeed")
|
||||
assert.Equal(t, "gw.e2e.netbird.selfhosted", created.Endpoint, "endpoint must be claimed verbatim")
|
||||
assert.Equal(t, created.Endpoint, created.ProxyAddress, "self-addressed: proxy address is the endpoint")
|
||||
assert.True(t, created.Dedicated, "a self-addressed pin is dedicated")
|
||||
|
||||
// No providers exist and no proxy declares the address, so both delete
|
||||
// guards are clear: the delete releases the claim and the account reads
|
||||
// as unbootstrapped defaults again.
|
||||
require.NoError(t, fresh.DeleteSettings(ctx), "guarded delete with both guards clear must succeed")
|
||||
|
||||
after, err := fresh.GetSettings(ctx)
|
||||
require.NoError(t, err, "get settings after delete must succeed")
|
||||
assert.Empty(t, after.Endpoint, "a deleted account must read as unbootstrapped")
|
||||
|
||||
// A second delete has nothing to remove.
|
||||
requireClientError(t, fresh.DeleteSettings(ctx))
|
||||
|
||||
// Re-creating is a fresh bootstrap — the released hostname is free to be
|
||||
// claimed again, or a different one chosen.
|
||||
recreated, err := fresh.CreateSettings(ctx, api.AgentNetworkSettingsCreateRequest{
|
||||
Endpoint: ptr("gw2.e2e.netbird.selfhosted"),
|
||||
})
|
||||
require.NoError(t, err, "bootstrap after delete must succeed")
|
||||
assert.Equal(t, "gw2.e2e.netbird.selfhosted", recreated.Endpoint, "the fresh bootstrap claims the new hostname")
|
||||
}
|
||||
|
||||
@@ -66,9 +66,7 @@ func TestProviderSkipTLSVerification(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// First create bootstraps the account cluster.
|
||||
insecureReq := newReq("skip-tls", insecureModel, true)
|
||||
insecureReq.BootstrapCluster = ptr(harness.AgentNetworkCluster)
|
||||
insecureProv, err := srv.CreateProvider(ctx, insecureReq)
|
||||
require.NoError(t, err, "create skip-tls provider")
|
||||
t.Cleanup(func() { _ = srv.DeleteProvider(context.Background(), insecureProv.Id) })
|
||||
|
||||
@@ -57,12 +57,11 @@ func TestVLLMProvider(t *testing.T) {
|
||||
// is enumerated so the router dispatches this model string to this provider.
|
||||
dummyKey := "sk-vllm-e2e"
|
||||
prov, err := srv.CreateProvider(ctx, api.AgentNetworkProviderRequest{
|
||||
Name: "vllm",
|
||||
ProviderId: "vllm",
|
||||
UpstreamUrl: vllm.URL,
|
||||
ApiKey: &dummyKey,
|
||||
Enabled: ptr(true),
|
||||
BootstrapCluster: ptr(harness.AgentNetworkCluster),
|
||||
Name: "vllm",
|
||||
ProviderId: "vllm",
|
||||
UpstreamUrl: vllm.URL,
|
||||
ApiKey: &dummyKey,
|
||||
Enabled: ptr(true),
|
||||
Models: &[]api.AgentNetworkProviderModel{
|
||||
{Id: harness.VLLMModel, InputPer1k: 0.001, OutputPer1k: 0.002},
|
||||
},
|
||||
|
||||
@@ -20,5 +20,9 @@ ENV NETBIRD_BIN="/usr/local/bin/netbird" \
|
||||
NB_ENABLE_CAPTURE="false" \
|
||||
NB_ENTRYPOINT_SERVICE_TIMEOUT="30"
|
||||
ENTRYPOINT [ "/usr/local/bin/netbird-entrypoint.sh" ]
|
||||
COPY client/netbird-entrypoint.sh /usr/local/bin/netbird-entrypoint.sh
|
||||
# --chmod because the build context is not always a git checkout. A suite in
|
||||
# another module builds from this module's extracted copy in the module cache,
|
||||
# where every file is 0444 — the cache drops the executable bit git records — and
|
||||
# a bare COPY then produces an entrypoint the runtime cannot exec.
|
||||
COPY --chmod=0755 client/netbird-entrypoint.sh /usr/local/bin/netbird-entrypoint.sh
|
||||
COPY --from=builder /out/netbird /usr/local/bin/netbird
|
||||
|
||||
@@ -126,17 +126,33 @@ func (c *Combined) DeleteGuardrail(ctx context.Context, id string) error {
|
||||
return anDelete(ctx, c, "/api/agent-network/guardrails/"+id)
|
||||
}
|
||||
|
||||
// GetSettings returns the account's agent-network settings row. It exists only
|
||||
// after the first provider create bootstraps it.
|
||||
// CreateSettings bootstraps the account's agent-network settings row,
|
||||
// assigning the immutable endpoint. Exactly one of req.ProxyAddress (labeled
|
||||
// endpoint beneath that cluster) and req.Endpoint (self-addressed dedicated
|
||||
// endpoint) must be set; a second bootstrap returns a conflict.
|
||||
func (c *Combined) CreateSettings(ctx context.Context, req api.AgentNetworkSettingsCreateRequest) (api.AgentNetworkSettings, error) {
|
||||
return anRequest[api.AgentNetworkSettings](ctx, c, http.MethodPost, "/api/agent-network/settings", req)
|
||||
}
|
||||
|
||||
// GetSettings returns the account's agent-network settings row. Before the
|
||||
// CreateSettings bootstrap it reads as the defaults with an empty endpoint.
|
||||
func (c *Combined) GetSettings(ctx context.Context) (api.AgentNetworkSettings, error) {
|
||||
return anRequest[api.AgentNetworkSettings](ctx, c, http.MethodGet, "/api/agent-network/settings", nil)
|
||||
}
|
||||
|
||||
// UpdateSettings applies the mutable collection toggles.
|
||||
// UpdateSettings applies the mutable collection toggles. The request must
|
||||
// echo the assigned endpoint and proxy address unchanged — the server rejects
|
||||
// a PUT that tries to change them.
|
||||
func (c *Combined) UpdateSettings(ctx context.Context, req api.AgentNetworkSettingsRequest) (api.AgentNetworkSettings, error) {
|
||||
return anRequest[api.AgentNetworkSettings](ctx, c, http.MethodPut, "/api/agent-network/settings", req)
|
||||
}
|
||||
|
||||
// DeleteSettings removes the account's settings row, releasing the endpoint.
|
||||
// Refused while providers exist or a proxy is actively serving the endpoint.
|
||||
func (c *Combined) DeleteSettings(ctx context.Context) error {
|
||||
return anDelete(ctx, c, "/api/agent-network/settings")
|
||||
}
|
||||
|
||||
// ListConsumption returns the account's consumption rows (possibly empty).
|
||||
func (c *Combined) ListConsumption(ctx context.Context) ([]api.AgentNetworkConsumption, error) {
|
||||
return anRequest[[]api.AgentNetworkConsumption](ctx, c, http.MethodGet, "/api/agent-network/consumption", nil)
|
||||
|
||||
@@ -32,12 +32,36 @@ type Client struct {
|
||||
container testcontainers.Container
|
||||
}
|
||||
|
||||
// clientOptions is what the ClientOption values assemble.
|
||||
type clientOptions struct {
|
||||
name string
|
||||
}
|
||||
|
||||
// ClientOption adjusts how StartClient runs the agent.
|
||||
type ClientOption func(*clientOptions)
|
||||
|
||||
// WithClientName names the agent, which sets both its network alias and its
|
||||
// container hostname. The hostname matters beyond addressing: the agent reports
|
||||
// it to management at registration, so it is the name the peer appears under in
|
||||
// the API.
|
||||
//
|
||||
// Required to run more than one agent against the same server — the default name
|
||||
// is shared, and two containers cannot hold the same alias on one network.
|
||||
func WithClientName(name string) ClientOption {
|
||||
return func(o *clientOptions) { o.name = name }
|
||||
}
|
||||
|
||||
// StartClient builds the client image and runs it on the combined server's
|
||||
// network, joining via the given setup key. The image entrypoint brings the
|
||||
// daemon up automatically; callers wait for connectivity with WaitConnected /
|
||||
// WaitProxyPeer.
|
||||
func StartClient(ctx context.Context, c *Combined, setupKey string) (*Client, error) {
|
||||
root, err := repoRoot()
|
||||
func StartClient(ctx context.Context, c *Combined, setupKey string, opts ...ClientOption) (*Client, error) {
|
||||
o := clientOptions{name: clientAlias}
|
||||
for _, opt := range opts {
|
||||
opt(&o)
|
||||
}
|
||||
|
||||
root, err := repoRoot(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -47,9 +71,13 @@ func StartClient(ctx context.Context, c *Combined, setupKey string) (*Client, er
|
||||
}
|
||||
|
||||
req := testcontainers.ContainerRequest{
|
||||
Image: clientImage,
|
||||
Image: clientImage,
|
||||
// The agent reports the container's hostname to management, so this is
|
||||
// the name the peer is addressable by in the API as well as on the
|
||||
// network. The entrypoint takes no hostname flag of its own.
|
||||
Hostname: o.name,
|
||||
Networks: []string{c.network.Name},
|
||||
NetworkAliases: map[string][]string{c.network.Name: {clientAlias}},
|
||||
NetworkAliases: map[string][]string{c.network.Name: {o.name}},
|
||||
Env: map[string]string{
|
||||
"NB_MANAGEMENT_URL": combinedExposedURL,
|
||||
"NB_SETUP_KEY": setupKey,
|
||||
|
||||
@@ -61,11 +61,68 @@ type Combined struct {
|
||||
workDir string
|
||||
}
|
||||
|
||||
// combinedOptions is what the CombinedOption values assemble.
|
||||
type combinedOptions struct {
|
||||
geolocation bool
|
||||
env map[string]string
|
||||
}
|
||||
|
||||
// CombinedOption adjusts how StartCombined boots the server. The defaults suit a
|
||||
// suite that only drives the API; the options exist for the ones that need more
|
||||
// of the product than that.
|
||||
type CombinedOption func(*combinedOptions)
|
||||
|
||||
// WithGeolocation leaves the GeoLite database download enabled. It is off by
|
||||
// default because the download adds startup latency that most suites get nothing
|
||||
// for. A suite asserting on location-based posture checks needs it: management
|
||||
// evaluates those rules against the database, and without it the rule fails
|
||||
// instead of passing without having been checked.
|
||||
func WithGeolocation() CombinedOption {
|
||||
return func(o *combinedOptions) { o.geolocation = true }
|
||||
}
|
||||
|
||||
// WithServerEnv adds environment variables to the combined container, overriding
|
||||
// the defaults on a key collision. For settings this harness does not model
|
||||
// directly, so a suite needing one does not have to fork the harness to get it.
|
||||
func WithServerEnv(env map[string]string) CombinedOption {
|
||||
return func(o *combinedOptions) {
|
||||
if o.env == nil {
|
||||
o.env = map[string]string{}
|
||||
}
|
||||
for k, v := range env {
|
||||
o.env[k] = v
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// combinedEnv is the combined container's environment: setup-PAT enabled so the
|
||||
// caller can mint an admin token through /api/setup, geolocation off unless the
|
||||
// suite asked for it, and whatever the suite added on top.
|
||||
func combinedEnv(o combinedOptions) map[string]string {
|
||||
env := map[string]string{
|
||||
"NB_SETUP_PAT_ENABLED": "true",
|
||||
}
|
||||
if !o.geolocation {
|
||||
// Skip the GeoLite DB download — it blocks startup and agent-network
|
||||
// ingest doesn't use geolocation.
|
||||
env["NB_DISABLE_GEOLOCATION"] = "true"
|
||||
}
|
||||
for k, v := range o.env {
|
||||
env[k] = v
|
||||
}
|
||||
return env
|
||||
}
|
||||
|
||||
// StartCombined builds the combined server from its multistage Dockerfile and
|
||||
// boots it with setup-PAT enabled on a fresh shared network, returning once the
|
||||
// API is serving. The caller still owns minting the admin PAT via Bootstrap.
|
||||
func StartCombined(ctx context.Context) (*Combined, error) {
|
||||
root, err := repoRoot()
|
||||
func StartCombined(ctx context.Context, opts ...CombinedOption) (*Combined, error) {
|
||||
var o combinedOptions
|
||||
for _, opt := range opts {
|
||||
opt(&o)
|
||||
}
|
||||
|
||||
root, err := repoRoot(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -88,7 +145,7 @@ func StartCombined(ctx context.Context) (*Combined, error) {
|
||||
return nil, fmt.Errorf("create work dir: %w", err)
|
||||
}
|
||||
|
||||
cfg := fmt.Sprintf(combinedConfigYAML, combinedExposedURL, containerIssuer)
|
||||
cfg := fmt.Sprintf(combinedConfigYAML, combinedExposedURL, !o.geolocation, containerIssuer)
|
||||
if err := os.WriteFile(filepath.Join(workDir, "config.yaml"), []byte(cfg), 0o644); err != nil { //nolint:gosec // non-secret config, bind-mounted and read by the container
|
||||
_ = net.Remove(ctx)
|
||||
return nil, fmt.Errorf("write combined config: %w", err)
|
||||
@@ -112,13 +169,8 @@ func StartCombined(ctx context.Context) (*Combined, error) {
|
||||
ExposedPorts: []string{combinedHTTPPort},
|
||||
Networks: []string{net.Name},
|
||||
NetworkAliases: map[string][]string{net.Name: {combinedAlias}},
|
||||
Env: map[string]string{
|
||||
"NB_SETUP_PAT_ENABLED": "true",
|
||||
// Skip the GeoLite DB download — it blocks startup and agent-network
|
||||
// ingest doesn't use geolocation.
|
||||
"NB_DISABLE_GEOLOCATION": "true",
|
||||
},
|
||||
Cmd: []string{"--config", "/nb/config.yaml"},
|
||||
Env: combinedEnv(o),
|
||||
Cmd: []string{"--config", "/nb/config.yaml"},
|
||||
HostConfigModifier: func(hc *container.HostConfig) {
|
||||
hc.Binds = append(hc.Binds, workDir+":/nb")
|
||||
},
|
||||
|
||||
@@ -15,6 +15,11 @@ package harness
|
||||
// server is required to load it — a broken path or malformed file fails startup
|
||||
// rather than silently falling back to the compiled-in rates, and TestMain then
|
||||
// fails with the container logs.
|
||||
//
|
||||
// disableGeoliteUpdate is a parameter rather than a fixed true because a suite
|
||||
// that exercises geolocation needs the database: management can only evaluate a
|
||||
// location rule with GeoLite loaded, and a rule it cannot evaluate fails rather
|
||||
// than passing vacuously. See WithGeolocation.
|
||||
const combinedConfigYAML = `server:
|
||||
listenAddress: ":8080"
|
||||
exposedAddress: "%s"
|
||||
@@ -25,7 +30,7 @@ const combinedConfigYAML = `server:
|
||||
authSecret: "e2e-relay-secret"
|
||||
dataDir: "/nb/data"
|
||||
disableAnonymousMetrics: true
|
||||
disableGeoliteUpdate: true
|
||||
disableGeoliteUpdate: %t
|
||||
auth:
|
||||
issuer: "%s"
|
||||
store:
|
||||
|
||||
161
e2e/harness/options_test.go
Normal file
161
e2e/harness/options_test.go
Normal file
@@ -0,0 +1,161 @@
|
||||
//go:build e2e
|
||||
|
||||
package harness
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
)
|
||||
|
||||
// The options exist so a suite can ask for a deployment this harness would not
|
||||
// otherwise give it. What they configure is a container environment and a config
|
||||
// file, both assembled before anything is started, so they are checkable without
|
||||
// Docker — which is the point: a wiring mistake here would otherwise only show up
|
||||
// as a puzzling failure minutes into a container run.
|
||||
|
||||
func TestCombinedEnvGeolocation(t *testing.T) {
|
||||
var off combinedOptions
|
||||
assert.Equal(t, "true", combinedEnv(off)["NB_DISABLE_GEOLOCATION"],
|
||||
"geolocation should be off by default")
|
||||
|
||||
var on combinedOptions
|
||||
WithGeolocation()(&on)
|
||||
assert.NotContains(t, combinedEnv(on), "NB_DISABLE_GEOLOCATION",
|
||||
"WithGeolocation must leave NB_DISABLE_GEOLOCATION unset, so the server downloads the database")
|
||||
assert.Equal(t, "true", combinedEnv(on)["NB_SETUP_PAT_ENABLED"],
|
||||
"the setup PAT must stay enabled whatever else is configured; Bootstrap depends on it")
|
||||
}
|
||||
|
||||
// The config file carries the same decision as the environment variable, and the
|
||||
// server needs both to agree: disableGeoliteUpdate suppresses the download even
|
||||
// when geolocation itself is enabled.
|
||||
func TestCombinedConfigGeolocation(t *testing.T) {
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
opts []CombinedOption
|
||||
want string
|
||||
}{
|
||||
{name: "default", want: "disableGeoliteUpdate: true"},
|
||||
{name: "with geolocation", opts: []CombinedOption{WithGeolocation()}, want: "disableGeoliteUpdate: false"},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
var o combinedOptions
|
||||
for _, opt := range tc.opts {
|
||||
opt(&o)
|
||||
}
|
||||
cfg := fmt.Sprintf(combinedConfigYAML, combinedExposedURL, !o.geolocation, containerIssuer)
|
||||
assert.Contains(t, cfg, tc.want, "geolocation not rendered as expected")
|
||||
// The issuer is the last verb; a mis-ordered argument list would put
|
||||
// the boolean here instead and the server would fail to start.
|
||||
assert.Contains(t, cfg, `issuer: "`+containerIssuer+`"`, "issuer not rendered")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestWithServerEnvOverrides(t *testing.T) {
|
||||
var o combinedOptions
|
||||
WithServerEnv(map[string]string{"NB_LOG_LEVEL": "debug"})(&o)
|
||||
WithServerEnv(map[string]string{"NB_SETUP_PAT_ENABLED": "false"})(&o)
|
||||
|
||||
env := combinedEnv(o)
|
||||
assert.Equal(t, "debug", env["NB_LOG_LEVEL"], "added variable missing")
|
||||
assert.Equal(t, "false", env["NB_SETUP_PAT_ENABLED"], "a suite must be able to override a default")
|
||||
}
|
||||
|
||||
// Two agents on one network cannot share an alias, so the name has to reach both
|
||||
// the alias and the hostname. The hostname is the one management records, so it is
|
||||
// also what the peer is addressable by through the API.
|
||||
func TestWithClientName(t *testing.T) {
|
||||
o := clientOptions{name: clientAlias}
|
||||
require.Equal(t, "client", o.name, "unexpected default client name")
|
||||
|
||||
WithClientName("peer2")(&o)
|
||||
assert.Equal(t, "peer2", o.name, "WithClientName did not take")
|
||||
}
|
||||
|
||||
// repoRoot has to recognise this module rather than merely finding a go.mod, or a
|
||||
// suite in another module gets its own root and a build context without the
|
||||
// component Dockerfiles in it.
|
||||
func TestIsModule(t *testing.T) {
|
||||
dir := t.TempDir()
|
||||
|
||||
other := filepath.Join(dir, "go.mod")
|
||||
require.NoError(t, os.WriteFile(other, []byte("module example.com/other\n\ngo 1.25\n"), 0o600))
|
||||
assert.False(t, isModule(other, modulePath), "another module's go.mod must not be taken for this repo")
|
||||
|
||||
ours := filepath.Join(dir, "ours.mod")
|
||||
require.NoError(t, os.WriteFile(ours, []byte("// a comment\n\nmodule "+modulePath+"\n\ngo 1.25\n"), 0o600))
|
||||
assert.True(t, isModule(ours, modulePath), "this repo's go.mod was not recognised")
|
||||
|
||||
assert.False(t, isModule(filepath.Join(dir, "absent.mod"), modulePath),
|
||||
"a missing go.mod must not report a match")
|
||||
}
|
||||
|
||||
// Running from inside the repo, repoRoot finds it by walking up — the module
|
||||
// lookup is only the fallback, and this asserts the walk still wins so an in-repo
|
||||
// run never depends on the module cache.
|
||||
func TestRepoRootFindsThisRepo(t *testing.T) {
|
||||
root, err := repoRoot(context.Background())
|
||||
require.NoError(t, err)
|
||||
assert.True(t, isModule(filepath.Join(root, "go.mod"), modulePath),
|
||||
"repoRoot returned %s, which is not this module", root)
|
||||
|
||||
for _, f := range []string{combinedDockerfile, clientDockerfile} {
|
||||
_, err := os.Stat(filepath.Join(root, f))
|
||||
assert.NoError(t, err, "%s is not present under the reported root %s", f, root)
|
||||
}
|
||||
}
|
||||
|
||||
// A caller that vendors its dependencies puts the go command in automatic vendor
|
||||
// mode, where `go list -m -f {{.Dir}}` succeeds and reports an EMPTY directory:
|
||||
// vendor/ holds packages, not module source. Without -mod=readonly the lookup
|
||||
// would come back empty and the harness would report a missing module for a
|
||||
// dependency that is present.
|
||||
func TestModuleDirResolvesUnderVendorMode(t *testing.T) {
|
||||
if _, err := exec.LookPath("go"); err != nil {
|
||||
t.Skip("no go tool on PATH")
|
||||
}
|
||||
ctx := context.Background()
|
||||
|
||||
base := t.TempDir()
|
||||
dep := filepath.Join(base, "dep")
|
||||
main := filepath.Join(base, "main")
|
||||
require.NoError(t, os.MkdirAll(dep, 0o750))
|
||||
require.NoError(t, os.MkdirAll(main, 0o750))
|
||||
|
||||
// A local replacement rather than a real dependency, so this needs no network.
|
||||
require.NoError(t, os.WriteFile(filepath.Join(dep, "go.mod"),
|
||||
[]byte("module example.com/dep\n\ngo 1.25\n"), 0o600))
|
||||
require.NoError(t, os.WriteFile(filepath.Join(dep, "dep.go"),
|
||||
[]byte("package dep\n"), 0o600))
|
||||
require.NoError(t, os.WriteFile(filepath.Join(main, "go.mod"),
|
||||
[]byte("module example.com/main\n\ngo 1.25\n\nrequire example.com/dep v0.0.0\n\nreplace example.com/dep v0.0.0 => ../dep\n"), 0o600))
|
||||
require.NoError(t, os.WriteFile(filepath.Join(main, "main.go"),
|
||||
[]byte("package main\n\nimport _ \"example.com/dep\"\n\nfunc main() {}\n"), 0o600))
|
||||
|
||||
t.Chdir(main)
|
||||
vendor := exec.CommandContext(ctx, "go", "mod", "vendor")
|
||||
out, err := vendor.CombinedOutput()
|
||||
require.NoError(t, err, "go mod vendor: %s", out)
|
||||
|
||||
dir, err := moduleDir(ctx, "example.com/dep")
|
||||
require.NoError(t, err, "the module must still resolve with a vendor directory present")
|
||||
assert.Equal(t, dep, dir, "resolved the wrong directory")
|
||||
}
|
||||
|
||||
// A cancelled context has to stop the lookup rather than leaving the caller
|
||||
// waiting on a subprocess it has already given up on.
|
||||
func TestModuleDirHonoursContext(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
_, err := moduleDir(ctx, modulePath)
|
||||
assert.ErrorIs(t, err, context.Canceled, "a cancelled context must stop the lookup")
|
||||
}
|
||||
@@ -3,27 +3,82 @@
|
||||
package harness
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"os"
|
||||
"os/exec"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// repoRoot walks up from the working directory to the module root (the
|
||||
// directory holding go.mod), so the Docker build context is correct no matter
|
||||
// which package the test runs from.
|
||||
func repoRoot() (string, error) {
|
||||
// modulePath is this module, used both to recognise the repo when walking up
|
||||
// from the working directory and to locate it when the suite lives elsewhere.
|
||||
const modulePath = "github.com/netbirdio/netbird"
|
||||
|
||||
// repoRoot returns the directory the component Dockerfiles are built from.
|
||||
//
|
||||
// Walking up from the working directory finds it for any test inside this repo,
|
||||
// no matter which package it runs from. A suite in another module gets a
|
||||
// different answer that way — its own module root, where combined/Dockerfile
|
||||
// does not exist — so the ancestor has to be this module and not merely some
|
||||
// module. When it is not, the build context is the extracted module directory of
|
||||
// whichever version that suite depends on, which is the right one: the server it
|
||||
// tests against is then built from the same revision as the client library it
|
||||
// was compiled with.
|
||||
func repoRoot(ctx context.Context) (string, error) {
|
||||
dir, err := os.Getwd()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
for {
|
||||
if _, statErr := os.Stat(filepath.Join(dir, "go.mod")); statErr == nil {
|
||||
if isModule(filepath.Join(dir, "go.mod"), modulePath) {
|
||||
return dir, nil
|
||||
}
|
||||
parent := filepath.Dir(dir)
|
||||
if parent == dir {
|
||||
return "", fmt.Errorf("go.mod not found above %s", dir)
|
||||
break
|
||||
}
|
||||
dir = parent
|
||||
}
|
||||
return moduleDir(ctx, modulePath)
|
||||
}
|
||||
|
||||
// isModule reports whether the go.mod at path declares the given module.
|
||||
func isModule(path, want string) bool {
|
||||
b, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return false
|
||||
}
|
||||
for _, line := range strings.Split(string(b), "\n") {
|
||||
if rest, ok := strings.CutPrefix(strings.TrimSpace(line), "module "); ok {
|
||||
return strings.TrimSpace(rest) == want
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// moduleDir asks the go tool where a module's source is, which for a dependent
|
||||
// module is its extracted copy in the module cache. The cache is read-only, and
|
||||
// a Docker build context is only ever read.
|
||||
//
|
||||
// -mod=readonly is required rather than cosmetic. A caller that vendors its
|
||||
// dependencies puts the go command in automatic vendor mode, where this lookup
|
||||
// succeeds with an EMPTY directory — vendor/ holds packages, not module source,
|
||||
// so there is nothing to report. Asking in readonly mode resolves against the
|
||||
// module graph instead, which answers for both a cached module and a local
|
||||
// replacement, and neither writes to go.mod.
|
||||
func moduleDir(ctx context.Context, module string) (string, error) {
|
||||
cmd := exec.CommandContext(ctx, "go", "list", "-mod=readonly", "-m", "-f", "{{.Dir}}", module)
|
||||
out, err := cmd.Output()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("locate %s: %w", module, err)
|
||||
}
|
||||
dir := strings.TrimSpace(string(out))
|
||||
if dir == "" {
|
||||
return "", fmt.Errorf("locate %s: the go tool reported no directory; run `go mod download %s`", module, module)
|
||||
}
|
||||
if _, err := os.Stat(dir); err != nil {
|
||||
return "", fmt.Errorf("locate %s: %w", module, err)
|
||||
}
|
||||
return dir, nil
|
||||
}
|
||||
|
||||
@@ -43,7 +43,7 @@ type Proxy struct {
|
||||
// or override any NB_PROXY_* var (e.g. NB_PROXY_TUNNEL_CACHE_TTL for tests that
|
||||
// need a short authorization-cache window).
|
||||
func StartProxy(ctx context.Context, c *Combined, proxyToken string, envOverrides ...map[string]string) (*Proxy, error) {
|
||||
root, err := repoRoot()
|
||||
root, err := repoRoot(ctx)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
5
go.mod
5
go.mod
@@ -99,7 +99,7 @@ require (
|
||||
github.com/pires/go-proxyproto v0.11.0
|
||||
github.com/pkg/sftp v1.13.9
|
||||
github.com/prometheus/client_golang v1.23.2
|
||||
github.com/quic-go/quic-go v0.55.0
|
||||
github.com/quic-go/quic-go v0.59.1
|
||||
github.com/redis/go-redis/v9 v9.7.3
|
||||
github.com/rs/xid v1.3.0
|
||||
github.com/shirou/gopsutil/v4 v4.25.8
|
||||
@@ -239,7 +239,6 @@ require (
|
||||
github.com/jackc/pgservicefile v0.0.0-20221227161230-091c0ba34f0a // indirect
|
||||
github.com/jackc/puddle/v2 v2.2.1 // indirect
|
||||
github.com/jackpal/go-nat-pmp v1.0.2 // indirect
|
||||
github.com/jchv/go-winloader v0.0.0-20250406163304-c1995be93bd1 // indirect
|
||||
github.com/jinzhu/inflection v1.0.0 // indirect
|
||||
github.com/jinzhu/now v1.1.5 // indirect
|
||||
github.com/jmespath/go-jmespath v0.4.0 // indirect
|
||||
@@ -340,4 +339,4 @@ 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.20260807055527-fc03f984d701
|
||||
replace github.com/wailsapp/wails/v3 => github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260810103952-24e716aea4db
|
||||
|
||||
11
go.sum
11
go.sum
@@ -349,8 +349,6 @@ github.com/jackc/puddle/v2 v2.2.1 h1:RhxXJtFG022u4ibrCSMSiu5aOq1i77R3OHKNJj77OAk
|
||||
github.com/jackc/puddle/v2 v2.2.1/go.mod h1:vriiEXHvEE654aYKXXjOvZM39qJ0q+azkZFrfEOc3H4=
|
||||
github.com/jackpal/go-nat-pmp v1.0.2 h1:KzKSgb7qkJvOUTqYl9/Hg/me3pWgBmERKrTGD7BdWus=
|
||||
github.com/jackpal/go-nat-pmp v1.0.2/go.mod h1:QPH045xvCAeXUZOxsnwmrtiCoxIr9eob+4orBN1SBKc=
|
||||
github.com/jchv/go-winloader v0.0.0-20250406163304-c1995be93bd1 h1:njuLRcjAuMKr7kI3D85AXWkw6/+v9PwtV6M6o11sWHQ=
|
||||
github.com/jchv/go-winloader v0.0.0-20250406163304-c1995be93bd1/go.mod h1:alcuEEnZsY1WQsagKhZDsoPCRoOijYqhZvPwLG0kzVs=
|
||||
github.com/jcmturner/aescts/v2 v2.0.0 h1:9YKLH6ey7H4eDBXW8khjYslgyqG2xZikXP0EQFKrle8=
|
||||
github.com/jcmturner/aescts/v2 v2.0.0/go.mod h1:AiaICIRyfYg35RUkr8yESTqvSy7csK90qZ5xfvvsoNs=
|
||||
github.com/jcmturner/dnsutils/v2 v2.0.0 h1:lltnkeZGL0wILNvrNiVCR6Ro5PGU/SeBvVO/8c/iPbo=
|
||||
@@ -490,8 +488,8 @@ github.com/netbirdio/service v0.0.0-20240911161631-f62744f42502 h1:3tHlFmhTdX9ax
|
||||
github.com/netbirdio/service v0.0.0-20240911161631-f62744f42502/go.mod h1:CIMRFEJVL+0DS1a3Nx06NaMn4Dz63Ng6O7dl0qH0zVM=
|
||||
github.com/netbirdio/signal-dispatcher/dispatcher v0.0.0-20250805121659-6b4ac470ca45 h1:ujgviVYmx243Ksy7NdSwrdGPSRNE3pb8kEDSpH0QuAQ=
|
||||
github.com/netbirdio/signal-dispatcher/dispatcher v0.0.0-20250805121659-6b4ac470ca45/go.mod h1:5/sjFmLb8O96B5737VCqhHyGRzNFIaN/Bu7ZodXc3qQ=
|
||||
github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260807055527-fc03f984d701 h1:QL9nupfRom0L9jcY7N9l/Bc6QK2PtC6pHzC+ftpTqpw=
|
||||
github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260807055527-fc03f984d701/go.mod h1:BzATbK71VFikMMMCo434wAi0QcaI03P+xeaWgDvQvjw=
|
||||
github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260810103952-24e716aea4db h1:gBOE2r4AW1soSmpYJC5/n9/1L8UQ8+HLjed8CY/TzZY=
|
||||
github.com/netbirdio/wails/v3 v3.0.0-beta.3.0.20260810103952-24e716aea4db/go.mod h1:bsdahLwBQxXjlmdPPeQyrTcDJfcqAr/ymFj0RXhwtWI=
|
||||
github.com/netbirdio/wireguard-go v0.0.0-20260628102922-2834bebf6c1a h1:3CWK+yTvRKOcC0Q8VCTGy4l60TEb27CQVS7LkMxwjmw=
|
||||
github.com/netbirdio/wireguard-go v0.0.0-20260628102922-2834bebf6c1a/go.mod h1:rpwXGsirqLqN2L0JDJQlwOboGHmptD5ZD6T2VmcqhTw=
|
||||
github.com/nxadm/tail v1.4.4/go.mod h1:kenIhsEOeOJmVchQTgglprH7qJGnHDVpk1VPCcaMI8A=
|
||||
@@ -582,8 +580,8 @@ github.com/prometheus/otlptranslator v1.0.0 h1:s0LJW/iN9dkIH+EnhiD3BlkkP5QVIUVEo
|
||||
github.com/prometheus/otlptranslator v1.0.0/go.mod h1:vRYWnXvI6aWGpsdY/mOT/cbeVRBlPWtBNDb7kGR3uKM=
|
||||
github.com/prometheus/procfs v0.19.2 h1:zUMhqEW66Ex7OXIiDkll3tl9a1ZdilUOd/F6ZXw4Vws=
|
||||
github.com/prometheus/procfs v0.19.2/go.mod h1:M0aotyiemPhBCM0z5w87kL22CxfcH05ZpYlu+b4J7mw=
|
||||
github.com/quic-go/quic-go v0.55.0 h1:zccPQIqYCXDt5NmcEabyYvOnomjs8Tlwl7tISjJh9Mk=
|
||||
github.com/quic-go/quic-go v0.55.0/go.mod h1:DR51ilwU1uE164KuWXhinFcKWGlEjzys2l8zUl5Ss1U=
|
||||
github.com/quic-go/quic-go v0.59.1 h1:0Gmua0HW1Tv7ANR7hUYwRyD0MG5OJfgvYSZasGZzBic=
|
||||
github.com/quic-go/quic-go v0.59.1/go.mod h1:upnsH4Ju1YkqpLXC305eW3yDZ4NfnNbmQRCMWS58IKU=
|
||||
github.com/redis/go-redis/v9 v9.7.3 h1:YpPyAayJV+XErNsatSElgRZZVCwXX9QzkKYNvO7x0wM=
|
||||
github.com/redis/go-redis/v9 v9.7.3/go.mod h1:bGUrSggJ9X9GUmZpZNEOQKaANxSGgOEBRltRTZHSvrA=
|
||||
github.com/rogpeppe/go-internal v1.14.1 h1:UQB4HGPB6osV0SQTLymcB4TgvyWu6ZyliaW0tI/otEQ=
|
||||
@@ -793,7 +791,6 @@ golang.org/x/sys v0.0.0-20190916202348-b4ddaad3f8a3/go.mod h1:h1NjWce9XRLGQEsW7w
|
||||
golang.org/x/sys v0.0.0-20191005200804-aed5e4c7ecf9/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
golang.org/x/sys v0.0.0-20191120155948-bd437916bb0e/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
golang.org/x/sys v0.0.0-20200323222414-85ca7c5b95cd/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
golang.org/x/sys v0.0.0-20200810151505-1b9f1253b3ed/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
golang.org/x/sys v0.0.0-20200930185726-fdedc70b468f/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
golang.org/x/sys v0.0.0-20201015000850-e3ed0017c211/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
golang.org/x/sys v0.0.0-20201119102817-f84b799fce68/go.mod h1:h1NjWce9XRLGQEsW7wpKNCjG9DtNlClVuFLEZdDNbEs=
|
||||
|
||||
@@ -102,8 +102,8 @@ func TestSettingsHandler_GetExposesCollectionToggles(t *testing.T) {
|
||||
|
||||
require.NoError(t, f.store.SaveAgentNetworkSettings(context.Background(), &agentNetworkTypes.Settings{
|
||||
AccountID: testAccountID,
|
||||
Cluster: "eu.proxy.netbird.io",
|
||||
Subdomain: "violet",
|
||||
Domain: "violet.eu.proxy.netbird.io",
|
||||
ProxyAddress: "eu.proxy.netbird.io",
|
||||
EnableLogCollection: true,
|
||||
EnablePromptCollection: true,
|
||||
RedactPii: false,
|
||||
|
||||
@@ -155,12 +155,7 @@ func (h *handler) createProvider(w http.ResponseWriter, r *http.Request) {
|
||||
provider := types.NewProvider(userAuth.AccountId)
|
||||
provider.FromAPIRequest(&req)
|
||||
|
||||
bootstrapCluster := ""
|
||||
if req.BootstrapCluster != nil {
|
||||
bootstrapCluster = *req.BootstrapCluster
|
||||
}
|
||||
|
||||
created, err := h.manager.CreateProvider(r.Context(), userAuth.UserId, provider, bootstrapCluster)
|
||||
created, err := h.manager.CreateProvider(r.Context(), userAuth.UserId, provider)
|
||||
if err != nil {
|
||||
util.WriteError(r.Context(), err, w)
|
||||
return
|
||||
|
||||
@@ -12,13 +12,55 @@ import (
|
||||
"github.com/netbirdio/netbird/shared/management/http/util"
|
||||
)
|
||||
|
||||
// addSettingsEndpoints registers the Agent Network settings routes. The
|
||||
// settings row is bootstrapped server-side on first provider create or on the
|
||||
// first PUT carrying a cluster; GET reads it and PUT applies a partial update
|
||||
// of the mutable collection toggles (cluster/subdomain stay immutable).
|
||||
// addSettingsEndpoints registers the Agent Network settings routes. POST
|
||||
// bootstraps the settings row, assigning the account's immutable endpoint;
|
||||
// GET reads it (defaults with an empty endpoint before bootstrap); PUT
|
||||
// carries every field, replacing the mutable collection toggles and rejecting
|
||||
// any change to the identity fields; DELETE removes the row — guarded so it
|
||||
// stays a bootstrap-repair operation — releasing the endpoint for a fresh
|
||||
// bootstrap.
|
||||
func (h *handler) addSettingsEndpoints(router *mux.Router) {
|
||||
router.HandleFunc("/agent-network/settings", h.getSettings).Methods("GET", "OPTIONS")
|
||||
router.HandleFunc("/agent-network/settings", h.createSettings).Methods("POST", "OPTIONS")
|
||||
router.HandleFunc("/agent-network/settings", h.updateSettings).Methods("PUT", "OPTIONS")
|
||||
router.HandleFunc("/agent-network/settings", h.deleteSettings).Methods("DELETE", "OPTIONS")
|
||||
}
|
||||
|
||||
// createSettings bootstraps the account's settings row. Exactly one of
|
||||
// proxy_address (labeled endpoint; the server allocates the label) and
|
||||
// endpoint (self-addressed, claimed verbatim) must be provided; optional
|
||||
// collection toggles ride along with defaults for omitted fields.
|
||||
func (h *handler) createSettings(w http.ResponseWriter, r *http.Request) {
|
||||
userAuth, err := nbcontext.GetUserAuthFromContext(r.Context())
|
||||
if err != nil {
|
||||
util.WriteError(r.Context(), err, w)
|
||||
return
|
||||
}
|
||||
|
||||
var req api.AgentNetworkSettingsCreateRequest
|
||||
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||
util.WriteErrorResponse("couldn't parse JSON request", http.StatusBadRequest, w)
|
||||
return
|
||||
}
|
||||
|
||||
settings := types.DefaultSettings(userAuth.AccountId)
|
||||
settings.FromAPICreateRequest(&req)
|
||||
|
||||
proxyAddress := ""
|
||||
if req.ProxyAddress != nil {
|
||||
proxyAddress = *req.ProxyAddress
|
||||
}
|
||||
endpoint := ""
|
||||
if req.Endpoint != nil {
|
||||
endpoint = *req.Endpoint
|
||||
}
|
||||
|
||||
created, err := h.manager.CreateSettings(r.Context(), userAuth.UserId, settings, proxyAddress, endpoint)
|
||||
if err != nil {
|
||||
util.WriteError(r.Context(), err, w)
|
||||
return
|
||||
}
|
||||
util.WriteJSONObject(r.Context(), w, created.ToAPIResponse())
|
||||
}
|
||||
|
||||
// updateSettings replaces the mutable settings fields on the account's row.
|
||||
@@ -48,6 +90,24 @@ func (h *handler) updateSettings(w http.ResponseWriter, r *http.Request) {
|
||||
util.WriteJSONObject(r.Context(), w, updated.ToAPIResponse())
|
||||
}
|
||||
|
||||
// deleteSettings removes the account's settings row, releasing the endpoint.
|
||||
// The manager refuses (412) while providers exist or a proxy is actively
|
||||
// serving the endpoint; a later POST bootstraps fresh, allocating a new
|
||||
// endpoint.
|
||||
func (h *handler) deleteSettings(w http.ResponseWriter, r *http.Request) {
|
||||
userAuth, err := nbcontext.GetUserAuthFromContext(r.Context())
|
||||
if err != nil {
|
||||
util.WriteError(r.Context(), err, w)
|
||||
return
|
||||
}
|
||||
|
||||
if err := h.manager.DeleteSettings(r.Context(), userAuth.AccountId, userAuth.UserId); err != nil {
|
||||
util.WriteError(r.Context(), err, w)
|
||||
return
|
||||
}
|
||||
util.WriteJSONObject(r.Context(), w, util.EmptyObject{})
|
||||
}
|
||||
|
||||
// getSettings returns the account's agent-network settings. Accounts that
|
||||
// haven't been bootstrapped yet read as the defaults with an empty cluster,
|
||||
// subdomain and endpoint; the manager synthesises that view.
|
||||
|
||||
@@ -1,20 +1,25 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
rpproxy "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
)
|
||||
|
||||
// TestSettingsHandler_GetUnbootstrappedReturnsDefaults pins the settings-read
|
||||
// convention shared with the account and DNS settings endpoints: settings
|
||||
// always read as a JSON object. Before bootstrap that object carries the
|
||||
// defaults with an empty cluster/subdomain/endpoint (the "not bootstrapped"
|
||||
// defaults with an empty endpoint/proxy_address (the "not bootstrapped"
|
||||
// signal) and no timestamps — never a 404 and never the legacy null body.
|
||||
func TestSettingsHandler_GetUnbootstrappedReturnsDefaults(t *testing.T) {
|
||||
f := newAgentNetworkHandlerFixture(t)
|
||||
@@ -27,9 +32,9 @@ func TestSettingsHandler_GetUnbootstrappedReturnsDefaults(t *testing.T) {
|
||||
|
||||
var got api.AgentNetworkSettings
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
|
||||
assert.Empty(t, got.Cluster, "cluster must be empty until bootstrapped")
|
||||
assert.Empty(t, got.Subdomain, "subdomain must be empty until bootstrapped")
|
||||
assert.Empty(t, got.Endpoint, "endpoint must be empty until bootstrapped, not a bare dot")
|
||||
assert.Empty(t, got.Endpoint, "endpoint must be empty until bootstrapped")
|
||||
assert.Empty(t, got.ProxyAddress, "proxy address must be empty until bootstrapped")
|
||||
assert.False(t, got.Dedicated, "an unbootstrapped account has no serving shape")
|
||||
assert.True(t, got.EnableLogCollection, "defaults must show log collection on, matching bootstrap")
|
||||
assert.False(t, got.EnablePromptCollection, "defaults must show prompt collection off")
|
||||
assert.False(t, got.RedactPii, "defaults must show redaction off")
|
||||
@@ -39,62 +44,149 @@ func TestSettingsHandler_GetUnbootstrappedReturnsDefaults(t *testing.T) {
|
||||
assert.Nil(t, got.UpdatedAt, "no timestamps before a row exists")
|
||||
}
|
||||
|
||||
// TestSettingsHandler_PutBootstrapsWithCluster covers the settings-first
|
||||
// bootstrap path: a PUT carrying a cluster on an unbootstrapped account
|
||||
// creates the row (cluster pinned, subdomain assigned) and applies the
|
||||
// mutable fields from the same request.
|
||||
func TestSettingsHandler_PutBootstrapsWithCluster(t *testing.T) {
|
||||
// TestSettingsHandler_PostBootstrapsLabeled covers the labeled bootstrap
|
||||
// shape: a POST carrying a proxy_address allocates a label beneath it, so the
|
||||
// endpoint hangs one label under the shared cluster's address and the pin is
|
||||
// not dedicated. Toggles riding along apply; omitted ones keep defaults.
|
||||
func TestSettingsHandler_PostBootstrapsLabeled(t *testing.T) {
|
||||
f := newAgentNetworkHandlerFixture(t)
|
||||
|
||||
rec := f.do(t, http.MethodPut, "/agent-network/settings",
|
||||
`{"cluster": "eu.proxy.netbird.io", "enable_log_collection": true, "enable_prompt_collection": true, "redact_pii": false, "access_log_retention_days": 30}`)
|
||||
require.Equal(t, http.StatusOK, rec.Code, "bootstrap PUT must succeed: %s", rec.Body.String())
|
||||
rec := f.do(t, http.MethodPost, "/agent-network/settings",
|
||||
`{"proxy_address": "eu.proxy.netbird.io", "enable_prompt_collection": true, "access_log_retention_days": 14}`)
|
||||
require.Equal(t, http.StatusOK, rec.Code, "bootstrap POST must succeed: %s", rec.Body.String())
|
||||
|
||||
var got api.AgentNetworkSettings
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
|
||||
assert.Equal(t, "eu.proxy.netbird.io", got.Cluster, "cluster must be pinned from the request")
|
||||
assert.NotEmpty(t, got.Subdomain, "subdomain must be assigned at bootstrap")
|
||||
assert.Equal(t, got.Subdomain+".eu.proxy.netbird.io", got.Endpoint, "endpoint must combine subdomain and cluster")
|
||||
assert.True(t, got.EnableLogCollection, "toggle from the bootstrap request must apply")
|
||||
assert.Equal(t, "eu.proxy.netbird.io", got.ProxyAddress, "proxy address must be pinned from the request")
|
||||
require.NotEmpty(t, got.Endpoint, "endpoint must be allocated at bootstrap")
|
||||
assert.True(t, strings.HasSuffix(got.Endpoint, ".eu.proxy.netbird.io"),
|
||||
"labeled endpoint must hang off the proxy address: %s", got.Endpoint)
|
||||
label := strings.TrimSuffix(got.Endpoint, ".eu.proxy.netbird.io")
|
||||
assert.NotContains(t, label, ".", "the allocated label must be a single DNS label: %s", label)
|
||||
assert.False(t, got.Dedicated, "a labeled pin is not dedicated")
|
||||
assert.True(t, got.EnableLogCollection, "omitted toggle must keep its default")
|
||||
assert.True(t, got.EnablePromptCollection, "toggle from the bootstrap request must apply")
|
||||
require.NotNil(t, got.AccessLogRetentionDays)
|
||||
assert.Equal(t, 30, *got.AccessLogRetentionDays, "retention from the bootstrap request must apply")
|
||||
assert.Equal(t, 14, *got.AccessLogRetentionDays, "retention from the bootstrap request must apply")
|
||||
assert.NotNil(t, got.CreatedAt, "a persisted row carries timestamps")
|
||||
|
||||
// The row is now readable via GET.
|
||||
rec = f.do(t, http.MethodGet, "/agent-network/settings", "")
|
||||
require.Equal(t, http.StatusOK, rec.Code, "GET after bootstrap must succeed")
|
||||
var read api.AgentNetworkSettings
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &read))
|
||||
assert.Equal(t, got.Endpoint, read.Endpoint, "GET must return the bootstrapped endpoint")
|
||||
}
|
||||
|
||||
// TestSettingsHandler_PutWithoutClusterOnUnbootstrapped pins that a PUT
|
||||
// without a cluster cannot conjure a settings row out of nothing — there is
|
||||
// no cluster to pin — and surfaces as 404 like the GET.
|
||||
func TestSettingsHandler_PutWithoutClusterOnUnbootstrapped(t *testing.T) {
|
||||
// TestSettingsHandler_PostBootstrapsSelfAddressed covers the dedicated shape:
|
||||
// a POST carrying an endpoint claims the hostname verbatim, the proxy address
|
||||
// equals it, and the pin reads as dedicated. The claim is legitimate before
|
||||
// any proxy declares the address (address-first).
|
||||
func TestSettingsHandler_PostBootstrapsSelfAddressed(t *testing.T) {
|
||||
f := newAgentNetworkHandlerFixture(t)
|
||||
|
||||
rec := f.do(t, http.MethodPost, "/agent-network/settings",
|
||||
`{"endpoint": "Brave-Otter.Gateway.Example.com"}`)
|
||||
require.Equal(t, http.StatusOK, rec.Code, "bootstrap POST must succeed: %s", rec.Body.String())
|
||||
|
||||
var got api.AgentNetworkSettings
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
|
||||
assert.Equal(t, "brave-otter.gateway.example.com", got.Endpoint,
|
||||
"endpoint must be claimed verbatim, lowercased")
|
||||
assert.Equal(t, got.Endpoint, got.ProxyAddress, "self-addressed: the proxy address is the endpoint")
|
||||
assert.True(t, got.Dedicated, "a self-addressed pin is dedicated")
|
||||
assert.True(t, got.EnableLogCollection, "omitted toggles must keep their defaults")
|
||||
}
|
||||
|
||||
// TestSettingsHandler_PostRequiresExactlyOneIdentityField pins the request
|
||||
// contract: proxy_address and endpoint are mutually exclusive and one is
|
||||
// required — both or neither is a validation error, not a guess.
|
||||
func TestSettingsHandler_PostRequiresExactlyOneIdentityField(t *testing.T) {
|
||||
f := newAgentNetworkHandlerFixture(t)
|
||||
|
||||
rec := f.do(t, http.MethodPost, "/agent-network/settings", `{}`)
|
||||
assert.Equal(t, http.StatusUnprocessableEntity, rec.Code,
|
||||
"empty POST must be rejected: got %d body=%s", rec.Code, rec.Body.String())
|
||||
|
||||
rec = f.do(t, http.MethodPost, "/agent-network/settings",
|
||||
`{"proxy_address": "eu.proxy.netbird.io", "endpoint": "brave-otter.gateway.example.com"}`)
|
||||
assert.Equal(t, http.StatusUnprocessableEntity, rec.Code,
|
||||
"POST with both identity fields must be rejected: got %d body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
|
||||
// TestSettingsHandler_PostRejectsMalformedHostnames pins per-write input
|
||||
// validation: shapes canonicalization cannot repair — trailing dots, embedded
|
||||
// whitespace, empty labels — are rejected with a validation error instead of
|
||||
// landing in an immutable column.
|
||||
func TestSettingsHandler_PostRejectsMalformedHostnames(t *testing.T) {
|
||||
f := newAgentNetworkHandlerFixture(t)
|
||||
|
||||
for name, body := range map[string]string{
|
||||
"trailing dot": `{"endpoint": "gateway.example.com."}`,
|
||||
"leading dot": `{"endpoint": ".gateway.example.com"}`,
|
||||
"inner whitespace": `{"endpoint": "gate way.example.com"}`,
|
||||
"empty label": `{"proxy_address": "eu..proxy.netbird.io"}`,
|
||||
} {
|
||||
rec := f.do(t, http.MethodPost, "/agent-network/settings", body)
|
||||
assert.Equal(t, http.StatusUnprocessableEntity, rec.Code,
|
||||
"%s must be rejected: got %d body=%s", name, rec.Code, rec.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
// TestSettingsHandler_PostConflictsOnSecondBootstrap pins that bootstrap is a
|
||||
// one-time create: a second POST returns 409 and leaves the row untouched.
|
||||
func TestSettingsHandler_PostConflictsOnSecondBootstrap(t *testing.T) {
|
||||
f := newAgentNetworkHandlerFixture(t)
|
||||
|
||||
rec := f.do(t, http.MethodPost, "/agent-network/settings", `{"proxy_address": "eu.proxy.netbird.io"}`)
|
||||
require.Equal(t, http.StatusOK, rec.Code, "first bootstrap must succeed: %s", rec.Body.String())
|
||||
var first api.AgentNetworkSettings
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &first))
|
||||
|
||||
rec = f.do(t, http.MethodPost, "/agent-network/settings", `{"proxy_address": "us.proxy.netbird.io"}`)
|
||||
assert.Equal(t, http.StatusConflict, rec.Code,
|
||||
"second bootstrap must 409: got %d body=%s", rec.Code, rec.Body.String())
|
||||
|
||||
rec = f.do(t, http.MethodGet, "/agent-network/settings", "")
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
var got api.AgentNetworkSettings
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
|
||||
assert.Equal(t, first.Endpoint, got.Endpoint, "the original endpoint must survive the rejected bootstrap")
|
||||
assert.Equal(t, first.ProxyAddress, got.ProxyAddress, "the original proxy address must survive")
|
||||
}
|
||||
|
||||
// TestSettingsHandler_PutBeforeBootstrapIs404 pins that a PUT cannot conjure a
|
||||
// settings row out of nothing — bootstrap is the explicit POST — and the
|
||||
// error points the caller there.
|
||||
func TestSettingsHandler_PutBeforeBootstrapIs404(t *testing.T) {
|
||||
f := newAgentNetworkHandlerFixture(t)
|
||||
|
||||
rec := f.do(t, http.MethodPut, "/agent-network/settings",
|
||||
`{"enable_log_collection": false, "enable_prompt_collection": false, "redact_pii": false}`)
|
||||
assert.Equal(t, http.StatusNotFound, rec.Code,
|
||||
"cluster-less PUT on an unbootstrapped account must 404: got %d body=%s", rec.Code, rec.Body.String())
|
||||
assert.Contains(t, rec.Body.String(), "cluster",
|
||||
"the error must point the caller at the bootstrap paths: %s", rec.Body.String())
|
||||
"PUT on an unbootstrapped account must 404: got %d body=%s", rec.Code, rec.Body.String())
|
||||
assert.Contains(t, rec.Body.String(), "/api/agent-network/settings",
|
||||
"the error must point the caller at the bootstrap POST: %s", rec.Body.String())
|
||||
}
|
||||
|
||||
// TestSettingsHandler_PutReplacesMutableFields pins the update contract shared
|
||||
// with the other PUT endpoints: the request replaces every mutable field, so a
|
||||
// toggle absent from the JSON lands as its zero value rather than being
|
||||
// preserved. Cluster and subdomain survive untouched.
|
||||
// with the other PUT endpoints: the request carries every field, replacing the
|
||||
// mutable ones. The identity fields ride along as a required echo of the
|
||||
// assigned values — compared, never written — so the endpoint and proxy
|
||||
// address survive every accepted update.
|
||||
func TestSettingsHandler_PutReplacesMutableFields(t *testing.T) {
|
||||
f := newAgentNetworkHandlerFixture(t)
|
||||
|
||||
rec := f.do(t, http.MethodPut, "/agent-network/settings",
|
||||
`{"cluster": "eu.proxy.netbird.io", "enable_log_collection": true, "enable_prompt_collection": true, "redact_pii": true, "access_log_retention_days": 14}`)
|
||||
require.Equal(t, http.StatusOK, rec.Code, "bootstrap PUT must succeed: %s", rec.Body.String())
|
||||
rec := f.do(t, http.MethodPost, "/agent-network/settings",
|
||||
`{"proxy_address": "eu.proxy.netbird.io", "enable_prompt_collection": true, "redact_pii": true, "access_log_retention_days": 14}`)
|
||||
require.Equal(t, http.StatusOK, rec.Code, "bootstrap POST must succeed: %s", rec.Body.String())
|
||||
|
||||
var before api.AgentNetworkSettings
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &before))
|
||||
|
||||
rec = f.do(t, http.MethodPut, "/agent-network/settings",
|
||||
`{"enable_log_collection": true, "enable_prompt_collection": false, "redact_pii": false}`)
|
||||
rec = f.do(t, http.MethodPut, "/agent-network/settings", fmt.Sprintf(
|
||||
`{"endpoint": %q, "proxy_address": %q, "enable_log_collection": true, "enable_prompt_collection": false, "redact_pii": false, "access_log_retention_days": 7}`,
|
||||
before.Endpoint, before.ProxyAddress))
|
||||
require.Equal(t, http.StatusOK, rec.Code, "update PUT must succeed: %s", rec.Body.String())
|
||||
|
||||
var got api.AgentNetworkSettings
|
||||
@@ -103,35 +195,201 @@ func TestSettingsHandler_PutReplacesMutableFields(t *testing.T) {
|
||||
assert.False(t, got.EnablePromptCollection, "sent toggle must apply")
|
||||
assert.False(t, got.RedactPii, "sent toggle must apply")
|
||||
require.NotNil(t, got.AccessLogRetentionDays)
|
||||
assert.Equal(t, 0, *got.AccessLogRetentionDays,
|
||||
"retention absent from the request must land as the zero value — PUT replaces all mutable fields")
|
||||
assert.Equal(t, before.Cluster, got.Cluster, "cluster must survive updates untouched")
|
||||
assert.Equal(t, before.Subdomain, got.Subdomain, "subdomain must survive updates untouched")
|
||||
assert.Equal(t, 7, *got.AccessLogRetentionDays, "sent retention must apply")
|
||||
assert.Equal(t, before.Endpoint, got.Endpoint, "endpoint must survive updates untouched")
|
||||
assert.Equal(t, before.ProxyAddress, got.ProxyAddress, "proxy address must survive updates untouched")
|
||||
}
|
||||
|
||||
// TestSettingsHandler_PutRejectsClusterChange pins cluster immutability: once
|
||||
// assigned, a differing cluster is rejected as a validation error instead of
|
||||
// being silently ignored, so callers never observe a value other than the one
|
||||
// they sent. Echoing the assigned cluster back stays valid, which lets
|
||||
// declarative clients send their full desired state idempotently.
|
||||
func TestSettingsHandler_PutRejectsClusterChange(t *testing.T) {
|
||||
// TestSettingsHandler_PutRejectsChangedIdentity pins the immutability contract:
|
||||
// the PUT carries the identity fields like every other field, but they are an
|
||||
// echo — a request carrying a different endpoint or proxy address is rejected
|
||||
// as a validation error and the row is left untouched. The comparison is
|
||||
// lenient about casing (the stored values are normalized lowercase), so a
|
||||
// client replaying a GET response with different casing is not rejected.
|
||||
func TestSettingsHandler_PutRejectsChangedIdentity(t *testing.T) {
|
||||
f := newAgentNetworkHandlerFixture(t)
|
||||
|
||||
rec := f.do(t, http.MethodPut, "/agent-network/settings",
|
||||
`{"cluster": "eu.proxy.netbird.io", "enable_log_collection": true, "enable_prompt_collection": false, "redact_pii": false}`)
|
||||
require.Equal(t, http.StatusOK, rec.Code, "bootstrap PUT must succeed: %s", rec.Body.String())
|
||||
rec := f.do(t, http.MethodPost, "/agent-network/settings",
|
||||
`{"proxy_address": "eu.proxy.netbird.io", "enable_prompt_collection": true}`)
|
||||
require.Equal(t, http.StatusOK, rec.Code, "bootstrap POST must succeed: %s", rec.Body.String())
|
||||
var before api.AgentNetworkSettings
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &before))
|
||||
|
||||
rec = f.do(t, http.MethodPut, "/agent-network/settings",
|
||||
`{"cluster": "us.proxy.netbird.io", "enable_log_collection": true, "enable_prompt_collection": false, "redact_pii": false}`)
|
||||
assert.Equal(t, http.StatusUnprocessableEntity, rec.Code,
|
||||
"cluster change must be rejected as a validation error: got %d body=%s", rec.Code, rec.Body.String())
|
||||
for name, body := range map[string]string{
|
||||
"changed endpoint": fmt.Sprintf(
|
||||
`{"endpoint": "other.gateway.example.com", "proxy_address": %q, "enable_log_collection": false, "enable_prompt_collection": false, "redact_pii": false, "access_log_retention_days": 30}`,
|
||||
before.ProxyAddress),
|
||||
"changed proxy_address": fmt.Sprintf(
|
||||
`{"endpoint": %q, "proxy_address": "us.proxy.netbird.io", "enable_log_collection": false, "enable_prompt_collection": false, "redact_pii": false, "access_log_retention_days": 30}`,
|
||||
before.Endpoint),
|
||||
"omitted identity": `{"enable_log_collection": false, "enable_prompt_collection": false, "redact_pii": false, "access_log_retention_days": 30}`,
|
||||
} {
|
||||
rec = f.do(t, http.MethodPut, "/agent-network/settings", body)
|
||||
assert.Equal(t, http.StatusUnprocessableEntity, rec.Code,
|
||||
"%s must be rejected: got %d body=%s", name, rec.Code, rec.Body.String())
|
||||
}
|
||||
|
||||
rec = f.do(t, http.MethodPut, "/agent-network/settings",
|
||||
`{"cluster": "eu.proxy.netbird.io", "enable_log_collection": true, "enable_prompt_collection": false, "redact_pii": true}`)
|
||||
require.Equal(t, http.StatusOK, rec.Code, "echoing the assigned cluster must stay valid: %s", rec.Body.String())
|
||||
// The rejected updates must not have applied anything — toggles included.
|
||||
rec = f.do(t, http.MethodGet, "/agent-network/settings", "")
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
var got api.AgentNetworkSettings
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
|
||||
assert.Equal(t, before.Endpoint, got.Endpoint, "rejected PUT must not change the endpoint")
|
||||
assert.Equal(t, before.ProxyAddress, got.ProxyAddress, "rejected PUT must not change the proxy address")
|
||||
assert.True(t, got.EnablePromptCollection, "rejected PUT must not apply its toggles")
|
||||
|
||||
// An uppercased echo of the assigned values still names the same host and
|
||||
// must be accepted.
|
||||
rec = f.do(t, http.MethodPut, "/agent-network/settings", fmt.Sprintf(
|
||||
`{"endpoint": %q, "proxy_address": %q, "enable_log_collection": true, "enable_prompt_collection": true, "redact_pii": false, "access_log_retention_days": 30}`,
|
||||
strings.ToUpper(before.Endpoint), strings.ToUpper(before.ProxyAddress)))
|
||||
assert.Equal(t, http.StatusOK, rec.Code,
|
||||
"an uppercased identity echo must be accepted: got %d body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
|
||||
// TestSettingsHandler_PutOmittedRetentionLandsAsZero documents a residual the
|
||||
// required-ness of access_log_retention_days does not remove. Marking the field
|
||||
// required changes the generated client type from *int to int, so a generated
|
||||
// client cannot omit it — but nothing validates OpenAPI required-ness at
|
||||
// runtime, so a hand-rolled body without the field still decodes as 0, which
|
||||
// the API documents as "keep indefinitely".
|
||||
//
|
||||
// That is the same latitude the three booleans already have, so it is left
|
||||
// consistent rather than special-cased. This test exists to make the gap
|
||||
// explicit: if request validation is ever added, this expectation is what
|
||||
// changes.
|
||||
func TestSettingsHandler_PutOmittedRetentionLandsAsZero(t *testing.T) {
|
||||
f := newAgentNetworkHandlerFixture(t)
|
||||
|
||||
rec := f.do(t, http.MethodPost, "/agent-network/settings",
|
||||
`{"proxy_address": "eu.proxy.netbird.io", "access_log_retention_days": 14}`)
|
||||
require.Equal(t, http.StatusOK, rec.Code, "bootstrap POST must succeed: %s", rec.Body.String())
|
||||
var before api.AgentNetworkSettings
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &before))
|
||||
|
||||
rec = f.do(t, http.MethodPut, "/agent-network/settings", fmt.Sprintf(
|
||||
`{"endpoint": %q, "proxy_address": %q, "enable_log_collection": true, "enable_prompt_collection": false, "redact_pii": false}`,
|
||||
before.Endpoint, before.ProxyAddress))
|
||||
require.Equal(t, http.StatusOK, rec.Code, "update PUT must succeed: %s", rec.Body.String())
|
||||
|
||||
var got api.AgentNetworkSettings
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
|
||||
assert.Equal(t, "eu.proxy.netbird.io", got.Cluster, "cluster must be unchanged")
|
||||
assert.True(t, got.RedactPii, "toggle sent alongside the echoed cluster must apply")
|
||||
require.NotNil(t, got.AccessLogRetentionDays)
|
||||
assert.Equal(t, 0, *got.AccessLogRetentionDays,
|
||||
"a non-conforming body that omits retention still replaces it with the zero value")
|
||||
}
|
||||
|
||||
// TestSettingsHandler_DeleteBeforeBootstrapIs404 pins that DELETE on an
|
||||
// account with no settings row is a 404, mirroring the PUT.
|
||||
func TestSettingsHandler_DeleteBeforeBootstrapIs404(t *testing.T) {
|
||||
f := newAgentNetworkHandlerFixture(t)
|
||||
|
||||
rec := f.do(t, http.MethodDelete, "/agent-network/settings", "")
|
||||
assert.Equal(t, http.StatusNotFound, rec.Code,
|
||||
"DELETE on an unbootstrapped account must 404: got %d body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
|
||||
// TestSettingsHandler_DeleteBlockedByProviders pins the first delete guard:
|
||||
// while any provider exists for the account, the delete is refused with 412
|
||||
// and the row survives. Providers route through the endpoint — the guard
|
||||
// keeps DELETE a bootstrap-repair operation rather than a way to abandon a
|
||||
// configured gateway.
|
||||
func TestSettingsHandler_DeleteBlockedByProviders(t *testing.T) {
|
||||
f := newAgentNetworkHandlerFixture(t)
|
||||
|
||||
rec := f.do(t, http.MethodPost, "/agent-network/settings", `{"proxy_address": "eu.proxy.netbird.io"}`)
|
||||
require.Equal(t, http.StatusOK, rec.Code, "bootstrap POST must succeed: %s", rec.Body.String())
|
||||
var before api.AgentNetworkSettings
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &before))
|
||||
|
||||
f.seedProvider(t, "prov-guard")
|
||||
|
||||
rec = f.do(t, http.MethodDelete, "/agent-network/settings", "")
|
||||
assert.Equal(t, http.StatusPreconditionFailed, rec.Code,
|
||||
"delete with a provider present must be refused: got %d body=%s", rec.Code, rec.Body.String())
|
||||
|
||||
rec = f.do(t, http.MethodGet, "/agent-network/settings", "")
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
var got api.AgentNetworkSettings
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got))
|
||||
assert.Equal(t, before.Endpoint, got.Endpoint, "the refused delete must leave the row intact")
|
||||
}
|
||||
|
||||
// TestSettingsHandler_DeleteBlockedByActiveProxy pins the second delete
|
||||
// guard: while a proxy is actively serving the endpoint — an active proxy
|
||||
// row declaring the endpoint hostname as its cluster address, the dedicated
|
||||
// shape — the delete is refused with 412. A proxy that has disconnected no
|
||||
// longer blocks: the guard is about a live serving path, not history.
|
||||
//
|
||||
// The proxy declares its address with mixed casing on purpose: Connect
|
||||
// stores the declared address verbatim while the settings row is normalized
|
||||
// lowercase, and hostnames are case-insensitive, so the guard must match
|
||||
// across the casing difference rather than be sidestepped by it.
|
||||
func TestSettingsHandler_DeleteBlockedByActiveProxy(t *testing.T) {
|
||||
f := newAgentNetworkHandlerFixture(t)
|
||||
|
||||
const endpoint = "gw.dedicated.example.com"
|
||||
rec := f.do(t, http.MethodPost, "/agent-network/settings", fmt.Sprintf(`{"endpoint": %q}`, endpoint))
|
||||
require.Equal(t, http.StatusOK, rec.Code, "bootstrap POST must succeed: %s", rec.Body.String())
|
||||
|
||||
now := time.Now()
|
||||
accountID := testAccountID
|
||||
proxyRow := &rpproxy.Proxy{
|
||||
ID: "proxy-guard",
|
||||
SessionID: "sess-1",
|
||||
ClusterAddress: "GW.Dedicated.Example.Com",
|
||||
AccountID: &accountID,
|
||||
LastSeen: now,
|
||||
ConnectedAt: &now,
|
||||
Status: rpproxy.StatusConnected,
|
||||
}
|
||||
require.NoError(t, f.store.SaveProxy(context.Background(), proxyRow))
|
||||
|
||||
rec = f.do(t, http.MethodDelete, "/agent-network/settings", "")
|
||||
assert.Equal(t, http.StatusPreconditionFailed, rec.Code,
|
||||
"delete with an active proxy at the endpoint must be refused: got %d body=%s", rec.Code, rec.Body.String())
|
||||
|
||||
// Once the proxy disconnects it no longer serves the endpoint, so the
|
||||
// delete goes through.
|
||||
require.NoError(t, f.store.DisconnectProxy(context.Background(), proxyRow.ID, proxyRow.SessionID))
|
||||
rec = f.do(t, http.MethodDelete, "/agent-network/settings", "")
|
||||
assert.Equal(t, http.StatusOK, rec.Code,
|
||||
"delete after the proxy disconnected must succeed: got %d body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
|
||||
// TestSettingsHandler_DeleteReleasesEndpointForFreshBootstrap pins the
|
||||
// full-reset semantic that gives replace-on-change clients (e.g. Terraform's
|
||||
// RequiresReplace) a real path: with both guards clear the delete succeeds,
|
||||
// the account reads as the defaults again, and a fresh bootstrap draws a
|
||||
// fresh label. The released hostname is not reserved — a fresh draw may even
|
||||
// legitimately re-pick it — so the assertions check the new row's shape, not
|
||||
// that the label differs.
|
||||
func TestSettingsHandler_DeleteReleasesEndpointForFreshBootstrap(t *testing.T) {
|
||||
f := newAgentNetworkHandlerFixture(t)
|
||||
|
||||
rec := f.do(t, http.MethodPost, "/agent-network/settings",
|
||||
`{"proxy_address": "eu.proxy.netbird.io", "enable_prompt_collection": true}`)
|
||||
require.Equal(t, http.StatusOK, rec.Code, "bootstrap POST must succeed: %s", rec.Body.String())
|
||||
|
||||
rec = f.do(t, http.MethodDelete, "/agent-network/settings", "")
|
||||
require.Equal(t, http.StatusOK, rec.Code,
|
||||
"delete with both guards clear must succeed: got %d body=%s", rec.Code, rec.Body.String())
|
||||
|
||||
rec = f.do(t, http.MethodGet, "/agent-network/settings", "")
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
var after api.AgentNetworkSettings
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &after))
|
||||
assert.Empty(t, after.Endpoint, "a deleted account must read as unbootstrapped defaults")
|
||||
assert.False(t, after.EnablePromptCollection, "the deleted row's toggles must not linger")
|
||||
|
||||
rec = f.do(t, http.MethodPost, "/agent-network/settings", `{"proxy_address": "eu.proxy.netbird.io"}`)
|
||||
require.Equal(t, http.StatusOK, rec.Code, "re-bootstrap after delete must succeed: %s", rec.Body.String())
|
||||
var second api.AgentNetworkSettings
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &second))
|
||||
require.NotEmpty(t, second.Endpoint, "the fresh bootstrap must allocate an endpoint")
|
||||
assert.True(t, strings.HasSuffix(second.Endpoint, ".eu.proxy.netbird.io"),
|
||||
"the fresh endpoint must hang beneath the requested proxy address: %s", second.Endpoint)
|
||||
assert.False(t, second.EnablePromptCollection,
|
||||
"the fresh row must carry bootstrap defaults, not the deleted row's toggles")
|
||||
assert.NotNil(t, second.CreatedAt, "the fresh row is persisted and carries timestamps")
|
||||
}
|
||||
|
||||
@@ -0,0 +1,37 @@
|
||||
package labelgen
|
||||
|
||||
// adjectives is the descriptor half of a generated label. It pairs with the
|
||||
// noun pool in words.go to form `<adjective>-<noun>` labels, and is kept
|
||||
// separate because words.go is almost entirely nouns — drawing both halves
|
||||
// from it produced unreadable pairs like "millet-hammock". Entries are
|
||||
// lowercase ASCII, 4-12 chars, free of hyphens and digits, screened for
|
||||
// offensive/brand/region-specific terms, and disjoint from the noun pool
|
||||
// (enforced by TestAdjectives_AreDisjointFromNouns).
|
||||
var adjectives = []string{
|
||||
"able", "active", "adept", "agile", "airy", "alert", "amiable", "ample",
|
||||
"ancient", "ardent", "artful", "astute", "balmy", "blithe", "bold", "bonny",
|
||||
"brave", "breezy", "brisk", "bubbly", "buoyant", "bushy", "candid", "canny",
|
||||
"cheery", "chilly", "chipper", "chunky", "civil", "classic", "clever", "comely",
|
||||
"compact", "cordial", "cosmic", "courtly", "crafty", "creamy", "crisp", "cuddly",
|
||||
"curious", "dainty", "dapper", "daring", "dashing", "deft", "dewy", "diligent",
|
||||
"downy", "dreamy", "dulcet", "durable", "dusky", "eager", "earnest", "earthy",
|
||||
"easy", "elated", "elegant", "epic", "fabled", "faithful", "fancy", "fearless",
|
||||
"feisty", "fervent", "fleet", "fluffy", "fond", "frisky", "frosty", "gallant",
|
||||
"genial", "genteel", "gentle", "giddy", "gilded", "glad", "glassy", "gleaming",
|
||||
"glossy", "graceful", "grand", "grainy", "hale", "hardy", "hearty", "hefty",
|
||||
"honest", "hopeful", "humble", "hushed", "immense", "jaunty", "jolly", "jovial",
|
||||
"joyful", "jubilant", "keen", "kindly", "kindred", "lanky", "leafy", "limber",
|
||||
"lively", "lofty", "loyal", "lucent", "lucid", "luminous", "lush", "maroon",
|
||||
"mellow", "merry", "mighty", "mindful", "mirthful", "misty", "modest", "muted",
|
||||
"nifty", "nimble", "noble", "patient", "peaceful", "pearly", "peppy", "perky",
|
||||
"petite", "placid", "playful", "pleasant", "plucky", "plush", "polite", "posh",
|
||||
"prancing", "pristine", "prompt", "proud", "prudent", "quaint", "quick", "quirky",
|
||||
"radiant", "ready", "regal", "restful", "robust", "rosy", "ruddy", "rugged",
|
||||
"sandy", "satin", "saucy", "savvy", "sedate", "serene", "shady", "shiny",
|
||||
"silken", "silky", "sincere", "sleek", "slender", "smart", "smooth", "snappy",
|
||||
"snug", "soaring", "sparkly", "spiffy", "spirited", "sprightly", "spry", "stalwart",
|
||||
"stately", "steady", "sterling", "stoic", "stormy", "stout", "sturdy", "sunlit",
|
||||
"supple", "svelte", "tawny", "tender", "tidy", "timeless", "trusty", "upbeat",
|
||||
"urbane", "valiant", "vast", "vernal", "vibrant", "vintage", "whimsy", "willing",
|
||||
"windy", "winsome", "wintry", "witty", "worthy", "zesty", "zippy",
|
||||
}
|
||||
@@ -64,3 +64,20 @@ func PickUnique(rng *rand.Rand, taken map[string]struct{}, fallbackSuffix string
|
||||
w := pool[rng.Intn(len(pool))]
|
||||
return fmt.Sprintf("%s-%s", w, fallbackSuffix)
|
||||
}
|
||||
|
||||
// PickTuple returns an adjective-noun label such as "brave-otter". It is still
|
||||
// a single DNS label.
|
||||
//
|
||||
// Unlike PickUnique it takes no `taken` set and has no fallback suffix. The
|
||||
// noun pool holds 857 entries, which is ample per cluster but a hard ceiling
|
||||
// once labels must be unique across one shared zone; pairing an adjective with
|
||||
// a noun spans len(adjectives) * 857 instead. Uniqueness is enforced by a
|
||||
// database constraint and retried by the caller, rather than guessed from a
|
||||
// pre-read set that a concurrent allocation can invalidate.
|
||||
func PickTuple(rng *rand.Rand) string {
|
||||
nouns := uniqueWords()
|
||||
if len(nouns) == 0 || len(adjectives) == 0 {
|
||||
return ""
|
||||
}
|
||||
return adjectives[rng.Intn(len(adjectives))] + "-" + nouns[rng.Intn(len(nouns))]
|
||||
}
|
||||
|
||||
@@ -99,3 +99,82 @@ func TestUniqueWords_DropsDuplicates(t *testing.T) {
|
||||
}
|
||||
assert.GreaterOrEqual(t, len(pool), 500, "Pool must contain at least 500 unique words")
|
||||
}
|
||||
|
||||
// TestPickTuple_ShapeAndPoolMembership locks the wire-visible shape: an
|
||||
// adjective and a noun, each from its own pool, joined by a single hyphen so
|
||||
// the result stays one DNS label.
|
||||
func TestPickTuple_ShapeAndPoolMembership(t *testing.T) {
|
||||
nouns := uniqueWords()
|
||||
inNouns := make(map[string]struct{}, len(nouns))
|
||||
for _, w := range nouns {
|
||||
inNouns[w] = struct{}{}
|
||||
}
|
||||
inAdjectives := make(map[string]struct{}, len(adjectives))
|
||||
for _, a := range adjectives {
|
||||
inAdjectives[a] = struct{}{}
|
||||
}
|
||||
|
||||
rng := rand.New(rand.NewSource(7))
|
||||
for i := 0; i < 200; i++ {
|
||||
got := PickTuple(rng)
|
||||
|
||||
parts := strings.Split(got, "-")
|
||||
require.Len(t, parts, 2, "PickTuple must produce exactly two hyphen-joined words; got %q", got)
|
||||
|
||||
_, adjOK := inAdjectives[parts[0]]
|
||||
assert.True(t, adjOK, "First half must be an adjective; %q not in adjectives (from %q)", parts[0], got)
|
||||
_, nounOK := inNouns[parts[1]]
|
||||
assert.True(t, nounOK, "Second half must be a noun; %q not in words (from %q)", parts[1], got)
|
||||
|
||||
assert.LessOrEqual(t, len(got), 63, "Label must fit a DNS label; got %q (%d chars)", got, len(got))
|
||||
}
|
||||
}
|
||||
|
||||
// TestAdjectives_AreDisjointFromNouns keeps the namespace a clean product and
|
||||
// prevents nonsense like "azure-azure": a handful of the noun pool's entries
|
||||
// are adjectival, and any overlap would let the same word land on both sides.
|
||||
func TestAdjectives_AreDisjointFromNouns(t *testing.T) {
|
||||
nouns := make(map[string]struct{}, len(uniqueWords()))
|
||||
for _, w := range uniqueWords() {
|
||||
nouns[w] = struct{}{}
|
||||
}
|
||||
for _, a := range adjectives {
|
||||
_, clash := nouns[a]
|
||||
assert.False(t, clash, "Adjective %q also appears in the noun pool; remove it from one list", a)
|
||||
}
|
||||
}
|
||||
|
||||
// TestAdjectives_AreDNSSafeAndDeduplicated mirrors the curation contract stated
|
||||
// in words.go: lowercase ASCII, 4-12 chars, no digits or hyphens, no repeats.
|
||||
func TestAdjectives_AreDNSSafeAndDeduplicated(t *testing.T) {
|
||||
seen := make(map[string]struct{}, len(adjectives))
|
||||
for _, a := range adjectives {
|
||||
_, dup := seen[a]
|
||||
assert.False(t, dup, "Duplicate adjective %q", a)
|
||||
seen[a] = struct{}{}
|
||||
|
||||
assert.Regexp(t, `^[a-z]{4,12}$`, a, "Adjective %q must be 4-12 lowercase ASCII letters", a)
|
||||
}
|
||||
assert.Greater(t, len(adjectives), 150, "Adjective pool too small to give a useful namespace")
|
||||
}
|
||||
|
||||
// TestPickTuple_DeterministicWithSeededRng documents that generation is a pure
|
||||
// function of the rng, which is what makes allocation retries reproducible in tests.
|
||||
func TestPickTuple_DeterministicWithSeededRng(t *testing.T) {
|
||||
a := PickTuple(rand.New(rand.NewSource(42)))
|
||||
b := PickTuple(rand.New(rand.NewSource(42)))
|
||||
assert.Equal(t, a, b, "Same seed must yield the same tuple")
|
||||
}
|
||||
|
||||
// TestPickTuple_SpansALargeNamespace guards the reason we moved to tuples: a
|
||||
// single-word pool caps the GLOBAL namespace at 857. Drawing many tuples must
|
||||
// yield overwhelmingly distinct values.
|
||||
func TestPickTuple_SpansALargeNamespace(t *testing.T) {
|
||||
rng := rand.New(rand.NewSource(11))
|
||||
seen := make(map[string]struct{}, 2000)
|
||||
for i := 0; i < 2000; i++ {
|
||||
seen[PickTuple(rng)] = struct{}{}
|
||||
}
|
||||
assert.Greater(t, len(seen), 1900,
|
||||
"2000 draws should be nearly all distinct across a ~200k namespace; got %d unique", len(seen))
|
||||
}
|
||||
|
||||
@@ -22,7 +22,6 @@ import (
|
||||
"github.com/netbirdio/netbird/management/server/permissions/modules"
|
||||
"github.com/netbirdio/netbird/management/server/permissions/operations"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
@@ -48,7 +47,7 @@ func ensureSessionKeys(p *types.Provider) error {
|
||||
type Manager interface {
|
||||
GetAllProviders(ctx context.Context, accountID, userID string) ([]*types.Provider, error)
|
||||
GetProvider(ctx context.Context, accountID, userID, providerID string) (*types.Provider, error)
|
||||
CreateProvider(ctx context.Context, userID string, provider *types.Provider, bootstrapCluster string) (*types.Provider, error)
|
||||
CreateProvider(ctx context.Context, userID string, provider *types.Provider) (*types.Provider, error)
|
||||
UpdateProvider(ctx context.Context, userID string, provider *types.Provider) (*types.Provider, error)
|
||||
DeleteProvider(ctx context.Context, accountID, userID, providerID string) error
|
||||
|
||||
@@ -71,7 +70,9 @@ type Manager interface {
|
||||
DeleteBudgetRule(ctx context.Context, accountID, userID, ruleID string) error
|
||||
|
||||
GetSettings(ctx context.Context, accountID, userID string) (*types.Settings, error)
|
||||
CreateSettings(ctx context.Context, userID string, settings *types.Settings, proxyAddress, endpoint string) (*types.Settings, error)
|
||||
UpdateSettings(ctx context.Context, userID string, settings *types.Settings) (*types.Settings, error)
|
||||
DeleteSettings(ctx context.Context, accountID, userID string) error
|
||||
|
||||
ListConsumption(ctx context.Context, accountID, userID string) ([]*types.Consumption, error)
|
||||
ListAccessLogs(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter) ([]*types.AgentNetworkAccessLog, int64, error)
|
||||
@@ -123,11 +124,10 @@ type managerImpl struct {
|
||||
proxyController proxy.Controller
|
||||
|
||||
// reconcileCache holds the last set of synthesised proxy mappings
|
||||
// per account so reconcile can emit precise Create/Update/Delete
|
||||
// updates instead of a full re-push on every mutation. Keyed by
|
||||
// accountID, then by synthesised service ID.
|
||||
// per account, each paired with the proxy that served it, so a change
|
||||
// of serving proxy can be diffed without re-deriving it.
|
||||
reconcileMu sync.Mutex
|
||||
reconcileCache map[string]map[string]*proto.ProxyMapping
|
||||
reconcileCache map[string]map[string]syntheticMapping
|
||||
|
||||
// labelRngMu guards labelRng. PickUnique consumes math/rand.Source
|
||||
// state; concurrent provider creates would otherwise race.
|
||||
@@ -151,7 +151,7 @@ func NewManager(
|
||||
accountManager: accountManager,
|
||||
permissionsManager: permissionsManager,
|
||||
proxyController: proxyController,
|
||||
reconcileCache: make(map[string]map[string]*proto.ProxyMapping),
|
||||
reconcileCache: make(map[string]map[string]syntheticMapping),
|
||||
labelRng: rand.New(rand.NewSource(time.Now().UnixNano())),
|
||||
}
|
||||
}
|
||||
@@ -170,19 +170,14 @@ func (m *managerImpl) GetProvider(ctx context.Context, accountID, userID, provid
|
||||
return m.store.GetAgentNetworkProviderByID(ctx, store.LockingStrengthNone, accountID, providerID)
|
||||
}
|
||||
|
||||
// CreateProvider persists a new provider for the account. bootstrapCluster
|
||||
// is used only when the per-account agent-network Settings row hasn't
|
||||
// been created yet; otherwise it is ignored (the cluster is pinned on
|
||||
// Settings and every provider in the account routes through it).
|
||||
func (m *managerImpl) CreateProvider(ctx context.Context, userID string, provider *types.Provider, bootstrapCluster string) (*types.Provider, error) {
|
||||
// CreateProvider persists a new provider for the account. Providers have no
|
||||
// settings side effects: the account's endpoint is bootstrapped separately and
|
||||
// explicitly via CreateSettings, and every provider in the account routes
|
||||
// through it.
|
||||
func (m *managerImpl) CreateProvider(ctx context.Context, userID string, provider *types.Provider) (*types.Provider, error) {
|
||||
if err := m.requirePermission(ctx, provider.AccountID, userID, modules.AgentNetworkProviders, operations.Create); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if strings.TrimSpace(bootstrapCluster) != "" {
|
||||
if err := m.requireSettingsBootstrapPermission(ctx, provider.AccountID, userID); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
// An empty api_key would silently produce a synthesised service
|
||||
// that 401s on every upstream request. Surface the misconfiguration
|
||||
@@ -206,16 +201,6 @@ func (m *managerImpl) CreateProvider(ctx context.Context, userID string, provide
|
||||
return nil, fmt.Errorf("save agent network provider: %w", err)
|
||||
}
|
||||
|
||||
if strings.TrimSpace(bootstrapCluster) != "" {
|
||||
if _, err := m.bootstrapSettingsIfNeeded(ctx, m.store, provider.AccountID, bootstrapCluster); err != nil {
|
||||
// The provider create has already succeeded; logging the
|
||||
// bootstrap miss matches the plan's PoC behaviour. The synth
|
||||
// path treats a missing settings row as a no-op, and the next
|
||||
// provider create retries the bootstrap.
|
||||
log.WithContext(ctx).Debugf("agent-network bootstrap settings for account %s on cluster %s: %v", provider.AccountID, bootstrapCluster, err)
|
||||
}
|
||||
}
|
||||
|
||||
m.accountManager.StoreEvent(ctx, userID, provider.ID, provider.AccountID, activity.AgentNetworkProviderCreated, provider.EventMeta())
|
||||
m.reconcile(ctx, provider.AccountID)
|
||||
|
||||
@@ -560,52 +545,44 @@ func (m *managerImpl) DeleteBudgetRule(ctx context.Context, accountID, userID, r
|
||||
}
|
||||
|
||||
// UpdateSettings replaces the mutable account-level settings — the collection
|
||||
// toggles and retention — on the account's row. When the account has no
|
||||
// settings row yet, a non-empty settings.Cluster bootstraps one (same path as
|
||||
// first provider create); without it the update fails with NotFound. On an
|
||||
// existing row the cluster and subdomain are immutable: a differing
|
||||
// settings.Cluster is rejected rather than silently ignored so callers never
|
||||
// observe a value other than what they sent. Because the collection toggles
|
||||
// change the synthesised service config (prompt-capture gating, access-log
|
||||
// emission), a reconcile is triggered so the proxy and peer network maps
|
||||
// converge on the new state.
|
||||
// toggles and retention — on the account's row. The identity fields (Domain,
|
||||
// ProxyAddress) are assigned at bootstrap (CreateSettings) and immutable: the
|
||||
// request carries them, matching the PUT convention of every other endpoint,
|
||||
// but they are only compared against the stored row — a request carrying
|
||||
// different values is rejected, and the stored values are never overwritten.
|
||||
// When the account has no settings row yet the update fails with NotFound.
|
||||
// Because the collection toggles change the synthesised service config
|
||||
// (prompt-capture gating, access-log emission), a reconcile is triggered so
|
||||
// the proxy and peer network maps converge on the new state.
|
||||
func (m *managerImpl) UpdateSettings(ctx context.Context, userID string, settings *types.Settings) (*types.Settings, error) {
|
||||
if err := m.requirePermission(ctx, settings.AccountID, userID, modules.AgentNetworkSettings, operations.Update); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
requestedCluster := strings.TrimSpace(settings.Cluster)
|
||||
|
||||
// The row lock from LockingStrengthUpdate only holds for the duration of
|
||||
// the surrounding transaction, so the read, the cluster-immutability
|
||||
// check, and the save must share one — otherwise concurrent PUTs could
|
||||
// interleave between them.
|
||||
// the surrounding transaction, so the read and the save must share one —
|
||||
// otherwise concurrent PUTs could interleave between them.
|
||||
var updated *types.Settings
|
||||
err := m.store.ExecuteInTransaction(ctx, func(tx store.Store) error {
|
||||
existing, err := tx.GetAgentNetworkSettings(ctx, store.LockingStrengthUpdate, settings.AccountID)
|
||||
switch {
|
||||
case err == nil:
|
||||
if requestedCluster != "" && requestedCluster != existing.Cluster {
|
||||
return status.Errorf(status.InvalidArgument, "cluster is immutable once assigned (current: %s)", existing.Cluster)
|
||||
}
|
||||
case isNotFound(err):
|
||||
if requestedCluster == "" {
|
||||
return status.Errorf(status.NotFound, "agent network settings have not been bootstrapped yet; pass cluster to bootstrap them, or create a provider with bootstrap_cluster set")
|
||||
}
|
||||
// Bootstrapping pins the cluster and subdomain — a settings
|
||||
// create on top of the update the caller already passed, matching
|
||||
// the gate on the provider-create bootstrap path.
|
||||
if err := m.requirePermission(ctx, settings.AccountID, userID, modules.AgentNetworkSettings, operations.Create); err != nil {
|
||||
return err
|
||||
}
|
||||
existing, err = m.bootstrapSettingsIfNeeded(ctx, tx, settings.AccountID, requestedCluster)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
return status.Errorf(status.NotFound, "agent network settings have not been bootstrapped yet; POST /api/agent-network/settings to bootstrap them")
|
||||
default:
|
||||
return fmt.Errorf("get agent network settings: %w", err)
|
||||
}
|
||||
|
||||
// The identity echo is compared leniently (trimmed, case-insensitive):
|
||||
// the stored values are normalized lowercase, and a client replaying a
|
||||
// GET response must never be rejected over casing it didn't choose.
|
||||
if !hostnamesEquivalent(settings.Domain, existing.Domain) {
|
||||
return status.Errorf(status.InvalidArgument, "endpoint is immutable: it must match the assigned endpoint %q; delete the settings to release it and bootstrap again", existing.Domain)
|
||||
}
|
||||
if !hostnamesEquivalent(settings.ProxyAddress, existing.ProxyAddress) {
|
||||
return status.Errorf(status.InvalidArgument, "proxy_address is immutable: it must match the assigned proxy address %q; delete the settings to release it and bootstrap again", existing.ProxyAddress)
|
||||
}
|
||||
|
||||
existing.EnableLogCollection = settings.EnableLogCollection
|
||||
existing.EnablePromptCollection = settings.EnablePromptCollection
|
||||
existing.RedactPii = settings.RedactPii
|
||||
@@ -632,6 +609,83 @@ func (m *managerImpl) UpdateSettings(ctx context.Context, userID string, setting
|
||||
return updated, nil
|
||||
}
|
||||
|
||||
// hostnamesEquivalent reports whether a caller-supplied hostname names the
|
||||
// same host as a stored (normalized, lowercase) one: equal after trimming and
|
||||
// case folding. No structural validation — an arbitrary mismatch and a
|
||||
// malformed value are both simply "not the assigned value".
|
||||
func hostnamesEquivalent(supplied, stored string) bool {
|
||||
return strings.EqualFold(strings.TrimSpace(supplied), stored)
|
||||
}
|
||||
|
||||
// DeleteSettings removes the account's settings row, releasing the endpoint.
|
||||
// Two guards make this a bootstrap-repair operation rather than a way to tear
|
||||
// down a serving gateway, both re-checked under the row lock:
|
||||
//
|
||||
// - No Agent Network providers may exist for the account. Providers route
|
||||
// through the endpoint; delete them first.
|
||||
// - No proxy may be actively serving the endpoint — that is, no active proxy
|
||||
// declares the endpoint hostname as its cluster address. This is the
|
||||
// dedicated (self-addressed) shape's guard: the proxy at the address IS
|
||||
// this account's gateway. A labeled endpoint hangs beneath a shared
|
||||
// cluster's address, and with the account's providers already gone the
|
||||
// shared proxy serves nothing of the account's, so the parent cluster
|
||||
// being up does not block the delete.
|
||||
//
|
||||
// Bootstrapping again after a delete allocates fresh — the released hostname
|
||||
// is not reserved. That full-reset semantic is what gives clients that model
|
||||
// immutability as replace-on-change (e.g. Terraform's RequiresReplace) a real
|
||||
// path: tear down providers, delete, re-create.
|
||||
func (m *managerImpl) DeleteSettings(ctx context.Context, accountID, userID string) error {
|
||||
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkSettings, operations.Delete); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
var deleted *types.Settings
|
||||
err := m.store.ExecuteInTransaction(ctx, func(tx store.Store) error {
|
||||
existing, err := tx.GetAgentNetworkSettings(ctx, store.LockingStrengthUpdate, accountID)
|
||||
switch {
|
||||
case err == nil:
|
||||
case isNotFound(err):
|
||||
return status.Errorf(status.NotFound, "agent network settings have not been bootstrapped yet; there is nothing to delete")
|
||||
default:
|
||||
return fmt.Errorf("get agent network settings: %w", err)
|
||||
}
|
||||
|
||||
providers, err := tx.GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, accountID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("get agent network providers: %w", err)
|
||||
}
|
||||
if len(providers) > 0 {
|
||||
return status.Errorf(status.PreconditionFailed, "agent network settings cannot be deleted while %d provider(s) exist; delete the providers first", len(providers))
|
||||
}
|
||||
|
||||
serving, err := tx.HasActiveProxyAtClusterAddress(ctx, existing.Domain)
|
||||
if err != nil {
|
||||
return fmt.Errorf("check for a proxy serving the endpoint: %w", err)
|
||||
}
|
||||
if serving {
|
||||
return status.Errorf(status.PreconditionFailed, "agent network settings cannot be deleted while a proxy is actively serving the endpoint %q", existing.Domain)
|
||||
}
|
||||
|
||||
if err := tx.DeleteAgentNetworkSettings(ctx, accountID); err != nil {
|
||||
return fmt.Errorf("delete agent network settings: %w", err)
|
||||
}
|
||||
deleted = existing
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
m.accountManager.StoreEvent(ctx, userID, accountID, accountID, activity.AgentNetworkSettingsDeleted, map[string]any{
|
||||
"endpoint": deleted.Domain,
|
||||
"proxy_address": deleted.ProxyAddress,
|
||||
})
|
||||
m.reconcile(ctx, accountID)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// isNotFound reports whether err is a status.NotFound error.
|
||||
func isNotFound(err error) bool {
|
||||
var sErr *status.Error
|
||||
@@ -678,74 +732,162 @@ func (m *managerImpl) GetSettings(ctx context.Context, accountID, userID string)
|
||||
}
|
||||
}
|
||||
|
||||
// requireSettingsBootstrapPermission gates the one-time settings bootstrap a
|
||||
// first provider create performs. Pinning the account's cluster and subdomain
|
||||
// is a settings write, so it needs the settings permission on top of the
|
||||
// provider one. No-op once the settings row exists.
|
||||
func (m *managerImpl) requireSettingsBootstrapPermission(ctx context.Context, accountID, userID string) error {
|
||||
_, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID)
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
if !isNotFound(err) {
|
||||
return fmt.Errorf("get agent network settings: %w", err)
|
||||
}
|
||||
return m.requirePermission(ctx, accountID, userID, modules.AgentNetworkSettings, operations.Create)
|
||||
}
|
||||
// maxDomainAllocationAttempts bounds the label search when bootstrapping a
|
||||
// labeled endpoint. Package-level (rather than function-local) so tests can
|
||||
// assert on the exhaustion path without duplicating the literal.
|
||||
const maxDomainAllocationAttempts = 10
|
||||
|
||||
// bootstrapSettingsIfNeeded creates the per-account agent-network
|
||||
// settings row when missing. The cluster comes from the create-time
|
||||
// hint the dashboard sends (auto-picked from the active cluster list);
|
||||
// the subdomain is picked from the curated wordlist avoiding
|
||||
// collisions on the same cluster. Idempotent: if a row already exists
|
||||
// it is returned untouched and the hint is ignored. st is the store to
|
||||
// operate on — pass the transaction store when calling from within one.
|
||||
func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, st store.Store, accountID, providerCluster string) (*types.Settings, error) {
|
||||
if accountID == "" {
|
||||
return nil, fmt.Errorf("bootstrap settings: account id is required")
|
||||
// CreateSettings bootstraps the per-account settings row, assigning the
|
||||
// account's immutable endpoint. Exactly one of proxyAddress and endpoint must
|
||||
// be non-empty: proxyAddress allocates a labeled endpoint one label beneath
|
||||
// the given cluster address; endpoint claims the given hostname verbatim as a
|
||||
// self-addressed (dedicated) endpoint — a legitimate claim before any proxy
|
||||
// declares the address (address-first). settings carries the account ID and
|
||||
// the initial collection toggles; its identity fields are assigned here.
|
||||
func (m *managerImpl) CreateSettings(ctx context.Context, userID string, settings *types.Settings, proxyAddress, endpoint string) (*types.Settings, error) {
|
||||
if settings == nil || settings.AccountID == "" {
|
||||
return nil, status.Errorf(status.InvalidArgument, "account id is required")
|
||||
}
|
||||
if strings.TrimSpace(providerCluster) == "" {
|
||||
return nil, fmt.Errorf("bootstrap settings: provider cluster is required")
|
||||
if err := m.requirePermission(ctx, settings.AccountID, userID, modules.AgentNetworkSettings, operations.Create); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
existing, err := st.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID)
|
||||
if err == nil {
|
||||
return existing, nil
|
||||
hasProxyAddress := strings.TrimSpace(proxyAddress) != ""
|
||||
hasEndpoint := strings.TrimSpace(endpoint) != ""
|
||||
if hasProxyAddress == hasEndpoint {
|
||||
return nil, status.Errorf(status.InvalidArgument, "exactly one of proxy_address and endpoint is required")
|
||||
}
|
||||
if !isNotFound(err) {
|
||||
|
||||
// Fail fast on an existing row for a clean 409; the insert below stays
|
||||
// the authority against concurrent bootstraps (the primary key wins).
|
||||
if _, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, settings.AccountID); err == nil {
|
||||
return nil, status.Errorf(status.AlreadyExists, "agent network settings already bootstrapped for account %s", settings.AccountID)
|
||||
} else if !isNotFound(err) {
|
||||
return nil, fmt.Errorf("get agent network settings: %w", err)
|
||||
}
|
||||
|
||||
siblings, err := st.GetAgentNetworkSettingsByCluster(ctx, store.LockingStrengthNone, providerCluster)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list agent network settings on cluster: %w", err)
|
||||
}
|
||||
taken := make(map[string]struct{}, len(siblings))
|
||||
for _, s := range siblings {
|
||||
taken[s.Subdomain] = struct{}{}
|
||||
}
|
||||
|
||||
suffix := accountID
|
||||
if len(suffix) > 4 {
|
||||
suffix = suffix[:4]
|
||||
}
|
||||
|
||||
m.labelRngMu.Lock()
|
||||
subdomain := labelgen.PickUnique(m.labelRng, taken, suffix)
|
||||
m.labelRngMu.Unlock()
|
||||
|
||||
now := time.Now().UTC()
|
||||
settings := types.DefaultSettings(accountID)
|
||||
settings.Cluster = providerCluster
|
||||
settings.Subdomain = subdomain
|
||||
settings.CreatedAt = now
|
||||
settings.UpdatedAt = now
|
||||
if err := st.SaveAgentNetworkSettings(ctx, settings); err != nil {
|
||||
return nil, fmt.Errorf("save agent network settings: %w", err)
|
||||
|
||||
var err error
|
||||
if hasEndpoint {
|
||||
err = m.bootstrapSelfAddressed(ctx, settings, endpoint)
|
||||
} else {
|
||||
err = m.bootstrapLabeled(ctx, settings, proxyAddress)
|
||||
}
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
m.accountManager.StoreEvent(ctx, userID, settings.AccountID, settings.AccountID, activity.AgentNetworkSettingsUpdated, map[string]any{
|
||||
"bootstrapped": true,
|
||||
"endpoint": settings.Domain,
|
||||
"dedicated": settings.Dedicated(),
|
||||
})
|
||||
m.reconcile(ctx, settings.AccountID)
|
||||
|
||||
return settings, nil
|
||||
}
|
||||
|
||||
// bootstrapSelfAddressed claims the given hostname as the account's endpoint,
|
||||
// served only by a proxy declaring exactly that address (Domain ==
|
||||
// ProxyAddress). The domain unique index is the arbiter of availability.
|
||||
func (m *managerImpl) bootstrapSelfAddressed(ctx context.Context, settings *types.Settings, endpoint string) error {
|
||||
hostname, err := types.NormalizeHostname(endpoint)
|
||||
if err != nil {
|
||||
return status.Errorf(status.InvalidArgument, "invalid endpoint: %s", err)
|
||||
}
|
||||
|
||||
settings.Domain = hostname
|
||||
settings.ProxyAddress = hostname
|
||||
if err := m.store.CreateAgentNetworkSettings(ctx, settings); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
// The violation is either the account primary key (a concurrent
|
||||
// bootstrap for the same account won) or the domain index
|
||||
// (another account holds the hostname). Distinguish by re-read.
|
||||
if _, getErr := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, settings.AccountID); getErr == nil {
|
||||
return status.Errorf(status.AlreadyExists, "agent network settings already bootstrapped for account %s", settings.AccountID)
|
||||
}
|
||||
return status.Errorf(status.AlreadyExists, "endpoint %s is already taken", hostname)
|
||||
}
|
||||
return fmt.Errorf("create agent network settings: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// bootstrapLabeled allocates a labeled endpoint one label beneath the given
|
||||
// cluster address: Domain = <label>.<proxyAddress>, served by whichever proxy
|
||||
// declares the parent. Labels are adjective-noun tuples; a candidate is
|
||||
// checked by read and the domain unique index stays the authority, so a
|
||||
// concurrent allocation of the same tuple surfaces as a unique violation and
|
||||
// another tuple is drawn.
|
||||
func (m *managerImpl) bootstrapLabeled(ctx context.Context, settings *types.Settings, proxyAddress string) error {
|
||||
parent, err := types.NormalizeHostname(proxyAddress)
|
||||
if err != nil {
|
||||
return status.Errorf(status.InvalidArgument, "invalid proxy_address: %s", err)
|
||||
}
|
||||
|
||||
for attempt := 1; attempt <= maxDomainAllocationAttempts; attempt++ {
|
||||
m.labelRngMu.Lock()
|
||||
label := labelgen.PickTuple(m.labelRng)
|
||||
m.labelRngMu.Unlock()
|
||||
if label == "" {
|
||||
// Only reachable if either word pool were emptied. An empty label
|
||||
// would produce a broken endpoint like ".example.com", so fail
|
||||
// loudly rather than looping or inserting.
|
||||
return fmt.Errorf("allocate agent network endpoint for account %s: label generator returned an empty label", settings.AccountID)
|
||||
}
|
||||
|
||||
candidate, err := types.NormalizeHostname(label + "." + parent)
|
||||
if err != nil {
|
||||
return status.Errorf(status.InvalidArgument, "proxy_address leaves no room for a label: %s", err)
|
||||
}
|
||||
|
||||
_, err = m.store.GetAgentNetworkSettingsByDomain(ctx, store.LockingStrengthNone, candidate)
|
||||
if err == nil {
|
||||
log.WithContext(ctx).Tracef("agent-network endpoint %q taken, retrying (attempt %d/%d)", candidate, attempt, maxDomainAllocationAttempts)
|
||||
continue
|
||||
}
|
||||
if !isNotFound(err) {
|
||||
return fmt.Errorf("check agent network endpoint availability: %w", err)
|
||||
}
|
||||
|
||||
settings.Domain = candidate
|
||||
settings.ProxyAddress = parent
|
||||
if err := m.store.CreateAgentNetworkSettings(ctx, settings); err != nil {
|
||||
if isUniqueConstraintError(err) {
|
||||
// A concurrent bootstrap for the same account may have won on
|
||||
// the primary key — return the conflict. A lost race on the
|
||||
// domain index just means the tuple was taken between the
|
||||
// read and the insert: draw another.
|
||||
if _, getErr := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, settings.AccountID); getErr == nil {
|
||||
return status.Errorf(status.AlreadyExists, "agent network settings already bootstrapped for account %s", settings.AccountID)
|
||||
}
|
||||
log.WithContext(ctx).Tracef("agent-network endpoint %q lost an allocation race, retrying (attempt %d/%d)", candidate, attempt, maxDomainAllocationAttempts)
|
||||
continue
|
||||
}
|
||||
return fmt.Errorf("create agent network settings: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
return fmt.Errorf("allocate agent network endpoint for account %s: %d attempts exhausted", settings.AccountID, maxDomainAllocationAttempts)
|
||||
}
|
||||
|
||||
// isUniqueConstraintError reports whether err is a database unique-constraint
|
||||
// violation, matched on the driver message because CreateAgentNetworkSettings
|
||||
// deliberately returns the driver error unwrapped.
|
||||
func isUniqueConstraintError(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
msg := err.Error()
|
||||
return strings.Contains(msg, "(SQLSTATE 23505)") || // postgres
|
||||
strings.Contains(msg, "Error 1062 (23000)") || // mysql
|
||||
strings.Contains(msg, "UNIQUE constraint failed") // sqlite
|
||||
}
|
||||
|
||||
// ListConsumption returns every consumption row recorded for the
|
||||
// account, ordered window-newest-first. Backs the dashboard's basic
|
||||
// counter view; permission gate is the same Read role that gates
|
||||
@@ -879,7 +1021,7 @@ func (*mockManager) GetProvider(_ context.Context, _, _, _ string) (*types.Provi
|
||||
return &types.Provider{}, nil
|
||||
}
|
||||
|
||||
func (*mockManager) CreateProvider(_ context.Context, _ string, p *types.Provider, _ string) (*types.Provider, error) {
|
||||
func (*mockManager) CreateProvider(_ context.Context, _ string, p *types.Provider) (*types.Provider, error) {
|
||||
return p, nil
|
||||
}
|
||||
|
||||
@@ -947,10 +1089,23 @@ func (*mockManager) GetSettings(_ context.Context, accountID, _ string) (*types.
|
||||
return types.DefaultSettings(accountID), nil
|
||||
}
|
||||
|
||||
func (*mockManager) CreateSettings(_ context.Context, _ string, s *types.Settings, proxyAddress, endpoint string) (*types.Settings, error) {
|
||||
if endpoint != "" {
|
||||
s.Domain = endpoint
|
||||
s.ProxyAddress = endpoint
|
||||
} else {
|
||||
s.Domain = "mock." + proxyAddress
|
||||
s.ProxyAddress = proxyAddress
|
||||
}
|
||||
return s, nil
|
||||
}
|
||||
|
||||
func (*mockManager) UpdateSettings(_ context.Context, _ string, s *types.Settings) (*types.Settings, error) {
|
||||
return s, nil
|
||||
}
|
||||
|
||||
func (*mockManager) DeleteSettings(_ context.Context, _, _ string) error { return nil }
|
||||
|
||||
func (*mockManager) ListConsumption(_ context.Context, _, _ string) ([]*types.Consumption, error) {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
@@ -1,134 +0,0 @@
|
||||
package agentnetwork
|
||||
|
||||
import (
|
||||
"context"
|
||||
"runtime"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
"github.com/netbirdio/netbird/management/server/account"
|
||||
"github.com/netbirdio/netbird/management/server/permissions"
|
||||
"github.com/netbirdio/netbird/management/server/permissions/modules"
|
||||
"github.com/netbirdio/netbird/management/server/permissions/operations"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
nbtypes "github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
// bootstrapFixture wires a real sqlite store to a gomock permissions manager
|
||||
// so tests can grant the provider permission while denying (or never
|
||||
// expecting) the settings one.
|
||||
type bootstrapFixture struct {
|
||||
manager Manager
|
||||
store store.Store
|
||||
perms *permissions.MockManager
|
||||
}
|
||||
|
||||
func newBootstrapFixture(t *testing.T) *bootstrapFixture {
|
||||
t.Helper()
|
||||
if runtime.GOOS == "windows" {
|
||||
t.Skip("sqlite store not properly supported on Windows yet")
|
||||
}
|
||||
t.Setenv("NETBIRD_STORE_ENGINE", string(nbtypes.SqliteStoreEngine))
|
||||
|
||||
st, cleanUp, err := store.NewTestStoreFromSQL(context.Background(), "", t.TempDir())
|
||||
require.NoError(t, err, "test store setup must succeed")
|
||||
t.Cleanup(cleanUp)
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
perms := permissions.NewMockManager(ctrl)
|
||||
|
||||
accounts := account.NewMockManager(ctrl)
|
||||
accounts.EXPECT().StoreEvent(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes()
|
||||
accounts.EXPECT().UpdateAccountPeers(gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes()
|
||||
accounts.EXPECT().BufferUpdateAccountPeers(gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes()
|
||||
|
||||
return &bootstrapFixture{
|
||||
manager: NewManager(st, perms, accounts, nil),
|
||||
store: st,
|
||||
perms: perms,
|
||||
}
|
||||
}
|
||||
|
||||
func (f *bootstrapFixture) expectPermission(accountID, userID string, module modules.Module, op operations.Operation, allowed bool) {
|
||||
f.perms.EXPECT().
|
||||
ValidateUserPermissions(gomock.Any(), accountID, userID, module, op).
|
||||
Return(allowed, context.Background(), nil)
|
||||
}
|
||||
|
||||
func newBootstrapProvider(accountID string) *types.Provider {
|
||||
p := types.NewProvider(accountID)
|
||||
p.Name = "openai"
|
||||
p.UpstreamURL = "https://api.openai.com"
|
||||
p.APIKey = "sk-test"
|
||||
p.Enabled = true
|
||||
return p
|
||||
}
|
||||
|
||||
// TestCreateProviderBootstrapRequiresSettingsPermission pins the gate on the
|
||||
// one-time settings bootstrap: creating the first provider with a
|
||||
// bootstrap_cluster pins the account's cluster and subdomain, which is a
|
||||
// settings write and must not ride on the providers permission alone.
|
||||
func TestCreateProviderBootstrapRequiresSettingsPermission(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
t.Run("denied without settings permission", func(t *testing.T) {
|
||||
f := newBootstrapFixture(t)
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, false)
|
||||
|
||||
_, err := f.manager.CreateProvider(ctx, "user1", newBootstrapProvider("account1"), "cluster1.example.com")
|
||||
require.Error(t, err, "bootstrap without settings permission must fail")
|
||||
var sErr *status.Error
|
||||
require.ErrorAs(t, err, &sErr)
|
||||
assert.Equal(t, status.PermissionDenied, sErr.Type(), "denial should surface as permission denied")
|
||||
|
||||
providers, err := f.store.GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, "account1")
|
||||
require.NoError(t, err)
|
||||
assert.Empty(t, providers, "provider must not be persisted when bootstrap is denied")
|
||||
_, err = f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1")
|
||||
assert.Error(t, err, "settings row must not be created when bootstrap is denied")
|
||||
})
|
||||
|
||||
t.Run("allowed with settings permission", func(t *testing.T) {
|
||||
f := newBootstrapFixture(t)
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
|
||||
created, err := f.manager.CreateProvider(ctx, "user1", newBootstrapProvider("account1"), "cluster1.example.com")
|
||||
require.NoError(t, err, "bootstrap with both permissions must succeed")
|
||||
require.NotNil(t, created)
|
||||
|
||||
settings, err := f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1")
|
||||
require.NoError(t, err, "bootstrap must create the settings row")
|
||||
assert.Equal(t, "cluster1.example.com", settings.Cluster, "settings should pin the bootstrap cluster")
|
||||
})
|
||||
|
||||
t.Run("existing settings need no settings permission", func(t *testing.T) {
|
||||
f := newBootstrapFixture(t)
|
||||
require.NoError(t, f.store.SaveAgentNetworkSettings(ctx, &types.Settings{
|
||||
AccountID: "account1",
|
||||
Cluster: "cluster1.example.com",
|
||||
Subdomain: "existing",
|
||||
}), "pre-existing settings row setup must succeed")
|
||||
|
||||
// Only the providers permission may be consulted: gomock fails the
|
||||
// test on any unexpected settings-permission call.
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
|
||||
|
||||
_, err := f.manager.CreateProvider(ctx, "user1", newBootstrapProvider("account1"), "cluster1.example.com")
|
||||
require.NoError(t, err, "create with existing settings must not require the settings permission")
|
||||
})
|
||||
|
||||
t.Run("no bootstrap cluster needs no settings permission", func(t *testing.T) {
|
||||
f := newBootstrapFixture(t)
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
|
||||
|
||||
_, err := f.manager.CreateProvider(ctx, "user1", newBootstrapProvider("account1"), "")
|
||||
require.NoError(t, err, "create without bootstrap must not require the settings permission")
|
||||
})
|
||||
}
|
||||
@@ -10,6 +10,17 @@ import (
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
// syntheticMapping pairs a synthesised proxy mapping with the address of the
|
||||
// proxy that serves it. The cluster is recorded rather than derived from the
|
||||
// mapping's domain: ProxyMapping does not carry it, and the previous derivation
|
||||
// -- everything after the first DNS label -- is wrong whenever the service's
|
||||
// domain is not one label under its proxy's address, which silently addressed
|
||||
// updates to a cluster no proxy declares.
|
||||
type syntheticMapping struct {
|
||||
mapping *proto.ProxyMapping
|
||||
cluster string
|
||||
}
|
||||
|
||||
// reconcile recomputes the synthesised reverse-proxy services for an
|
||||
// account, diffs them against the previously-synthesised set in the
|
||||
// in-memory cache, and emits Create / Update / Delete proxy mappings
|
||||
@@ -45,18 +56,21 @@ func (m *managerImpl) reconcile(ctx context.Context, accountID string) {
|
||||
}
|
||||
|
||||
oidcCfg := m.proxyController.GetOIDCValidationConfig()
|
||||
current := make(map[string]*proto.ProxyMapping, len(services))
|
||||
current := make(map[string]syntheticMapping, len(services))
|
||||
for _, svc := range services {
|
||||
if svc == nil || svc.ID == "" {
|
||||
continue
|
||||
}
|
||||
current[svc.ID] = svc.ToProtoMapping(rpservice.Update, "", oidcCfg)
|
||||
current[svc.ID] = syntheticMapping{
|
||||
mapping: svc.ToProtoMapping(rpservice.Update, "", oidcCfg),
|
||||
cluster: svc.ProxyCluster,
|
||||
}
|
||||
}
|
||||
|
||||
m.reconcileMu.Lock()
|
||||
previous := m.reconcileCache[accountID]
|
||||
if previous == nil {
|
||||
previous = make(map[string]*proto.ProxyMapping)
|
||||
previous = make(map[string]syntheticMapping)
|
||||
}
|
||||
|
||||
creates, updates, deletes := diffMappings(previous, current)
|
||||
@@ -67,34 +81,36 @@ func (m *managerImpl) reconcile(ctx context.Context, accountID string) {
|
||||
}
|
||||
m.reconcileMu.Unlock()
|
||||
|
||||
for _, mapping := range creates {
|
||||
mapping.Type = proto.ProxyMappingUpdateType_UPDATE_TYPE_CREATED
|
||||
m.proxyController.SendServiceUpdateToCluster(ctx, accountID, mapping, clusterFromMapping(mapping))
|
||||
for _, entry := range creates {
|
||||
entry.mapping.Type = proto.ProxyMappingUpdateType_UPDATE_TYPE_CREATED
|
||||
m.proxyController.SendServiceUpdateToCluster(ctx, accountID, entry.mapping, entry.cluster)
|
||||
}
|
||||
for _, mapping := range updates {
|
||||
mapping.Type = proto.ProxyMappingUpdateType_UPDATE_TYPE_MODIFIED
|
||||
m.proxyController.SendServiceUpdateToCluster(ctx, accountID, mapping, clusterFromMapping(mapping))
|
||||
for _, entry := range updates {
|
||||
entry.mapping.Type = proto.ProxyMappingUpdateType_UPDATE_TYPE_MODIFIED
|
||||
m.proxyController.SendServiceUpdateToCluster(ctx, accountID, entry.mapping, entry.cluster)
|
||||
}
|
||||
for _, mapping := range deletes {
|
||||
mapping.Type = proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED
|
||||
m.proxyController.SendServiceUpdateToCluster(ctx, accountID, mapping, clusterFromMapping(mapping))
|
||||
for _, entry := range deletes {
|
||||
entry.mapping.Type = proto.ProxyMappingUpdateType_UPDATE_TYPE_REMOVED
|
||||
m.proxyController.SendServiceUpdateToCluster(ctx, accountID, entry.mapping, entry.cluster)
|
||||
}
|
||||
}
|
||||
|
||||
// diffMappings classifies the previous→current transition for a
|
||||
// single account into Create / Update / Delete sets.
|
||||
// diffMappings classifies the previous→current transition for a single
|
||||
// account into Create / Update / Delete sets.
|
||||
//
|
||||
// Cluster moves (current.cluster != previous.cluster) are surfaced as
|
||||
// a Delete on the old cluster + Create on the new — handled by
|
||||
// emitting both a delete (on previous mapping) and a create (on the
|
||||
// current mapping) for that service ID.
|
||||
func diffMappings(previous, current map[string]*proto.ProxyMapping) (creates, updates, deletes []*proto.ProxyMapping) {
|
||||
// A change of serving proxy for the same service ID is surfaced as a Delete
|
||||
// addressed to the old proxy plus a Create addressed to the new one, so the
|
||||
// mapping actually moves. Comparing the recorded cluster is what makes that
|
||||
// detectable: with a placement-free endpoint the mapping's domain is identical
|
||||
// before and after the move, so nothing about the mapping itself reveals it.
|
||||
func diffMappings(previous, current map[string]syntheticMapping) (creates, updates, deletes []syntheticMapping) {
|
||||
for id, cur := range current {
|
||||
prev, existed := previous[id]
|
||||
switch {
|
||||
case !existed:
|
||||
creates = append(creates, cur)
|
||||
case prev.GetDomain() == "" || cur.GetAccountId() == prev.GetAccountId() && currentClusterChanged(prev, cur):
|
||||
case prev.mapping.GetDomain() == "" ||
|
||||
cur.mapping.GetAccountId() == prev.mapping.GetAccountId() && prev.cluster != cur.cluster:
|
||||
deletes = append(deletes, prev)
|
||||
creates = append(creates, cur)
|
||||
default:
|
||||
@@ -108,24 +124,3 @@ func diffMappings(previous, current map[string]*proto.ProxyMapping) (creates, up
|
||||
}
|
||||
return creates, updates, deletes
|
||||
}
|
||||
|
||||
func currentClusterChanged(prev, cur *proto.ProxyMapping) bool {
|
||||
return clusterFromMapping(prev) != clusterFromMapping(cur)
|
||||
}
|
||||
|
||||
// clusterFromMapping returns the cluster the mapping should be sent
|
||||
// to. ProxyMapping doesn't carry the cluster directly, so we rely on
|
||||
// the synthesised service's domain (`<slug>.<cluster>`) and split on
|
||||
// the first '.'.
|
||||
func clusterFromMapping(m *proto.ProxyMapping) string {
|
||||
if m == nil {
|
||||
return ""
|
||||
}
|
||||
domain := m.GetDomain()
|
||||
for i := 0; i < len(domain); i++ {
|
||||
if domain[i] == '.' {
|
||||
return domain[i+1:]
|
||||
}
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
@@ -21,7 +21,7 @@ func newReconcileMgr(t *testing.T, ctrl *gomock.Controller) (*managerImpl, *stor
|
||||
return &managerImpl{
|
||||
store: mockStore,
|
||||
proxyController: mockProxy,
|
||||
reconcileCache: make(map[string]map[string]*proto.ProxyMapping),
|
||||
reconcileCache: make(map[string]map[string]syntheticMapping),
|
||||
}, mockStore, mockProxy
|
||||
}
|
||||
|
||||
@@ -52,9 +52,9 @@ func newReconcileTestPolicy(providerID, sourceGroupID string) *types.Policy {
|
||||
|
||||
func newReconcileTestSettings() *types.Settings {
|
||||
return &types.Settings{
|
||||
AccountID: "acct-1",
|
||||
Cluster: "eu.proxy.netbird.io",
|
||||
Subdomain: "violet",
|
||||
AccountID: "acct-1",
|
||||
Domain: "violet.eu.proxy.netbird.io",
|
||||
ProxyAddress: "eu.proxy.netbird.io",
|
||||
}
|
||||
}
|
||||
|
||||
@@ -196,7 +196,7 @@ func TestReconcile_PolicyRemoved_EmitsDelete(t *testing.T) {
|
||||
func TestReconcile_NilProxyController_NoOp(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
mgr := &managerImpl{
|
||||
reconcileCache: make(map[string]map[string]*proto.ProxyMapping),
|
||||
reconcileCache: make(map[string]map[string]syntheticMapping),
|
||||
}
|
||||
// Must not panic; must not query the store.
|
||||
mgr.reconcile(ctx, "acct-1")
|
||||
@@ -212,21 +212,78 @@ func TestReconcile_EmptyAccountID_NoOp(t *testing.T) {
|
||||
mgr.reconcile(ctx, "")
|
||||
}
|
||||
|
||||
func TestClusterFromMapping(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
domain string
|
||||
want string
|
||||
}{
|
||||
{"simple", "openai.eu.proxy.netbird.io", "eu.proxy.netbird.io"},
|
||||
{"deeply nested", "a.b.c.d", "b.c.d"},
|
||||
{"no dot", "openai", ""},
|
||||
{"empty", "", ""},
|
||||
// TestDiffMappings_ServingProxyChange — when the proxy serving an account
|
||||
// changes, the same service ID must be deleted on the old proxy and created on
|
||||
// the new one. The cluster cannot be recovered from the mapping's domain: with a
|
||||
// placement-free endpoint the domain does not change at all when the serving
|
||||
// proxy does, so a domain-derived cluster sees no change and emits a plain
|
||||
// update, addressed to a proxy that does not exist.
|
||||
func TestDiffMappings_ServingProxyChange(t *testing.T) {
|
||||
previous := map[string]syntheticMapping{
|
||||
"svc-1": {
|
||||
mapping: &proto.ProxyMapping{Id: "svc-1", AccountId: "acct-1", Domain: "brave-otter.gateway.example.com"},
|
||||
cluster: "proxy.example.com",
|
||||
},
|
||||
}
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
got := clusterFromMapping(&proto.ProxyMapping{Domain: tt.domain})
|
||||
assert.Equal(t, tt.want, got)
|
||||
})
|
||||
current := map[string]syntheticMapping{
|
||||
"svc-1": {
|
||||
mapping: &proto.ProxyMapping{Id: "svc-1", AccountId: "acct-1", Domain: "brave-otter.gateway.example.com"},
|
||||
cluster: "brave-otter.gateway.example.com",
|
||||
},
|
||||
}
|
||||
|
||||
creates, updates, deletes := diffMappings(previous, current)
|
||||
|
||||
if assert.Len(t, deletes, 1, "the old proxy must be told to drop the mapping") {
|
||||
assert.Equal(t, "proxy.example.com", deletes[0].cluster)
|
||||
}
|
||||
if assert.Len(t, creates, 1, "the new proxy must be told to add it") {
|
||||
assert.Equal(t, "brave-otter.gateway.example.com", creates[0].cluster)
|
||||
}
|
||||
assert.Empty(t, updates, "a serving-proxy move is a delete plus a create, not an update")
|
||||
}
|
||||
|
||||
// TestDiffMappings_UnchangedClusterIsAnUpdate keeps the ordinary path: same
|
||||
// service, same proxy, changed contents.
|
||||
func TestDiffMappings_UnchangedClusterIsAnUpdate(t *testing.T) {
|
||||
previous := map[string]syntheticMapping{
|
||||
"svc-1": {
|
||||
mapping: &proto.ProxyMapping{Id: "svc-1", AccountId: "acct-1", Domain: "otter.proxy.example.com"},
|
||||
cluster: "proxy.example.com",
|
||||
},
|
||||
}
|
||||
current := map[string]syntheticMapping{
|
||||
"svc-1": {
|
||||
mapping: &proto.ProxyMapping{Id: "svc-1", AccountId: "acct-1", Domain: "otter.proxy.example.com"},
|
||||
cluster: "proxy.example.com",
|
||||
},
|
||||
}
|
||||
|
||||
creates, updates, deletes := diffMappings(previous, current)
|
||||
|
||||
assert.Empty(t, creates)
|
||||
assert.Empty(t, deletes)
|
||||
if assert.Len(t, updates, 1) {
|
||||
assert.Equal(t, "proxy.example.com", updates[0].cluster)
|
||||
}
|
||||
}
|
||||
|
||||
// TestDiffMappings_RemovedServiceIsDeletedOnItsOwnCluster — a service that has
|
||||
// gone away is deleted on the cluster it was last served by, which is recorded
|
||||
// rather than re-derived.
|
||||
func TestDiffMappings_RemovedServiceIsDeletedOnItsOwnCluster(t *testing.T) {
|
||||
previous := map[string]syntheticMapping{
|
||||
"svc-1": {
|
||||
mapping: &proto.ProxyMapping{Id: "svc-1", AccountId: "acct-1", Domain: "brave-otter.gateway.example.com"},
|
||||
cluster: "brave-otter.gateway.example.com",
|
||||
},
|
||||
}
|
||||
|
||||
creates, updates, deletes := diffMappings(previous, map[string]syntheticMapping{})
|
||||
|
||||
assert.Empty(t, creates)
|
||||
assert.Empty(t, updates)
|
||||
if assert.Len(t, deletes, 1) {
|
||||
assert.Equal(t, "brave-otter.gateway.example.com", deletes[0].cluster)
|
||||
}
|
||||
}
|
||||
|
||||
@@ -0,0 +1,225 @@
|
||||
package agentnetwork
|
||||
|
||||
import (
|
||||
"context"
|
||||
"runtime"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
"github.com/netbirdio/netbird/management/server/account"
|
||||
"github.com/netbirdio/netbird/management/server/permissions"
|
||||
"github.com/netbirdio/netbird/management/server/permissions/modules"
|
||||
"github.com/netbirdio/netbird/management/server/permissions/operations"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
nbtypes "github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
// bootstrapFixture wires a real sqlite store to a gomock permissions manager
|
||||
// so tests can grant or deny the settings permission per case.
|
||||
type bootstrapFixture struct {
|
||||
manager Manager
|
||||
store store.Store
|
||||
perms *permissions.MockManager
|
||||
}
|
||||
|
||||
func newBootstrapFixture(t *testing.T) *bootstrapFixture {
|
||||
t.Helper()
|
||||
if runtime.GOOS == "windows" {
|
||||
t.Skip("sqlite store not properly supported on Windows yet")
|
||||
}
|
||||
t.Setenv("NETBIRD_STORE_ENGINE", string(nbtypes.SqliteStoreEngine))
|
||||
|
||||
st, cleanUp, err := store.NewTestStoreFromSQL(context.Background(), "", t.TempDir())
|
||||
require.NoError(t, err, "test store setup must succeed")
|
||||
t.Cleanup(cleanUp)
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
perms := permissions.NewMockManager(ctrl)
|
||||
|
||||
accounts := account.NewMockManager(ctrl)
|
||||
accounts.EXPECT().StoreEvent(gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes()
|
||||
accounts.EXPECT().UpdateAccountPeers(gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes()
|
||||
accounts.EXPECT().BufferUpdateAccountPeers(gomock.Any(), gomock.Any(), gomock.Any()).AnyTimes()
|
||||
|
||||
return &bootstrapFixture{
|
||||
manager: NewManager(st, perms, accounts, nil),
|
||||
store: st,
|
||||
perms: perms,
|
||||
}
|
||||
}
|
||||
|
||||
func (f *bootstrapFixture) expectPermission(accountID, userID string, module modules.Module, op operations.Operation, allowed bool) {
|
||||
f.perms.EXPECT().
|
||||
ValidateUserPermissions(gomock.Any(), accountID, userID, module, op).
|
||||
Return(allowed, context.Background(), nil)
|
||||
}
|
||||
|
||||
func (f *bootstrapFixture) createSettings(ctx context.Context, accountID, userID, proxyAddress, endpoint string) (*types.Settings, error) {
|
||||
return f.manager.CreateSettings(ctx, userID, types.DefaultSettings(accountID), proxyAddress, endpoint)
|
||||
}
|
||||
|
||||
// TestCreateSettingsRequiresPermission pins the gate: bootstrap assigns the
|
||||
// account's immutable endpoint, a settings write requiring the settings
|
||||
// Create permission — and a denial leaves no row behind.
|
||||
func TestCreateSettingsRequiresPermission(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
f := newBootstrapFixture(t)
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, false)
|
||||
|
||||
_, err := f.createSettings(ctx, "account1", "user1", "cluster1.example.com", "")
|
||||
require.Error(t, err, "bootstrap without the settings permission must fail")
|
||||
var sErr *status.Error
|
||||
require.ErrorAs(t, err, &sErr)
|
||||
assert.Equal(t, status.PermissionDenied, sErr.Type(), "denial should surface as permission denied")
|
||||
|
||||
_, err = f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1")
|
||||
assert.Error(t, err, "settings row must not be created when bootstrap is denied")
|
||||
}
|
||||
|
||||
// TestCreateSettingsLabeled pins the labeled shape: the server allocates an
|
||||
// adjective-noun label beneath the proxy address, the pin is not dedicated,
|
||||
// and the domain records the full endpoint hostname.
|
||||
func TestCreateSettingsLabeled(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
f := newBootstrapFixture(t)
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
|
||||
created, err := f.createSettings(ctx, "account1", "user1", "Cluster1.Example.com", "")
|
||||
require.NoError(t, err, "labeled bootstrap must succeed")
|
||||
assert.Equal(t, "cluster1.example.com", created.ProxyAddress, "proxy address must be pinned lowercased")
|
||||
require.True(t, strings.HasSuffix(created.Domain, ".cluster1.example.com"),
|
||||
"domain must hang one label beneath the proxy address: %s", created.Domain)
|
||||
label := strings.TrimSuffix(created.Domain, ".cluster1.example.com")
|
||||
assert.NotContains(t, label, ".", "the allocated label must be a single DNS label: %s", label)
|
||||
assert.False(t, created.Dedicated(), "a labeled pin is not dedicated")
|
||||
assert.Equal(t, created.Domain, created.Endpoint(), "the endpoint is the domain column")
|
||||
|
||||
stored, err := f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1")
|
||||
require.NoError(t, err, "bootstrap must persist the row")
|
||||
assert.Equal(t, created.Domain, stored.Domain)
|
||||
assert.Equal(t, created.ProxyAddress, stored.ProxyAddress)
|
||||
}
|
||||
|
||||
// TestCreateSettingsSelfAddressed pins the dedicated shape: the endpoint is
|
||||
// claimed verbatim (normalized), Domain == ProxyAddress, and the claim
|
||||
// succeeds with no proxy declaring the address yet (address-first).
|
||||
func TestCreateSettingsSelfAddressed(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
f := newBootstrapFixture(t)
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
|
||||
created, err := f.createSettings(ctx, "account1", "user1", "", "Brave-Otter.GW.Example.com")
|
||||
require.NoError(t, err, "self-addressed bootstrap must succeed")
|
||||
assert.Equal(t, "brave-otter.gw.example.com", created.Domain, "endpoint must be claimed lowercased")
|
||||
assert.Equal(t, created.Domain, created.ProxyAddress, "self-addressed: proxy address is the endpoint")
|
||||
assert.True(t, created.Dedicated(), "a self-addressed pin is dedicated")
|
||||
}
|
||||
|
||||
// TestCreateSettingsIdentityFieldValidation pins the request contract: exactly
|
||||
// one of proxyAddress and endpoint, and both must be well-formed hostnames.
|
||||
func TestCreateSettingsIdentityFieldValidation(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
|
||||
cases := map[string]struct {
|
||||
proxyAddress string
|
||||
endpoint string
|
||||
}{
|
||||
"neither": {"", ""},
|
||||
"both": {"cluster1.example.com", "gw.example.com"},
|
||||
"trailing dot endpoint": {"", "gw.example.com."},
|
||||
"leading dot endpoint": {"", ".gw.example.com"},
|
||||
"whitespace inside": {"", "g w.example.com"},
|
||||
"empty label in parent": {"eu..example.com", ""},
|
||||
"hyphen-edged label": {"", "-gw.example.com"},
|
||||
}
|
||||
for name, tc := range cases {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
f := newBootstrapFixture(t)
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
|
||||
_, err := f.createSettings(ctx, "account1", "user1", tc.proxyAddress, tc.endpoint)
|
||||
require.Error(t, err, "invalid identity input must be rejected")
|
||||
var sErr *status.Error
|
||||
require.ErrorAs(t, err, &sErr)
|
||||
assert.Equal(t, status.InvalidArgument, sErr.Type(), "rejection must be a validation error")
|
||||
|
||||
_, err = f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1")
|
||||
assert.Error(t, err, "no row may be left behind by a rejected bootstrap")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestCreateSettingsConflictsOnSecondBootstrap pins that bootstrap is a
|
||||
// one-time create per account: a second call is a conflict, whatever shape it
|
||||
// asks for, and the original row survives untouched.
|
||||
func TestCreateSettingsConflictsOnSecondBootstrap(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
f := newBootstrapFixture(t)
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
|
||||
first, err := f.createSettings(ctx, "account1", "user1", "cluster1.example.com", "")
|
||||
require.NoError(t, err)
|
||||
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
_, err = f.createSettings(ctx, "account1", "user1", "", "other.example.com")
|
||||
require.Error(t, err, "second bootstrap must fail")
|
||||
var sErr *status.Error
|
||||
require.ErrorAs(t, err, &sErr)
|
||||
assert.Equal(t, status.AlreadyExists, sErr.Type(), "second bootstrap must surface as a conflict")
|
||||
|
||||
stored, err := f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, first.Domain, stored.Domain, "the original endpoint must survive the rejected bootstrap")
|
||||
}
|
||||
|
||||
// TestCreateSettingsEndpointTaken pins global hostname uniqueness: a hostname
|
||||
// held by one account cannot be claimed by another, in either direction —
|
||||
// self-addressed onto self-addressed, or self-addressed onto an allocated
|
||||
// labeled endpoint.
|
||||
func TestCreateSettingsEndpointTaken(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
f := newBootstrapFixture(t)
|
||||
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkSettings, operations.Create, true)
|
||||
first, err := f.createSettings(ctx, "account1", "user1", "", "gw.example.com")
|
||||
require.NoError(t, err)
|
||||
|
||||
f.expectPermission("account2", "user2", modules.AgentNetworkSettings, operations.Create, true)
|
||||
_, err = f.createSettings(ctx, "account2", "user2", "", "gw.example.com")
|
||||
require.Error(t, err, "a taken hostname must be refused")
|
||||
var sErr *status.Error
|
||||
require.ErrorAs(t, err, &sErr)
|
||||
assert.Equal(t, status.AlreadyExists, sErr.Type(), "the refusal must surface as a conflict")
|
||||
|
||||
f.expectPermission("account3", "user3", modules.AgentNetworkSettings, operations.Create, true)
|
||||
_, err = f.createSettings(ctx, "account3", "user3", "", first.Domain)
|
||||
require.Error(t, err, "claiming another account's endpoint must be refused")
|
||||
}
|
||||
|
||||
// TestCreateProviderHasNoSettingsSideEffects pins the decoupling: provider
|
||||
// create needs only the providers permission (gomock fails the test on any
|
||||
// settings-permission call) and never creates a settings row.
|
||||
func TestCreateProviderHasNoSettingsSideEffects(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
f := newBootstrapFixture(t)
|
||||
f.expectPermission("account1", "user1", modules.AgentNetworkProviders, operations.Create, true)
|
||||
|
||||
provider := types.NewProvider("account1")
|
||||
provider.Name = "openai"
|
||||
provider.UpstreamURL = "https://api.openai.com"
|
||||
provider.APIKey = "sk-test"
|
||||
provider.Enabled = true
|
||||
|
||||
created, err := f.manager.CreateProvider(ctx, "user1", provider)
|
||||
require.NoError(t, err, "provider create must succeed on the providers permission alone")
|
||||
require.NotNil(t, created)
|
||||
|
||||
_, err = f.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, "account1")
|
||||
assert.Error(t, err, "provider create must not conjure a settings row")
|
||||
}
|
||||
@@ -89,7 +89,7 @@ func SynthesizeServicesForCluster(ctx context.Context, s store.Store, clusterAdd
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
settingsRows, err := s.GetAgentNetworkSettingsByCluster(ctx, store.LockingStrengthNone, clusterAddr)
|
||||
settingsRows, err := s.GetAgentNetworkSettingsByProxyAddress(ctx, store.LockingStrengthNone, clusterAddr)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list agent network settings on cluster: %w", err)
|
||||
}
|
||||
@@ -116,53 +116,41 @@ func SynthesizeServicesForCluster(ctx context.Context, s store.Store, clusterAdd
|
||||
}
|
||||
|
||||
// SynthesizeServiceForDomain resolves a single agent-network service by its
|
||||
// public endpoint domain. It lists the (few) settings rows on the domain's
|
||||
// cluster, matches the one whose endpoint equals the domain, and synthesises
|
||||
// only that account — avoiding full per-account synthesis for every tenant on
|
||||
// the cluster, which is what auth/session paths previously paid. Returns nil
|
||||
// (no error) when no account owns the domain.
|
||||
// public endpoint domain — a point query on the settings domain unique index,
|
||||
// then synthesis of just that account. Returns nil (no error) when no account
|
||||
// owns the domain.
|
||||
func SynthesizeServiceForDomain(ctx context.Context, s store.Store, domain string) (*rpservice.Service, error) {
|
||||
domain = strings.TrimSpace(domain)
|
||||
cluster := clusterFromDomain(domain)
|
||||
if domain != "" && cluster != "" {
|
||||
settingsRows, err := s.GetAgentNetworkSettingsByCluster(ctx, store.LockingStrengthNone, cluster)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list agent network settings on cluster: %w", err)
|
||||
domain = strings.ToLower(strings.TrimSpace(domain))
|
||||
if domain == "" {
|
||||
return nil, nil //nolint:nilnil // optional lookup: no account owns the domain
|
||||
}
|
||||
|
||||
settings, err := s.GetAgentNetworkSettingsByDomain(ctx, store.LockingStrengthNone, domain)
|
||||
if err != nil {
|
||||
if isNotFound(err) {
|
||||
return nil, nil //nolint:nilnil // optional lookup: no account owns the domain
|
||||
}
|
||||
for _, settings := range settingsRows {
|
||||
if settings == nil || settings.Endpoint() != domain {
|
||||
continue
|
||||
}
|
||||
services, serr := SynthesizeServices(ctx, s, settings.AccountID)
|
||||
if serr != nil {
|
||||
return nil, serr
|
||||
}
|
||||
for _, svc := range services {
|
||||
if svc != nil && svc.Domain == domain {
|
||||
return svc, nil
|
||||
}
|
||||
}
|
||||
break
|
||||
return nil, fmt.Errorf("get agent network settings by domain: %w", err)
|
||||
}
|
||||
|
||||
services, err := SynthesizeServices(ctx, s, settings.AccountID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
for _, svc := range services {
|
||||
if svc != nil && svc.Domain == domain {
|
||||
return svc, nil
|
||||
}
|
||||
}
|
||||
return nil, nil //nolint:nilnil // optional lookup: no account owns the domain
|
||||
}
|
||||
|
||||
// clusterFromDomain returns the cluster portion of an endpoint domain (every
|
||||
// label after the first).
|
||||
func clusterFromDomain(domain string) string {
|
||||
if i := strings.IndexByte(domain, '.'); i >= 0 {
|
||||
return domain[i+1:]
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
// SynthesizeServices builds the in-memory reverse-proxy service that
|
||||
// fronts the account's agent-network gateway. Returns nil when the
|
||||
// account has no settings row, no enabled providers, or no enabled
|
||||
// policies — in any of those cases there's nothing useful to expose.
|
||||
//
|
||||
// One service per (account, settings.Cluster) is emitted. The router
|
||||
// One service per (account, settings.ProxyAddress) is emitted. The router
|
||||
// middleware encodes a denormalised model→provider routing table
|
||||
// (auth headers + decrypted API keys baked in); the policy_check
|
||||
// middleware encodes per-provider authorised group IDs derived from
|
||||
@@ -175,7 +163,7 @@ func SynthesizeServices(ctx context.Context, s store.Store, accountID string) ([
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if !ok || strings.TrimSpace(settings.Cluster) == "" {
|
||||
if !ok || strings.TrimSpace(settings.ProxyAddress) == "" {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
@@ -934,7 +922,7 @@ func buildAccountService(
|
||||
middlewares []rpservice.MiddlewareConfig,
|
||||
sessionPriv, sessionPub string,
|
||||
) *rpservice.Service {
|
||||
cluster := settings.Cluster
|
||||
cluster := settings.ProxyAddress
|
||||
domain := settings.Endpoint()
|
||||
serviceID := SynthesizedServiceIDPrefix + accountID
|
||||
|
||||
|
||||
@@ -147,7 +147,7 @@ func TestReconcile_RealStore_PushesPrivateAfterStatusToggle(t *testing.T) {
|
||||
store: s,
|
||||
accountManager: noopAccountManager{},
|
||||
proxyController: ctrl,
|
||||
reconcileCache: make(map[string]map[string]*proto.ProxyMapping),
|
||||
reconcileCache: make(map[string]map[string]syntheticMapping),
|
||||
}
|
||||
|
||||
m.reconcile(ctx, testAccountID) // initial, provider enabled
|
||||
|
||||
@@ -19,15 +19,14 @@ import (
|
||||
const (
|
||||
testAccountID = "acct-1"
|
||||
testCluster = "eu.proxy.netbird.io"
|
||||
testSubdomain = "violet"
|
||||
testEndpoint = "violet.eu.proxy.netbird.io"
|
||||
)
|
||||
|
||||
func newSynthTestSettings() *types.Settings {
|
||||
return &types.Settings{
|
||||
AccountID: testAccountID,
|
||||
Cluster: testCluster,
|
||||
Subdomain: testSubdomain,
|
||||
AccountID: testAccountID,
|
||||
Domain: testEndpoint,
|
||||
ProxyAddress: testCluster,
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package types
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -12,13 +13,23 @@ import (
|
||||
// the long-term aggregate and are retained independently.
|
||||
const DefaultAccessLogRetentionDays = 30
|
||||
|
||||
// Settings is the per-account agent-network configuration row. One
|
||||
// row per account. Cluster + Subdomain are immutable once written and
|
||||
// produce the public endpoint agents call (`<subdomain>.<cluster>`).
|
||||
// Settings is the per-account agent-network configuration row. One row per
|
||||
// account. Domain and ProxyAddress are assigned at bootstrap and immutable
|
||||
// thereafter; a persisted row is always fully allocated — there is no "row
|
||||
// exists, endpoint pending" state.
|
||||
type Settings struct {
|
||||
AccountID string `gorm:"primaryKey"`
|
||||
Cluster string
|
||||
Subdomain string `gorm:"index:idx_agent_network_settings_cluster_subdomain"`
|
||||
|
||||
// Domain is the gateway endpoint hostname agents call. Globally unique
|
||||
// across accounts. Sized explicitly because MySQL cannot index an
|
||||
// unbounded TEXT column; 255 covers the RFC 1035 253-octet bound.
|
||||
Domain string `gorm:"type:varchar(255);uniqueIndex:idx_agent_network_settings_domain"`
|
||||
|
||||
// ProxyAddress is the declared cluster address of the proxy serving this
|
||||
// account's gateway. Either equal to Domain — a proxy dedicated to this
|
||||
// account, declaring the tenant's own hostname — or Domain's immediate
|
||||
// parent, with the endpoint one label beneath it on a shared cluster.
|
||||
ProxyAddress string `gorm:"type:varchar(255);index:idx_agent_network_settings_proxy_address"`
|
||||
|
||||
// Account-level collection controls sourced by the synthesizer.
|
||||
// EnableLogCollection gates the per-request access-log trail and defaults
|
||||
@@ -45,9 +56,9 @@ func (Settings) TableName() string { return "agent_network_settings" }
|
||||
|
||||
// DefaultSettings returns the settings an account observes before its row is
|
||||
// bootstrapped: log collection on with the default retention, everything else
|
||||
// off, and no cluster/subdomain assigned yet. Bootstrap persists exactly these
|
||||
// values plus the assigned cluster and subdomain, so the pre-bootstrap read
|
||||
// and the freshly bootstrapped row agree.
|
||||
// off, and no domain or proxy address assigned yet. Bootstrap persists exactly
|
||||
// these values plus the assigned domain and proxy address, so the
|
||||
// pre-bootstrap read and the freshly bootstrapped row agree.
|
||||
func DefaultSettings(accountID string) *Settings {
|
||||
return &Settings{
|
||||
AccountID: accountID,
|
||||
@@ -56,14 +67,15 @@ func DefaultSettings(accountID string) *Settings {
|
||||
}
|
||||
}
|
||||
|
||||
// Endpoint returns the bare hostname agents reach this account at:
|
||||
// `<subdomain>.<cluster>`. Empty until both halves are assigned at bootstrap.
|
||||
func (s *Settings) Endpoint() string {
|
||||
if s.Cluster == "" || s.Subdomain == "" {
|
||||
return ""
|
||||
}
|
||||
return s.Subdomain + "." + s.Cluster
|
||||
}
|
||||
// Endpoint returns the bare hostname agents reach this account at — the
|
||||
// Domain column. Empty until the row is bootstrapped.
|
||||
func (s *Settings) Endpoint() string { return s.Domain }
|
||||
|
||||
// Dedicated reports whether the account's gateway is served by a proxy
|
||||
// dedicated to it — the self-addressed shape, where the serving proxy declares
|
||||
// the endpoint hostname itself. The alternative (labeled) shape has the
|
||||
// endpoint one label beneath a shared cluster's address.
|
||||
func (s *Settings) Dedicated() bool { return s.Domain != "" && s.Domain == s.ProxyAddress }
|
||||
|
||||
// ToAPIResponse renders the settings as the API representation. The
|
||||
// timestamps are omitted while zero — a default (not yet bootstrapped) view
|
||||
@@ -71,9 +83,9 @@ func (s *Settings) Endpoint() string {
|
||||
func (s *Settings) ToAPIResponse() *api.AgentNetworkSettings {
|
||||
retention := s.AccessLogRetentionDays
|
||||
resp := &api.AgentNetworkSettings{
|
||||
Cluster: s.Cluster,
|
||||
Subdomain: s.Subdomain,
|
||||
Endpoint: s.Endpoint(),
|
||||
ProxyAddress: s.ProxyAddress,
|
||||
Dedicated: s.Dedicated(),
|
||||
EnableLogCollection: s.EnableLogCollection,
|
||||
EnablePromptCollection: s.EnablePromptCollection,
|
||||
RedactPii: s.RedactPii,
|
||||
@@ -90,19 +102,91 @@ func (s *Settings) ToAPIResponse() *api.AgentNetworkSettings {
|
||||
return resp
|
||||
}
|
||||
|
||||
// FromAPIRequest applies the request onto the receiver. The mutable
|
||||
// collection fields are always replaced with the request values. Cluster
|
||||
// participates only in bootstrap and the immutability check (see
|
||||
// Manager.UpdateSettings); Subdomain is server-assigned and never taken
|
||||
// from a request.
|
||||
// FromAPIRequest applies the update request onto the receiver: every mutable
|
||||
// field is replaced with the request value, and the identity fields (Domain,
|
||||
// ProxyAddress) carry the request's echo of the assigned values. The identity
|
||||
// fields are never written to the stored row — UpdateSettings compares them
|
||||
// against it and rejects the request when they differ, so PUT keeps the
|
||||
// house convention of requiring every field while the endpoint and proxy
|
||||
// address stay immutable.
|
||||
//
|
||||
// Every field is required by the schema, so none is presence-sensitive.
|
||||
// AccessLogRetentionDays in particular must stay required: the caller receives
|
||||
// a zero-valued Settings, and UpdateSettings copies each field onto the stored
|
||||
// row unconditionally, so an omitted value would be written as 0 — which the
|
||||
// API documents as "keep indefinitely". Making retention optional would
|
||||
// therefore let a client silently maximise log retention by leaving it out.
|
||||
func (s *Settings) FromAPIRequest(req *api.AgentNetworkSettingsRequest) {
|
||||
if req.Cluster != nil {
|
||||
s.Cluster = strings.TrimSpace(*req.Cluster)
|
||||
}
|
||||
s.Domain = req.Endpoint
|
||||
s.ProxyAddress = req.ProxyAddress
|
||||
s.EnableLogCollection = req.EnableLogCollection
|
||||
s.EnablePromptCollection = req.EnablePromptCollection
|
||||
s.RedactPii = req.RedactPii
|
||||
s.AccessLogRetentionDays = req.AccessLogRetentionDays
|
||||
}
|
||||
|
||||
// FromAPICreateRequest applies the optional collection toggles of a bootstrap
|
||||
// request onto the receiver (typically DefaultSettings), leaving defaults in
|
||||
// place for omitted fields. The identity fields are resolved by the manager
|
||||
// from the request's proxy_address / endpoint, not copied here.
|
||||
func (s *Settings) FromAPICreateRequest(req *api.AgentNetworkSettingsCreateRequest) {
|
||||
if req.EnableLogCollection != nil {
|
||||
s.EnableLogCollection = *req.EnableLogCollection
|
||||
}
|
||||
if req.EnablePromptCollection != nil {
|
||||
s.EnablePromptCollection = *req.EnablePromptCollection
|
||||
}
|
||||
if req.RedactPii != nil {
|
||||
s.RedactPii = *req.RedactPii
|
||||
}
|
||||
if req.AccessLogRetentionDays != nil {
|
||||
s.AccessLogRetentionDays = *req.AccessLogRetentionDays
|
||||
}
|
||||
}
|
||||
|
||||
// maxHostnameLength is the RFC 1035 bound on a full domain name.
|
||||
const maxHostnameLength = 253
|
||||
|
||||
// NormalizeHostname lowercases and trims a caller-supplied hostname and
|
||||
// validates its shape: non-empty DNS labels of letters, digits and inner
|
||||
// hyphens, joined by single dots, within length bounds. Shapes that
|
||||
// canonicalization cannot repair — leading/trailing dots, empty labels,
|
||||
// whitespace inside the name — are rejected rather than guessed at, because
|
||||
// the value lands in an immutable column.
|
||||
func NormalizeHostname(raw string) (string, error) {
|
||||
hostname := strings.ToLower(strings.TrimSpace(raw))
|
||||
if hostname == "" {
|
||||
return "", fmt.Errorf("hostname is empty")
|
||||
}
|
||||
if len(hostname) > maxHostnameLength {
|
||||
return "", fmt.Errorf("hostname exceeds %d characters", maxHostnameLength)
|
||||
}
|
||||
for _, label := range strings.Split(hostname, ".") {
|
||||
if err := validateHostnameLabel(label); err != nil {
|
||||
return "", fmt.Errorf("invalid hostname %q: %w", hostname, err)
|
||||
}
|
||||
}
|
||||
return hostname, nil
|
||||
}
|
||||
|
||||
func validateHostnameLabel(label string) error {
|
||||
if label == "" {
|
||||
return fmt.Errorf("empty label (leading, trailing or doubled dot)")
|
||||
}
|
||||
if len(label) > 63 {
|
||||
return fmt.Errorf("label %q exceeds 63 characters", label)
|
||||
}
|
||||
if label[0] == '-' || label[len(label)-1] == '-' {
|
||||
return fmt.Errorf("label %q must not start or end with a hyphen", label)
|
||||
}
|
||||
for _, r := range label {
|
||||
switch {
|
||||
case r >= 'a' && r <= 'z':
|
||||
case r >= '0' && r <= '9':
|
||||
case r == '-':
|
||||
default:
|
||||
return fmt.Errorf("label %q contains invalid character %q", label, r)
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -0,0 +1,25 @@
|
||||
// Package activity records that a principal used a reverse proxy service, so
|
||||
// that activity accounting counts people and devices which reach services
|
||||
// through the proxy but never touch the dashboard or the management API.
|
||||
package activity
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/netbirdio/netbird/management/server/peer"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
)
|
||||
|
||||
// Manager records reverse proxy usage against the timestamps activity
|
||||
// accounting reads. Both methods are best effort from the caller's point of
|
||||
// view: a lost record is corrected by the next request, and no authorization
|
||||
// decision reads them back.
|
||||
type Manager interface {
|
||||
// RecordUserLogin records a completed SSO sign-in to a proxied service.
|
||||
// Service users have no interactive login and are ignored.
|
||||
RecordUserLogin(ctx context.Context, accountID string, user *types.User) error
|
||||
// RecordPeerSeen records that a peer reached a private service over the
|
||||
// mesh, which is what lets its owner count as active. Peers activity
|
||||
// accounting excludes, and peers already seen recently, are ignored.
|
||||
RecordPeerSeen(ctx context.Context, accountID string, peer *peer.Peer) error
|
||||
}
|
||||
@@ -0,0 +1,66 @@
|
||||
package manager
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/activity"
|
||||
"github.com/netbirdio/netbird/management/server/peer"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
)
|
||||
|
||||
// peerSeenInterval is how stale a peer's LastSeen must be before reaching a
|
||||
// private service refreshes it. Positive tunnel validations are cached on the
|
||||
// proxy for five minutes, so without a floor a busy peer would rewrite its row
|
||||
// behind every request; an hour still sits well inside the window activity
|
||||
// accounting asks about.
|
||||
const peerSeenInterval = time.Hour
|
||||
|
||||
type managerImpl struct {
|
||||
store store.Store
|
||||
}
|
||||
|
||||
// NewManager returns the activity manager backed by the management store.
|
||||
func NewManager(store store.Store) activity.Manager {
|
||||
return &managerImpl{store: store}
|
||||
}
|
||||
|
||||
// RecordUserLogin stamps the login the same way the dashboard and device login
|
||||
// paths do, so a person who only ever reaches proxied services still has a
|
||||
// login on record.
|
||||
func (m *managerImpl) RecordUserLogin(ctx context.Context, accountID string, user *types.User) error {
|
||||
if user == nil || user.IsServiceUser {
|
||||
return nil
|
||||
}
|
||||
|
||||
return m.store.SaveUserLastLogin(ctx, accountID, user.Id, time.Now().UTC())
|
||||
}
|
||||
|
||||
// RecordPeerSeen stamps LastSeen, the column a peer activates its owner
|
||||
// through. The peer the caller already holds answers the throttle without a
|
||||
// query, so a peer seen inside the interval costs nothing to skip; the same
|
||||
// cutoff goes to the store, which enforces it inside the UPDATE so concurrent
|
||||
// requests for one peer cannot each write off their own stale read.
|
||||
func (m *managerImpl) RecordPeerSeen(ctx context.Context, accountID string, peer *peer.Peer) error {
|
||||
if peer == nil || !countsTowardActivity(peer) {
|
||||
return nil
|
||||
}
|
||||
|
||||
staleBefore := time.Now().UTC().Add(-peerSeenInterval)
|
||||
if peer.Status != nil && peer.Status.LastSeen.After(staleBefore) {
|
||||
return nil
|
||||
}
|
||||
|
||||
_, err := m.store.RefreshPeerLastSeen(ctx, accountID, peer.ID, staleBefore)
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
// countsTowardActivity reports whether the peer represents a device a person
|
||||
// actually runs. Embedded proxy peers are infrastructure and browser (WASM)
|
||||
// clients are ephemeral sessions, so activity accounting ignores both and a
|
||||
// write for them could never count.
|
||||
func countsTowardActivity(peer *peer.Peer) bool {
|
||||
return !peer.ProxyMeta.Embedded && peer.Meta.KernelVersion != "wasm"
|
||||
}
|
||||
@@ -0,0 +1,149 @@
|
||||
package manager
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/server/peer"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
)
|
||||
|
||||
// recordingStore captures the two writes the activity manager makes. The
|
||||
// embedded interface satisfies the rest and panics if anything else is called,
|
||||
// which keeps the manager honest about its surface.
|
||||
type recordingStore struct {
|
||||
store.Store
|
||||
logins []loginWrite
|
||||
seen []seenWrite
|
||||
}
|
||||
|
||||
type loginWrite struct {
|
||||
accountID string
|
||||
userID string
|
||||
at time.Time
|
||||
}
|
||||
|
||||
type seenWrite struct {
|
||||
accountID string
|
||||
peerID string
|
||||
staleBefore time.Time
|
||||
}
|
||||
|
||||
func (s *recordingStore) SaveUserLastLogin(_ context.Context, accountID, userID string, lastLogin time.Time) error {
|
||||
s.logins = append(s.logins, loginWrite{accountID: accountID, userID: userID, at: lastLogin})
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *recordingStore) RefreshPeerLastSeen(_ context.Context, accountID, peerID string, staleBefore time.Time) (bool, error) {
|
||||
s.seen = append(s.seen, seenWrite{accountID: accountID, peerID: peerID, staleBefore: staleBefore})
|
||||
return true, nil
|
||||
}
|
||||
|
||||
func TestRecordUserLogin(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
user *types.User
|
||||
expectWrite bool
|
||||
}{
|
||||
{
|
||||
name: "regular user is recorded",
|
||||
user: &types.User{Id: "user1", AccountID: "account1"},
|
||||
expectWrite: true,
|
||||
},
|
||||
{
|
||||
// Activity accounting never counts service users, so a row for one
|
||||
// would be noise.
|
||||
name: "service user is ignored",
|
||||
user: &types.User{Id: "svc1", AccountID: "account1", IsServiceUser: true},
|
||||
expectWrite: false,
|
||||
},
|
||||
{
|
||||
name: "missing user is ignored",
|
||||
user: nil,
|
||||
expectWrite: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
st := &recordingStore{}
|
||||
require.NoError(t, NewManager(st).RecordUserLogin(context.Background(), "account1", tt.user))
|
||||
|
||||
if !tt.expectWrite {
|
||||
assert.Empty(t, st.logins, "no login should have been recorded")
|
||||
return
|
||||
}
|
||||
|
||||
require.Len(t, st.logins, 1, "exactly one login should have been recorded")
|
||||
assert.Equal(t, "account1", st.logins[0].accountID, "login must be recorded against the service account")
|
||||
assert.Equal(t, tt.user.Id, st.logins[0].userID, "login must be recorded against the signing-in user")
|
||||
assert.Equal(t, time.UTC, st.logins[0].at.Location(), "timestamps are written in UTC")
|
||||
assert.WithinDuration(t, time.Now().UTC(), st.logins[0].at, time.Minute, "login should be stamped now")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestRecordPeerSeen(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
peer *peer.Peer
|
||||
expectWrite bool
|
||||
}{
|
||||
{
|
||||
name: "peer seen long ago is recorded",
|
||||
peer: &peer.Peer{ID: "peer1", Status: &peer.PeerStatus{LastSeen: time.Now().Add(-3 * time.Hour)}},
|
||||
expectWrite: true,
|
||||
},
|
||||
{
|
||||
name: "peer never seen is recorded",
|
||||
peer: &peer.Peer{ID: "peer1", Status: &peer.PeerStatus{}},
|
||||
expectWrite: true,
|
||||
},
|
||||
{
|
||||
// The throttle. The caller already holds the peer, so skipping a
|
||||
// recently seen one costs nothing.
|
||||
name: "peer seen inside the interval is skipped",
|
||||
peer: &peer.Peer{ID: "peer1", Status: &peer.PeerStatus{LastSeen: time.Now().Add(-10 * time.Minute)}},
|
||||
expectWrite: false,
|
||||
},
|
||||
{
|
||||
name: "embedded proxy peer is skipped",
|
||||
peer: &peer.Peer{ID: "peer1", ProxyMeta: peer.ProxyMeta{Embedded: true}, Status: &peer.PeerStatus{LastSeen: time.Now().Add(-3 * time.Hour)}},
|
||||
expectWrite: false,
|
||||
},
|
||||
{
|
||||
name: "browser client is skipped",
|
||||
peer: &peer.Peer{ID: "peer1", Meta: peer.PeerSystemMeta{KernelVersion: "wasm"}, Status: &peer.PeerStatus{LastSeen: time.Now().Add(-3 * time.Hour)}},
|
||||
expectWrite: false,
|
||||
},
|
||||
{
|
||||
name: "missing peer is ignored",
|
||||
peer: nil,
|
||||
expectWrite: false,
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
st := &recordingStore{}
|
||||
require.NoError(t, NewManager(st).RecordPeerSeen(context.Background(), "account1", tt.peer))
|
||||
|
||||
if !tt.expectWrite {
|
||||
assert.Empty(t, st.seen, "no activity should have been recorded")
|
||||
return
|
||||
}
|
||||
|
||||
require.Len(t, st.seen, 1, "exactly one activity write should have been recorded")
|
||||
assert.Equal(t, "account1", st.seen[0].accountID, "activity must be recorded against the service account")
|
||||
assert.Equal(t, tt.peer.ID, st.seen[0].peerID, "activity must be recorded against the calling peer")
|
||||
assert.Equal(t, time.UTC, st.seen[0].staleBefore.Location(), "cutoffs are passed in UTC")
|
||||
assert.WithinDuration(t, time.Now().UTC().Add(-peerSeenInterval), st.seen[0].staleBefore, time.Minute,
|
||||
"the store must enforce the same interval the local check applies")
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -2,24 +2,28 @@ package manager
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"net"
|
||||
"strings"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
agentnetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain"
|
||||
"github.com/netbirdio/netbird/management/server/account"
|
||||
"github.com/netbirdio/netbird/management/server/activity"
|
||||
"github.com/netbirdio/netbird/management/server/permissions"
|
||||
"github.com/netbirdio/netbird/management/server/permissions/modules"
|
||||
"github.com/netbirdio/netbird/management/server/permissions/operations"
|
||||
nbstore "github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
type store interface {
|
||||
GetAccount(ctx context.Context, accountID string) (*types.Account, error)
|
||||
GetAgentNetworkSettings(ctx context.Context, lockStrength nbstore.LockingStrength, accountID string) (*agentnetworkTypes.Settings, error)
|
||||
|
||||
GetCustomDomain(ctx context.Context, accountID string, domainID string) (*domain.Domain, error)
|
||||
ListFreeDomains(ctx context.Context, accountID string) ([]string, error)
|
||||
@@ -311,17 +315,21 @@ func (m Manager) getClusterAllowList(ctx context.Context, accountID string) ([]s
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get public cluster addresses: %w", err)
|
||||
}
|
||||
reserved, err := m.reservedGatewayAddress(ctx, accountID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
seen := make(map[string]struct{}, len(byopAddresses)+len(publicAddresses))
|
||||
merged := make([]string, 0, len(byopAddresses)+len(publicAddresses))
|
||||
for _, addr := range byopAddresses {
|
||||
if _, ok := seen[addr]; ok {
|
||||
if _, ok := seen[addr]; ok || addr == reserved {
|
||||
continue
|
||||
}
|
||||
seen[addr] = struct{}{}
|
||||
merged = append(merged, addr)
|
||||
}
|
||||
for _, addr := range publicAddresses {
|
||||
if _, ok := seen[addr]; ok {
|
||||
if _, ok := seen[addr]; ok || addr == reserved {
|
||||
continue
|
||||
}
|
||||
seen[addr] = struct{}{}
|
||||
@@ -330,6 +338,31 @@ func (m Manager) getClusterAllowList(ctx context.Context, accountID string) ([]s
|
||||
return merged, nil
|
||||
}
|
||||
|
||||
// reservedGatewayAddress returns the account's agent-network gateway address
|
||||
// when its settings pin is self-addressed — a proxy dedicated to serving
|
||||
// exactly the gateway. Dropping that address from the cluster allow list keeps
|
||||
// it from being offered as a cluster for ordinary services, and because the
|
||||
// free-domain suffix match is depth-independent, dropping the address rejects
|
||||
// every name beneath it as well as the bare one. Only the account's own
|
||||
// gateway address can ever appear in its allow list (another tenant's gateway
|
||||
// proxy is account-scoped to them), so this single-address exclusion is
|
||||
// sufficient. Returns "" when the account has no settings row or a labeled
|
||||
// (shared-cluster) pin.
|
||||
func (m Manager) reservedGatewayAddress(ctx context.Context, accountID string) (string, error) {
|
||||
settings, err := m.store.GetAgentNetworkSettings(ctx, nbstore.LockingStrengthNone, accountID)
|
||||
if err != nil {
|
||||
var sErr *status.Error
|
||||
if errors.As(err, &sErr) && sErr.Type() == status.NotFound {
|
||||
return "", nil
|
||||
}
|
||||
return "", fmt.Errorf("get agent network settings: %w", err)
|
||||
}
|
||||
if settings == nil || !settings.Dedicated() {
|
||||
return "", nil
|
||||
}
|
||||
return settings.ProxyAddress, nil
|
||||
}
|
||||
|
||||
func extractClusterFromCustomDomains(serviceDomain string, customDomains []*domain.Domain) (string, bool) {
|
||||
bestCluster := ""
|
||||
bestLen := -1
|
||||
|
||||
@@ -7,6 +7,12 @@ import (
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
agentnetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/domain"
|
||||
nbstore "github.com/netbirdio/netbird/management/server/store"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
type mockProxyManager struct {
|
||||
@@ -55,7 +61,7 @@ func TestGetClusterAllowList_BYOPMergedWithPublic(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
mgr := Manager{proxyManager: pm}
|
||||
mgr := Manager{store: &stubStore{}, proxyManager: pm}
|
||||
result, err := mgr.getClusterAllowList(context.Background(), "acc-123")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []string{"byop.example.com", "eu.proxy.netbird.io"}, result)
|
||||
@@ -71,7 +77,7 @@ func TestGetClusterAllowList_DeduplicatesBYOPAndPublic(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
mgr := Manager{proxyManager: pm}
|
||||
mgr := Manager{store: &stubStore{}, proxyManager: pm}
|
||||
result, err := mgr.getClusterAllowList(context.Background(), "acc-123")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []string{"shared.example.com", "byop.example.com", "eu.proxy.netbird.io"}, result)
|
||||
@@ -87,7 +93,7 @@ func TestGetClusterAllowList_NoBYOP_FallbackToShared(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
mgr := Manager{proxyManager: pm}
|
||||
mgr := Manager{store: &stubStore{}, proxyManager: pm}
|
||||
result, err := mgr.getClusterAllowList(context.Background(), "acc-123")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []string{"eu.proxy.netbird.io", "us.proxy.netbird.io"}, result)
|
||||
@@ -100,7 +106,7 @@ func TestGetClusterAllowList_BYOPError_ReturnsError(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
mgr := Manager{proxyManager: pm}
|
||||
mgr := Manager{store: &stubStore{}, proxyManager: pm}
|
||||
result, err := mgr.getClusterAllowList(context.Background(), "acc-123")
|
||||
require.Error(t, err)
|
||||
assert.Nil(t, result)
|
||||
@@ -117,7 +123,7 @@ func TestGetClusterAllowList_PublicError_ReturnsError(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
mgr := Manager{proxyManager: pm}
|
||||
mgr := Manager{store: &stubStore{}, proxyManager: pm}
|
||||
result, err := mgr.getClusterAllowList(context.Background(), "acc-123")
|
||||
require.Error(t, err)
|
||||
assert.Nil(t, result)
|
||||
@@ -134,7 +140,7 @@ func TestGetClusterAllowList_BYOPEmptySlice_FallbackToShared(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
mgr := Manager{proxyManager: pm}
|
||||
mgr := Manager{store: &stubStore{}, proxyManager: pm}
|
||||
result, err := mgr.getClusterAllowList(context.Background(), "acc-123")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []string{"eu.proxy.netbird.io"}, result)
|
||||
@@ -150,8 +156,138 @@ func TestGetClusterAllowList_PublicEmpty_BYOPOnly(t *testing.T) {
|
||||
},
|
||||
}
|
||||
|
||||
mgr := Manager{proxyManager: pm}
|
||||
mgr := Manager{store: &stubStore{}, proxyManager: pm}
|
||||
result, err := mgr.getClusterAllowList(context.Background(), "acc-123")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []string{"byop.example.com"}, result)
|
||||
}
|
||||
|
||||
// stubStore satisfies the manager's narrow store interface for allow-list
|
||||
// tests. Only the agent-network settings lookup participates; the default (a
|
||||
// nil func) reads as "no settings row", the state most accounts are in.
|
||||
type stubStore struct {
|
||||
getAgentNetworkSettingsFunc func(ctx context.Context, accountID string) (*agentnetworkTypes.Settings, error)
|
||||
}
|
||||
|
||||
func (s *stubStore) GetAccount(context.Context, string) (*types.Account, error) {
|
||||
panic("not used in allow-list tests")
|
||||
}
|
||||
|
||||
func (s *stubStore) GetAgentNetworkSettings(ctx context.Context, _ nbstore.LockingStrength, accountID string) (*agentnetworkTypes.Settings, error) {
|
||||
if s.getAgentNetworkSettingsFunc != nil {
|
||||
return s.getAgentNetworkSettingsFunc(ctx, accountID)
|
||||
}
|
||||
return nil, status.Errorf(status.NotFound, "agent network settings for account %s not found", accountID)
|
||||
}
|
||||
|
||||
func (s *stubStore) GetCustomDomain(context.Context, string, string) (*domain.Domain, error) {
|
||||
panic("not used in allow-list tests")
|
||||
}
|
||||
|
||||
func (s *stubStore) ListFreeDomains(context.Context, string) ([]string, error) {
|
||||
panic("not used in allow-list tests")
|
||||
}
|
||||
|
||||
func (s *stubStore) ListCustomDomains(context.Context, string) ([]*domain.Domain, error) {
|
||||
panic("not used in allow-list tests")
|
||||
}
|
||||
|
||||
func (s *stubStore) CreateCustomDomain(context.Context, string, string, string, bool) (*domain.Domain, error) {
|
||||
panic("not used in allow-list tests")
|
||||
}
|
||||
|
||||
func (s *stubStore) UpdateCustomDomain(context.Context, string, *domain.Domain) (*domain.Domain, error) {
|
||||
panic("not used in allow-list tests")
|
||||
}
|
||||
|
||||
func (s *stubStore) DeleteCustomDomain(context.Context, string, string) error {
|
||||
panic("not used in allow-list tests")
|
||||
}
|
||||
|
||||
// TestGetClusterAllowList_DedicatedGatewayAddressExcluded pins invariant (B)'s
|
||||
// chokepoint: a self-addressed settings pin reserves the account's gateway
|
||||
// address, so it is dropped from the allow list — which, because the
|
||||
// free-domain suffix match is depth-independent, rejects every name beneath
|
||||
// it as well as the bare one. Other addresses are unaffected.
|
||||
func TestGetClusterAllowList_DedicatedGatewayAddressExcluded(t *testing.T) {
|
||||
pm := &mockProxyManager{
|
||||
getActiveClusterAddressesForAccountFunc: func(_ context.Context, _ string) ([]string, error) {
|
||||
return []string{"brave-otter.gateway.example.com", "byop.example.com"}, nil
|
||||
},
|
||||
getActiveClusterAddressesFunc: func(_ context.Context) ([]string, error) {
|
||||
return []string{"eu.proxy.netbird.io"}, nil
|
||||
},
|
||||
}
|
||||
st := &stubStore{
|
||||
getAgentNetworkSettingsFunc: func(_ context.Context, accountID string) (*agentnetworkTypes.Settings, error) {
|
||||
assert.Equal(t, "acc-123", accountID,
|
||||
"the exclusion must look up the requesting account's own settings")
|
||||
return &agentnetworkTypes.Settings{
|
||||
AccountID: accountID,
|
||||
Domain: "brave-otter.gateway.example.com",
|
||||
ProxyAddress: "brave-otter.gateway.example.com",
|
||||
}, nil
|
||||
},
|
||||
}
|
||||
|
||||
mgr := Manager{store: st, proxyManager: pm}
|
||||
result, err := mgr.getClusterAllowList(context.Background(), "acc-123")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []string{"byop.example.com", "eu.proxy.netbird.io"}, result,
|
||||
"the dedicated gateway address must be reserved from cluster selection")
|
||||
}
|
||||
|
||||
// TestGetClusterAllowList_LabeledPinDoesNotExclude pins the counterpart: a
|
||||
// labeled pin means the gateway rides on a shared cluster serving ordinary
|
||||
// services too, so nothing is reserved.
|
||||
func TestGetClusterAllowList_LabeledPinDoesNotExclude(t *testing.T) {
|
||||
pm := &mockProxyManager{
|
||||
getActiveClusterAddressesForAccountFunc: func(_ context.Context, _ string) ([]string, error) {
|
||||
return []string{"byop.example.com"}, nil
|
||||
},
|
||||
getActiveClusterAddressesFunc: func(_ context.Context) ([]string, error) {
|
||||
return []string{"eu.proxy.netbird.io"}, nil
|
||||
},
|
||||
}
|
||||
st := &stubStore{
|
||||
getAgentNetworkSettingsFunc: func(_ context.Context, accountID string) (*agentnetworkTypes.Settings, error) {
|
||||
return &agentnetworkTypes.Settings{
|
||||
AccountID: accountID,
|
||||
Domain: "violet.eu.proxy.netbird.io",
|
||||
ProxyAddress: "eu.proxy.netbird.io",
|
||||
}, nil
|
||||
},
|
||||
}
|
||||
|
||||
mgr := Manager{store: st, proxyManager: pm}
|
||||
result, err := mgr.getClusterAllowList(context.Background(), "acc-123")
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, []string{"byop.example.com", "eu.proxy.netbird.io"}, result,
|
||||
"a labeled pin reserves nothing")
|
||||
}
|
||||
|
||||
// TestGetClusterAllowList_SettingsLookupError_ReturnsError pins that a store
|
||||
// outage is surfaced rather than silently treated as "nothing reserved" —
|
||||
// failing open here would offer a reserved gateway address for ordinary
|
||||
// services.
|
||||
func TestGetClusterAllowList_SettingsLookupError_ReturnsError(t *testing.T) {
|
||||
pm := &mockProxyManager{
|
||||
getActiveClusterAddressesForAccountFunc: func(_ context.Context, _ string) ([]string, error) {
|
||||
return []string{"byop.example.com"}, nil
|
||||
},
|
||||
getActiveClusterAddressesFunc: func(_ context.Context) ([]string, error) {
|
||||
return []string{"eu.proxy.netbird.io"}, nil
|
||||
},
|
||||
}
|
||||
st := &stubStore{
|
||||
getAgentNetworkSettingsFunc: func(_ context.Context, _ string) (*agentnetworkTypes.Settings, error) {
|
||||
return nil, status.Errorf(status.Internal, "store outage")
|
||||
},
|
||||
}
|
||||
|
||||
mgr := Manager{store: st, proxyManager: pm}
|
||||
result, err := mgr.getClusterAllowList(context.Background(), "acc-123")
|
||||
require.Error(t, err)
|
||||
assert.Nil(t, result)
|
||||
assert.Contains(t, err.Error(), "agent network settings")
|
||||
}
|
||||
|
||||
@@ -27,6 +27,8 @@ import (
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
|
||||
accesslogsmanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs/manager"
|
||||
proxyactivity "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/activity"
|
||||
proxyactivitymanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/activity/manager"
|
||||
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
|
||||
"github.com/netbirdio/netbird/management/server/activity"
|
||||
@@ -231,6 +233,7 @@ func (s *BaseServer) ReverseProxyGRPCServer() *nbgrpc.ProxyServiceServer {
|
||||
proxyService := nbgrpc.NewProxyServiceServer(s.AccessLogsManager(), s.ProxyTokenStore(), s.PKCEVerifierStore(), s.proxyOIDCConfig(), s.PeersManager(), s.UsersManager(), s.IdpManager(), s.ProxyManager(), s.Store())
|
||||
s.AfterInit(func(s *BaseServer) {
|
||||
proxyService.SetServiceManager(s.ServiceManager())
|
||||
proxyService.SetActivityManager(s.ProxyActivityManager())
|
||||
proxyService.SetProxyController(s.ServiceProxyController())
|
||||
proxyService.SetAgentNetworkSynthesizer(newAgentNetworkSynthesizer(s.Store()))
|
||||
proxyService.SetAgentNetworkLimitsService(s.AgentNetworkManager())
|
||||
@@ -290,6 +293,13 @@ func (s *BaseServer) PKCEVerifierStore() *nbgrpc.PKCEVerifierStore {
|
||||
})
|
||||
}
|
||||
|
||||
// ProxyActivityManager records reverse proxy usage for activity accounting.
|
||||
func (s *BaseServer) ProxyActivityManager() proxyactivity.Manager {
|
||||
return Create(s, func() proxyactivity.Manager {
|
||||
return proxyactivitymanager.NewManager(s.Store())
|
||||
})
|
||||
}
|
||||
|
||||
func (s *BaseServer) AccessLogsManager() accesslogs.Manager {
|
||||
return Create(s, func() accesslogs.Manager {
|
||||
accessLogManager := accesslogsmanager.NewManager(s.Store(), s.PermissionsManager(), s.GeoLocationManager())
|
||||
|
||||
@@ -13,7 +13,6 @@ import (
|
||||
"math"
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"net/url"
|
||||
"os"
|
||||
"strconv"
|
||||
@@ -26,7 +25,6 @@ import (
|
||||
log "github.com/sirupsen/logrus"
|
||||
"golang.org/x/oauth2"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/metadata"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
"github.com/netbirdio/netbird/shared/management/domain"
|
||||
@@ -34,6 +32,7 @@ import (
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/peers"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/activity"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
|
||||
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/sessionkey"
|
||||
@@ -63,6 +62,17 @@ type ProxyTokenChecker interface {
|
||||
IsProxyAccessTokenValid(ctx context.Context, tokenID string) (bool, error)
|
||||
}
|
||||
|
||||
// ProxyConnectAuthorizer authorizes a proxy's claim to the cluster address it
|
||||
// declares at connect time. Implementations are supplied by integrations; none
|
||||
// is installed by default, so every well-formed claim is authorized — the
|
||||
// declared address is otherwise only checked for availability. token is nil
|
||||
// when the connection carries no proxy access token. A returned status error
|
||||
// is sent to the proxy unchanged; any other error is wrapped as
|
||||
// PermissionDenied.
|
||||
type ProxyConnectAuthorizer interface {
|
||||
AuthorizeProxyConnect(ctx context.Context, token *types.ProxyAccessToken, proxyID, address string) error
|
||||
}
|
||||
|
||||
// ProxyServiceServer implements the ProxyService gRPC server
|
||||
// AgentNetworkSynthesizer produces in-memory reverse-proxy services from
|
||||
// Agent Network provider/policy state for the proxy snapshot path; synthesised
|
||||
@@ -101,6 +111,9 @@ type ProxyServiceServer struct {
|
||||
// and the post-flight consumption write (RecordLLMUsage). Optional — when
|
||||
// nil both RPCs return Unimplemented.
|
||||
agentNetworkLimits AgentNetworkLimitsService
|
||||
// connectAuthorizer authorizes address claims at proxy connect time.
|
||||
// Optional — when nil every well-formed claim is authorized.
|
||||
connectAuthorizer ProxyConnectAuthorizer
|
||||
// ProxyController for service updates and cluster management
|
||||
proxyController proxy.Controller
|
||||
|
||||
@@ -116,6 +129,9 @@ type ProxyServiceServer struct {
|
||||
// Manager for IdP-enriched user data (may be nil when no IdP is configured)
|
||||
idpManager idp.Manager
|
||||
|
||||
// Manager that records reverse proxy usage for activity accounting
|
||||
activityManager activity.Manager
|
||||
|
||||
// Store for one-time authentication tokens
|
||||
tokenStore *OneTimeTokenStore
|
||||
|
||||
@@ -136,10 +152,6 @@ type ProxyServiceServer struct {
|
||||
// initial snapshot delivery. Configurable via NB_PROXY_SNAPSHOT_BATCH_SIZE.
|
||||
snapshotBatchSize int
|
||||
|
||||
authAttemptLimiter *authFailureLimiter
|
||||
authClientLimiter *authFailureLimiter
|
||||
authFailureMAC []byte
|
||||
|
||||
cancel context.CancelFunc
|
||||
}
|
||||
|
||||
@@ -210,10 +222,6 @@ func NewProxyServiceServer(accessLogMgr accesslogs.Manager, tokenStore *OneTimeT
|
||||
snapshotBatchSize: snapshotBatchSizeFromEnv(),
|
||||
cancel: cancel,
|
||||
}
|
||||
s.authAttemptLimiter = newAuthFailureLimiter()
|
||||
s.authClientLimiter = newAuthClientLimiter()
|
||||
s.authFailureMAC = make([]byte, sha256.Size)
|
||||
_, _ = rand.Read(s.authFailureMAC)
|
||||
go s.cleanupStaleProxies(ctx)
|
||||
return s
|
||||
}
|
||||
@@ -246,6 +254,13 @@ func (s *ProxyServiceServer) SetServiceManager(manager rpservice.Manager) {
|
||||
s.serviceManager = manager
|
||||
}
|
||||
|
||||
// SetActivityManager wires the manager that records reverse proxy usage.
|
||||
func (s *ProxyServiceServer) SetActivityManager(manager activity.Manager) {
|
||||
s.mu.Lock()
|
||||
defer s.mu.Unlock()
|
||||
s.activityManager = manager
|
||||
}
|
||||
|
||||
// SetAgentNetworkSynthesizer wires the agent-network service synthesiser.
|
||||
// Optional — when nil the snapshot path skips agent-network synthesis. The
|
||||
// modules layer injects this after both the proxy server and the agent-network
|
||||
@@ -272,6 +287,23 @@ func (s *ProxyServiceServer) agentNetworkSynthesizer() AgentNetworkSynthesizer {
|
||||
return s.agentNetworkSynth
|
||||
}
|
||||
|
||||
// SetProxyConnectAuthorizer wires the connect-time address-claim authorizer.
|
||||
// Optional — when nil (the default) every well-formed claim is authorized,
|
||||
// which is the behavior without the hook. The modules layer injects this
|
||||
// after the proxy server is constructed, like the other setters.
|
||||
func (s *ProxyServiceServer) SetProxyConnectAuthorizer(authorizer ProxyConnectAuthorizer) {
|
||||
s.mu.Lock()
|
||||
s.connectAuthorizer = authorizer
|
||||
s.mu.Unlock()
|
||||
}
|
||||
|
||||
// proxyConnectAuthorizer returns the connect authorizer under read lock.
|
||||
func (s *ProxyServiceServer) proxyConnectAuthorizer() ProxyConnectAuthorizer {
|
||||
s.mu.RLock()
|
||||
defer s.mu.RUnlock()
|
||||
return s.connectAuthorizer
|
||||
}
|
||||
|
||||
// CheckLLMPolicyLimits is the pre-flight policy gate the proxy calls before
|
||||
// forwarding an LLM request upstream. Delegates to the agent-network selector,
|
||||
// which scores applicable policies by remaining headroom and returns the
|
||||
@@ -456,8 +488,9 @@ func recvSyncInit(stream proto.ProxyService_SyncMappingsServer) (*proto.SyncMapp
|
||||
return init, nil
|
||||
}
|
||||
|
||||
// validateProxyConnect validates the proxy ID and address, and checks cluster
|
||||
// address availability for account-scoped tokens.
|
||||
// validateProxyConnect validates the proxy ID and address, checks cluster
|
||||
// address availability for account-scoped tokens, and finally consults the
|
||||
// connect authorizer (when installed) on the address claim.
|
||||
func (s *ProxyServiceServer) validateProxyConnect(proxyID, address string, ctx context.Context) (proxyConnectParams, error) {
|
||||
if proxyID == "" {
|
||||
return proxyConnectParams{}, status.Errorf(codes.InvalidArgument, "proxy_id is required")
|
||||
@@ -477,6 +510,19 @@ func (s *ProxyServiceServer) validateProxyConnect(proxyID, address string, ctx c
|
||||
}
|
||||
}
|
||||
|
||||
// The authorizer runs last, outside the account-scoped branch, so it also
|
||||
// sees management-wide and token-less connects. PermissionDenied keeps an
|
||||
// authorization rejection distinguishable from the AlreadyExists address
|
||||
// conflict above in proxy logs.
|
||||
if authorizer := s.proxyConnectAuthorizer(); authorizer != nil {
|
||||
if err := authorizer.AuthorizeProxyConnect(ctx, token, proxyID, address); err != nil {
|
||||
if _, ok := status.FromError(err); ok {
|
||||
return proxyConnectParams{}, err
|
||||
}
|
||||
return proxyConnectParams{}, status.Errorf(codes.PermissionDenied, "proxy connect not authorized: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
return proxyConnectParams{proxyID: proxyID, address: address}, nil
|
||||
}
|
||||
|
||||
@@ -1182,18 +1228,6 @@ func (s *ProxyServiceServer) Authenticate(ctx context.Context, req *proto.Authen
|
||||
return nil, err
|
||||
}
|
||||
|
||||
failureKey := s.authFailureKey(req)
|
||||
limitFailures := failureKey != "" && s.authAttemptLimiter != nil && len(s.authFailureMAC) > 0
|
||||
if limitFailures && s.authAttemptLimiter.isLimited(failureKey) {
|
||||
return nil, status.Errorf(codes.ResourceExhausted, "too many failed authentication attempts for this credential, please try again later")
|
||||
}
|
||||
|
||||
clientKey := s.authClientKey(ctx, req.GetId())
|
||||
limitClient := clientKey != "" && s.authClientLimiter != nil
|
||||
if limitClient && s.authClientLimiter.isLimited(clientKey) {
|
||||
return nil, status.Errorf(codes.ResourceExhausted, "too many failed authentication attempts from this client, please try again later")
|
||||
}
|
||||
|
||||
service, err := s.serviceManager.GetServiceByID(ctx, req.GetAccountId(), req.GetId())
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Debugf("failed to get service from store: %v", err)
|
||||
@@ -1201,14 +1235,6 @@ func (s *ProxyServiceServer) Authenticate(ctx context.Context, req *proto.Authen
|
||||
}
|
||||
|
||||
authenticated, userId, method := s.authenticateRequest(ctx, req, service)
|
||||
if !authenticated {
|
||||
if limitFailures {
|
||||
s.authAttemptLimiter.recordFailure(failureKey)
|
||||
}
|
||||
if limitClient {
|
||||
s.authClientLimiter.recordFailure(clientKey)
|
||||
}
|
||||
}
|
||||
|
||||
// Non-OIDC schemes (PIN/Password/Header) authenticate against per-service
|
||||
// secrets and have no user-level group context, so groups stay nil. Email
|
||||
@@ -1224,40 +1250,6 @@ func (s *ProxyServiceServer) Authenticate(ctx context.Context, req *proto.Authen
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *ProxyServiceServer) authClientKey(ctx context.Context, serviceID string) string {
|
||||
md, ok := metadata.FromIncomingContext(ctx)
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
values := md.Get(proxyauth.ClientIPMetadataKey)
|
||||
if len(values) == 0 {
|
||||
return ""
|
||||
}
|
||||
addr, err := netip.ParseAddr(strings.TrimSpace(values[0]))
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
return serviceID + "|" + addr.Unmap().String()
|
||||
}
|
||||
|
||||
func (s *ProxyServiceServer) authFailureKey(req *proto.AuthenticateRequest) string {
|
||||
var secret string
|
||||
switch v := req.GetRequest().(type) {
|
||||
case *proto.AuthenticateRequest_Pin:
|
||||
secret = "pin|" + v.Pin.GetPin()
|
||||
case *proto.AuthenticateRequest_Password:
|
||||
secret = "password|" + v.Password.GetPassword()
|
||||
case *proto.AuthenticateRequest_HeaderAuth:
|
||||
secret = "header|" + v.HeaderAuth.GetHeaderName() + "|" + v.HeaderAuth.GetHeaderValue()
|
||||
default:
|
||||
return ""
|
||||
}
|
||||
|
||||
mac := hmac.New(sha256.New, s.authFailureMAC)
|
||||
mac.Write([]byte(secret))
|
||||
return req.GetId() + "|" + hex.EncodeToString(mac.Sum(nil))
|
||||
}
|
||||
|
||||
func (s *ProxyServiceServer) authenticateRequest(ctx context.Context, req *proto.AuthenticateRequest, service *rpservice.Service) (bool, string, proxyauth.Method) {
|
||||
switch v := req.GetRequest().(type) {
|
||||
case *proto.AuthenticateRequest_Pin:
|
||||
@@ -1736,7 +1728,7 @@ func (s *ProxyServiceServer) GenerateSessionToken(ctx context.Context, domain, u
|
||||
|
||||
groupIDs, groupNames := pairGroupIDsAndNames(userGroups)
|
||||
|
||||
return sessionkey.SignToken(
|
||||
token, err := sessionkey.SignToken(
|
||||
service.SessionPrivateKey,
|
||||
userID,
|
||||
user.Email,
|
||||
@@ -1746,6 +1738,25 @@ func (s *ProxyServiceServer) GenerateSessionToken(ctx context.Context, domain, u
|
||||
groupNames,
|
||||
proxyauth.DefaultSessionExpiry,
|
||||
)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
s.recordUserLogin(ctx, service.AccountID, user)
|
||||
|
||||
return token, nil
|
||||
}
|
||||
|
||||
// recordUserLogin hands the sign-in to the activity manager. The RPC must not
|
||||
// fail on it, so the error is logged and dropped here rather than returned.
|
||||
func (s *ProxyServiceServer) recordUserLogin(ctx context.Context, accountID string, user *types.User) {
|
||||
if s.activityManager == nil {
|
||||
return
|
||||
}
|
||||
|
||||
if err := s.activityManager.RecordUserLogin(ctx, accountID, user); err != nil {
|
||||
log.WithContext(ctx).Debugf("record proxy login for user %s: %v", user.Id, err)
|
||||
}
|
||||
}
|
||||
|
||||
// ValidateUserGroupAccess checks if a user has access to a service.
|
||||
@@ -2095,6 +2106,8 @@ func (s *ProxyServiceServer) ValidateTunnelPeer(ctx context.Context, req *proto.
|
||||
return nil, err
|
||||
}
|
||||
|
||||
s.recordPeerSeen(ctx, service.AccountID, peer)
|
||||
|
||||
log.WithFields(log.Fields{
|
||||
"domain": domain,
|
||||
"tunnel_ip": tunnelIPStr,
|
||||
@@ -2112,6 +2125,18 @@ func (s *ProxyServiceServer) ValidateTunnelPeer(ctx context.Context, req *proto.
|
||||
}, nil
|
||||
}
|
||||
|
||||
// recordPeerSeen hands the mesh request to the activity manager. The RPC must
|
||||
// not fail on it, so the error is logged and dropped here rather than returned.
|
||||
func (s *ProxyServiceServer) recordPeerSeen(ctx context.Context, accountID string, peer *peer.Peer) {
|
||||
if s.activityManager == nil {
|
||||
return
|
||||
}
|
||||
|
||||
if err := s.activityManager.RecordPeerSeen(ctx, accountID, peer); err != nil {
|
||||
log.WithContext(ctx).Debugf("record proxy activity for peer %s: %v", peer.ID, err)
|
||||
}
|
||||
}
|
||||
|
||||
// resolvePeerOwner returns the user a peer is linked to, once per request so
|
||||
// the status gate and the identity resolution below share a single lookup.
|
||||
// Unlinked peers (machine agents) have no owner. A lookup that fails returns
|
||||
|
||||
@@ -1,189 +0,0 @@
|
||||
package grpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"fmt"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"golang.org/x/time/rate"
|
||||
"google.golang.org/grpc/codes"
|
||||
"google.golang.org/grpc/metadata"
|
||||
"google.golang.org/grpc/status"
|
||||
|
||||
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
proxyauth "github.com/netbirdio/netbird/proxy/auth"
|
||||
"github.com/netbirdio/netbird/shared/hash/argon2id"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
const authAttemptsHeaderName = "X-API-Key"
|
||||
|
||||
func newAuthAttemptsTestServer(t *testing.T) *ProxyServiceServer {
|
||||
t.Helper()
|
||||
|
||||
firstHash, err := argon2id.Hash("first-key")
|
||||
require.NoError(t, err)
|
||||
secondHash, err := argon2id.Hash("second-key")
|
||||
require.NoError(t, err)
|
||||
|
||||
svc := &rpservice.Service{
|
||||
ID: "svc1",
|
||||
Domain: "example.com",
|
||||
Auth: rpservice.AuthConfig{
|
||||
HeaderAuths: []*rpservice.HeaderAuthConfig{
|
||||
{Enabled: true, Header: authAttemptsHeaderName, Value: firstHash},
|
||||
{Enabled: true, Header: authAttemptsHeaderName, Value: secondHash},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
ctrl := gomock.NewController(t)
|
||||
mgr := rpservice.NewMockManager(ctrl)
|
||||
mgr.EXPECT().GetServiceByID(gomock.Any(), gomock.Any(), gomock.Any()).Return(svc, nil).AnyTimes()
|
||||
|
||||
limiter := newAuthFailureLimiter()
|
||||
t.Cleanup(limiter.stop)
|
||||
clientLimiter := newAuthClientLimiter()
|
||||
t.Cleanup(clientLimiter.stop)
|
||||
|
||||
mac := make([]byte, sha256.Size)
|
||||
_, err = rand.Read(mac)
|
||||
require.NoError(t, err)
|
||||
|
||||
return &ProxyServiceServer{
|
||||
serviceManager: mgr,
|
||||
authAttemptLimiter: limiter,
|
||||
authClientLimiter: clientLimiter,
|
||||
authFailureMAC: mac,
|
||||
}
|
||||
}
|
||||
|
||||
func clientIPContext(ip string) context.Context {
|
||||
return metadata.NewIncomingContext(context.Background(), metadata.Pairs(proxyauth.ClientIPMetadataKey, ip))
|
||||
}
|
||||
|
||||
func authAttemptsRequest(credential string) *proto.AuthenticateRequest {
|
||||
return &proto.AuthenticateRequest{
|
||||
Id: "svc1",
|
||||
AccountId: "acc1",
|
||||
Request: &proto.AuthenticateRequest_HeaderAuth{
|
||||
HeaderAuth: &proto.HeaderAuthRequest{
|
||||
HeaderName: authAttemptsHeaderName,
|
||||
HeaderValue: credential,
|
||||
},
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthenticate_ValidCredentialIsNeverRateLimited(t *testing.T) {
|
||||
s := newAuthAttemptsTestServer(t)
|
||||
|
||||
for i := 0; i < proxyAuthFailureBurst*2; i++ {
|
||||
resp, err := s.Authenticate(context.Background(), authAttemptsRequest("first-key"))
|
||||
require.NoError(t, err, "a valid credential must never be throttled (attempt %d)", i)
|
||||
require.True(t, resp.GetSuccess())
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthenticate_FailedCredentialIsRateLimited(t *testing.T) {
|
||||
s := newAuthAttemptsTestServer(t)
|
||||
|
||||
for i := 0; i < proxyAuthFailureBurst; i++ {
|
||||
resp, err := s.Authenticate(context.Background(), authAttemptsRequest("wrong-key"))
|
||||
require.NoError(t, err, "attempt %d should be within the failure budget", i)
|
||||
require.False(t, resp.GetSuccess())
|
||||
}
|
||||
|
||||
_, err := s.Authenticate(context.Background(), authAttemptsRequest("wrong-key"))
|
||||
require.Error(t, err)
|
||||
assert.Equal(t, codes.ResourceExhausted, status.Code(err))
|
||||
}
|
||||
|
||||
func TestAuthenticate_ThrottledCredentialDoesNotAffectOthers(t *testing.T) {
|
||||
s := newAuthAttemptsTestServer(t)
|
||||
|
||||
for i := 0; i < proxyAuthFailureBurst+2; i++ {
|
||||
_, _ = s.Authenticate(context.Background(), authAttemptsRequest("wrong-key"))
|
||||
}
|
||||
|
||||
resp, err := s.Authenticate(context.Background(), authAttemptsRequest("first-key"))
|
||||
require.NoError(t, err, "one throttled credential must not block a valid one")
|
||||
assert.True(t, resp.GetSuccess())
|
||||
|
||||
resp, err = s.Authenticate(context.Background(), authAttemptsRequest("second-key"))
|
||||
require.NoError(t, err)
|
||||
assert.True(t, resp.GetSuccess())
|
||||
|
||||
_, err = s.Authenticate(context.Background(), authAttemptsRequest("another-wrong-key"))
|
||||
require.NoError(t, err, "a different failing credential has its own budget")
|
||||
}
|
||||
|
||||
func TestAuthenticate_DistinctCredentialsThrottledPerClient(t *testing.T) {
|
||||
const budget = 3
|
||||
|
||||
s := newAuthAttemptsTestServer(t)
|
||||
s.authClientLimiter.stop()
|
||||
s.authClientLimiter = newAuthLimiter(rate.Every(time.Hour), budget)
|
||||
t.Cleanup(s.authClientLimiter.stop)
|
||||
|
||||
ctx := clientIPContext("198.51.100.7")
|
||||
|
||||
for i := 0; i < budget; i++ {
|
||||
resp, err := s.Authenticate(ctx, authAttemptsRequest(fmt.Sprintf("garbage-%d", i)))
|
||||
require.NoError(t, err, "attempt %d should be within the client budget", i)
|
||||
require.False(t, resp.GetSuccess())
|
||||
}
|
||||
|
||||
_, err := s.Authenticate(ctx, authAttemptsRequest("garbage-final"))
|
||||
require.Error(t, err, "a client rotating distinct credentials must be throttled")
|
||||
assert.Equal(t, codes.ResourceExhausted, status.Code(err))
|
||||
|
||||
other := clientIPContext("198.51.100.8")
|
||||
resp, err := s.Authenticate(other, authAttemptsRequest("first-key"))
|
||||
require.NoError(t, err, "a different client must be unaffected")
|
||||
assert.True(t, resp.GetSuccess())
|
||||
}
|
||||
|
||||
func TestAuthenticate_OneStaleCredentialDoesNotExhaustSharedClientBudget(t *testing.T) {
|
||||
s := newAuthAttemptsTestServer(t)
|
||||
ctx := clientIPContext("198.51.100.9")
|
||||
|
||||
for i := 0; i < proxyAuthFailureBurst*4; i++ {
|
||||
_, _ = s.Authenticate(ctx, authAttemptsRequest("stale-key"))
|
||||
}
|
||||
|
||||
resp, err := s.Authenticate(ctx, authAttemptsRequest("first-key"))
|
||||
require.NoError(t, err, "one client stuck on a stale key must not block others behind the same NAT")
|
||||
assert.True(t, resp.GetSuccess())
|
||||
}
|
||||
|
||||
func TestAuthenticate_ProxyWithoutClientIPIsNotClientLimited(t *testing.T) {
|
||||
s := newAuthAttemptsTestServer(t)
|
||||
|
||||
for i := 0; i < proxyAuthFailureBurst*2; i++ {
|
||||
_, _ = s.Authenticate(context.Background(), authAttemptsRequest(fmt.Sprintf("garbage-%d", i)))
|
||||
}
|
||||
|
||||
resp, err := s.Authenticate(context.Background(), authAttemptsRequest("first-key"))
|
||||
require.NoError(t, err, "an old proxy must not have its clients share one budget")
|
||||
assert.True(t, resp.GetSuccess())
|
||||
}
|
||||
|
||||
func TestAuthenticate_MalformedClientIPIsIgnored(t *testing.T) {
|
||||
s := newAuthAttemptsTestServer(t)
|
||||
ctx := clientIPContext("not-an-ip")
|
||||
|
||||
for i := 0; i < proxyAuthFailureBurst*2; i++ {
|
||||
_, _ = s.Authenticate(ctx, authAttemptsRequest(fmt.Sprintf("garbage-%d", i)))
|
||||
}
|
||||
|
||||
resp, err := s.Authenticate(ctx, authAttemptsRequest("first-key"))
|
||||
require.NoError(t, err)
|
||||
assert.True(t, resp.GetSuccess())
|
||||
}
|
||||
@@ -18,16 +18,12 @@ const (
|
||||
proxyAuthLimiterCleanup = 5 * time.Minute
|
||||
// proxyAuthLimiterTTL is how long a limiter is kept after the last failure.
|
||||
proxyAuthLimiterTTL = 15 * time.Minute
|
||||
|
||||
proxyAuthClientBurst = 30
|
||||
)
|
||||
|
||||
// defaultProxyAuthFailureRate is the token replenishment rate for failed auth attempts.
|
||||
// One token every 12 seconds = 5 per minute.
|
||||
var defaultProxyAuthFailureRate = rate.Every(12 * time.Second)
|
||||
|
||||
var defaultProxyAuthClientRate = rate.Limit(1)
|
||||
|
||||
// clientIP identifies a client by its IP address for rate limiting purposes.
|
||||
type clientIP = string
|
||||
|
||||
@@ -41,7 +37,6 @@ type authFailureLimiter struct {
|
||||
mu sync.Mutex
|
||||
limiters map[clientIP]*limiterEntry
|
||||
failureRate rate.Limit
|
||||
burst int
|
||||
cancel context.CancelFunc
|
||||
}
|
||||
|
||||
@@ -50,19 +45,10 @@ func newAuthFailureLimiter() *authFailureLimiter {
|
||||
}
|
||||
|
||||
func newAuthFailureLimiterWithRate(failureRate rate.Limit) *authFailureLimiter {
|
||||
return newAuthLimiter(failureRate, proxyAuthFailureBurst)
|
||||
}
|
||||
|
||||
func newAuthClientLimiter() *authFailureLimiter {
|
||||
return newAuthLimiter(defaultProxyAuthClientRate, proxyAuthClientBurst)
|
||||
}
|
||||
|
||||
func newAuthLimiter(failureRate rate.Limit, burst int) *authFailureLimiter {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
l := &authFailureLimiter{
|
||||
limiters: make(map[clientIP]*limiterEntry),
|
||||
failureRate: failureRate,
|
||||
burst: burst,
|
||||
cancel: cancel,
|
||||
}
|
||||
go l.cleanupLoop(ctx)
|
||||
@@ -91,7 +77,7 @@ func (l *authFailureLimiter) recordFailure(ip clientIP) {
|
||||
entry, exists := l.limiters[ip]
|
||||
if !exists {
|
||||
entry = &limiterEntry{
|
||||
limiter: rate.NewLimiter(l.failureRate, l.burst),
|
||||
limiter: rate.NewLimiter(l.failureRate, proxyAuthFailureBurst),
|
||||
}
|
||||
l.limiters[ip] = entry
|
||||
}
|
||||
|
||||
@@ -0,0 +1,168 @@
|
||||
package grpc
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"github.com/golang/mock/gomock"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
"google.golang.org/grpc/codes"
|
||||
grpcstatus "google.golang.org/grpc/status"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
)
|
||||
|
||||
// capturingAuthorizer records the arguments of the last AuthorizeProxyConnect
|
||||
// call and returns a fixed error.
|
||||
type capturingAuthorizer struct {
|
||||
called int
|
||||
token *types.ProxyAccessToken
|
||||
proxyID string
|
||||
address string
|
||||
err error
|
||||
}
|
||||
|
||||
func (a *capturingAuthorizer) AuthorizeProxyConnect(_ context.Context, token *types.ProxyAccessToken, proxyID, address string) error {
|
||||
a.called++
|
||||
a.token = token
|
||||
a.proxyID = proxyID
|
||||
a.address = address
|
||||
return a.err
|
||||
}
|
||||
|
||||
// authorizerServer builds a ProxyServiceServer whose proxy manager reports
|
||||
// every cluster address as available, so the authorizer is the only thing
|
||||
// standing between a claim and success.
|
||||
func authorizerServer(t *testing.T) *ProxyServiceServer {
|
||||
t.Helper()
|
||||
ctrl := gomock.NewController(t)
|
||||
mgr := proxy.NewMockManager(ctrl)
|
||||
mgr.EXPECT().IsClusterAddressAvailable(gomock.Any(), gomock.Any(), gomock.Any()).Return(true, nil).AnyTimes()
|
||||
return &ProxyServiceServer{proxyManager: mgr}
|
||||
}
|
||||
|
||||
// TestValidateProxyConnect_NilAuthorizerUnchanged guards the no-behavior-change
|
||||
// claim: with no authorizer installed — the OSS default — a well-formed claim
|
||||
// succeeds and malformed input is rejected exactly as before the hook existed.
|
||||
func TestValidateProxyConnect_NilAuthorizerUnchanged(t *testing.T) {
|
||||
s := authorizerServer(t)
|
||||
|
||||
params, err := s.validateProxyConnect("proxy-1", "cluster.example.com", scopedCtx("acc-1"))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "proxy-1", params.proxyID)
|
||||
assert.Equal(t, "cluster.example.com", params.address)
|
||||
|
||||
_, err = s.validateProxyConnect("", "cluster.example.com", scopedCtx("acc-1"))
|
||||
require.Error(t, err)
|
||||
st, ok := grpcstatus.FromError(err)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, codes.InvalidArgument, st.Code(), "missing proxy_id must stay InvalidArgument")
|
||||
|
||||
_, err = s.validateProxyConnect("proxy-1", "not a hostname", scopedCtx("acc-1"))
|
||||
require.Error(t, err)
|
||||
st, ok = grpcstatus.FromError(err)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, codes.InvalidArgument, st.Code(), "invalid address must stay InvalidArgument")
|
||||
}
|
||||
|
||||
// TestValidateProxyConnect_AuthorizerReceivesClaim pins the hook contract: the
|
||||
// authorizer sees the presented token and the claimed proxy ID and address,
|
||||
// and an authorized claim proceeds.
|
||||
func TestValidateProxyConnect_AuthorizerReceivesClaim(t *testing.T) {
|
||||
s := authorizerServer(t)
|
||||
auth := &capturingAuthorizer{}
|
||||
s.SetProxyConnectAuthorizer(auth)
|
||||
|
||||
params, err := s.validateProxyConnect("proxy-1", "cluster.example.com", scopedCtx("acc-1"))
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, "cluster.example.com", params.address)
|
||||
|
||||
require.Equal(t, 1, auth.called, "authorizer must be consulted exactly once per connect")
|
||||
assert.Equal(t, "proxy-1", auth.proxyID)
|
||||
assert.Equal(t, "cluster.example.com", auth.address)
|
||||
require.NotNil(t, auth.token, "the presented token must be handed to the authorizer")
|
||||
require.NotNil(t, auth.token.AccountID)
|
||||
assert.Equal(t, "acc-1", *auth.token.AccountID)
|
||||
}
|
||||
|
||||
// TestValidateProxyConnect_PlainErrorBecomesPermissionDenied pins the error
|
||||
// mapping: a non-status error from the authorizer surfaces as
|
||||
// PermissionDenied — distinguishable from the AlreadyExists used for address
|
||||
// conflicts — and the claim does not proceed even though the address itself
|
||||
// was available.
|
||||
func TestValidateProxyConnect_PlainErrorBecomesPermissionDenied(t *testing.T) {
|
||||
s := authorizerServer(t)
|
||||
s.SetProxyConnectAuthorizer(&capturingAuthorizer{err: errors.New("not the assigned credential")})
|
||||
|
||||
_, err := s.validateProxyConnect("proxy-1", "cluster.example.com", scopedCtx("acc-1"))
|
||||
require.Error(t, err)
|
||||
st, ok := grpcstatus.FromError(err)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, codes.PermissionDenied, st.Code())
|
||||
assert.Contains(t, st.Message(), "not the assigned credential", "the authorizer's reason must survive into the status message")
|
||||
}
|
||||
|
||||
// TestValidateProxyConnect_StatusErrorPassesThrough pins that an authorizer
|
||||
// which chooses its own status code is not second-guessed.
|
||||
func TestValidateProxyConnect_StatusErrorPassesThrough(t *testing.T) {
|
||||
s := authorizerServer(t)
|
||||
s.SetProxyConnectAuthorizer(&capturingAuthorizer{
|
||||
err: grpcstatus.Errorf(codes.ResourceExhausted, "try later"),
|
||||
})
|
||||
|
||||
_, err := s.validateProxyConnect("proxy-1", "cluster.example.com", scopedCtx("acc-1"))
|
||||
require.Error(t, err)
|
||||
st, ok := grpcstatus.FromError(err)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, codes.ResourceExhausted, st.Code(), "a status error must pass through unchanged")
|
||||
assert.Equal(t, "try later", st.Message())
|
||||
}
|
||||
|
||||
// TestValidateProxyConnect_AuthorizerSeesGlobalAndTokenlessConnects pins the
|
||||
// call-site placement: the authorizer sits outside the account-scoped branch,
|
||||
// so management-wide tokens (AccountID == nil) and connections without any
|
||||
// token are also presented to it rather than bypassing policy.
|
||||
func TestValidateProxyConnect_AuthorizerSeesGlobalAndTokenlessConnects(t *testing.T) {
|
||||
s := &ProxyServiceServer{} // no proxy manager: neither path may reach the availability check
|
||||
|
||||
auth := &capturingAuthorizer{}
|
||||
s.SetProxyConnectAuthorizer(auth)
|
||||
|
||||
_, err := s.validateProxyConnect("proxy-1", "cluster.example.com", globalCtx())
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 1, auth.called, "a management-wide token must still be presented to the authorizer")
|
||||
require.NotNil(t, auth.token)
|
||||
assert.Nil(t, auth.token.AccountID)
|
||||
|
||||
_, err = s.validateProxyConnect("proxy-1", "cluster.example.com", context.Background())
|
||||
require.NoError(t, err)
|
||||
require.Equal(t, 2, auth.called, "a token-less connect must still be presented to the authorizer")
|
||||
assert.Nil(t, auth.token, "no token in context must surface as a nil token, not a zero value")
|
||||
}
|
||||
|
||||
// TestValidateProxyConnect_AuthorizerRunsLast pins the ordering: input
|
||||
// validation and the availability check precede policy, so the authorizer is
|
||||
// never consulted about a claim that is malformed or already rejected.
|
||||
func TestValidateProxyConnect_AuthorizerRunsLast(t *testing.T) {
|
||||
ctrl := gomock.NewController(t)
|
||||
mgr := proxy.NewMockManager(ctrl)
|
||||
mgr.EXPECT().IsClusterAddressAvailable(gomock.Any(), gomock.Any(), gomock.Any()).Return(false, nil)
|
||||
s := &ProxyServiceServer{proxyManager: mgr}
|
||||
|
||||
auth := &capturingAuthorizer{}
|
||||
s.SetProxyConnectAuthorizer(auth)
|
||||
|
||||
_, err := s.validateProxyConnect("proxy-1", "not a hostname", scopedCtx("acc-1"))
|
||||
require.Error(t, err)
|
||||
assert.Zero(t, auth.called, "a malformed address must be rejected before policy runs")
|
||||
|
||||
_, err = s.validateProxyConnect("proxy-1", "cluster.example.com", scopedCtx("acc-1"))
|
||||
require.Error(t, err)
|
||||
st, ok := grpcstatus.FromError(err)
|
||||
require.True(t, ok)
|
||||
assert.Equal(t, codes.AlreadyExists, st.Code(), "an address conflict must keep its own status")
|
||||
assert.Zero(t, auth.called, "a conflicting address must be rejected before policy runs")
|
||||
}
|
||||
@@ -5,6 +5,7 @@ import (
|
||||
"errors"
|
||||
"net"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
@@ -155,6 +156,27 @@ type mockTunnelPeersManager struct {
|
||||
groupsErr error
|
||||
}
|
||||
|
||||
// mockActivityManager records what the RPC handed to the activity manager. The
|
||||
// policy (throttling, exclusions) is the manager's and is tested there; these
|
||||
// tests only pin which requests reach it.
|
||||
type mockActivityManager struct {
|
||||
seenMarks []seenMark
|
||||
}
|
||||
|
||||
type seenMark struct {
|
||||
accountID string
|
||||
peerID string
|
||||
}
|
||||
|
||||
func (m *mockActivityManager) RecordUserLogin(_ context.Context, _ string, _ *types.User) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *mockActivityManager) RecordPeerSeen(_ context.Context, accountID string, peer *peer.Peer) error {
|
||||
m.seenMarks = append(m.seenMarks, seenMark{accountID: accountID, peerID: peer.ID})
|
||||
return nil
|
||||
}
|
||||
|
||||
func (m *mockTunnelPeersManager) GetPeerByTunnelIP(_ context.Context, _ string, _ net.IP) (*peer.Peer, error) {
|
||||
return m.peer, m.peerErr
|
||||
}
|
||||
@@ -745,6 +767,78 @@ func TestValidateTunnelPeerOwnerStatus(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
// TestValidateTunnelPeerRecordsActivity pins that a granted mesh request is
|
||||
// handed to the activity manager. Which of those the manager then writes is its
|
||||
// own decision, covered by its tests.
|
||||
func TestValidateTunnelPeerRecordsActivity(t *testing.T) {
|
||||
const (
|
||||
domain = "app.example.com"
|
||||
accountID = "account1"
|
||||
peerID = "peer1"
|
||||
)
|
||||
|
||||
activityManager := &mockActivityManager{}
|
||||
server := &ProxyServiceServer{
|
||||
activityManager: activityManager,
|
||||
serviceManager: &mockReverseProxyManager{
|
||||
proxiesByAccount: map[string][]*service.Service{
|
||||
accountID: {{Domain: domain, AccountID: accountID}},
|
||||
},
|
||||
},
|
||||
peersManager: &mockTunnelPeersManager{
|
||||
peer: &peer.Peer{ID: peerID, Name: "agent", Status: &peer.PeerStatus{LastSeen: time.Now().Add(-3 * time.Hour)}},
|
||||
},
|
||||
usersManager: &mockUsersManager{users: map[string]*types.User{}},
|
||||
}
|
||||
|
||||
resp, err := server.ValidateTunnelPeer(context.Background(), &proto.ValidateTunnelPeerRequest{
|
||||
Domain: domain,
|
||||
TunnelIp: "100.64.0.1",
|
||||
})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.True(t, resp.GetValid(), "peer should be granted access")
|
||||
|
||||
require.Len(t, activityManager.seenMarks, 1, "a granted peer should reach the activity manager once")
|
||||
assert.Equal(t, accountID, activityManager.seenMarks[0].accountID, "activity must be attributed to the service account")
|
||||
assert.Equal(t, peerID, activityManager.seenMarks[0].peerID, "activity must be attributed to the calling peer")
|
||||
}
|
||||
|
||||
// TestValidateTunnelPeerDeniedRecordsNoActivity keeps the write on the granted
|
||||
// path only: a refused peer is not evidence its owner was active.
|
||||
func TestValidateTunnelPeerDeniedRecordsNoActivity(t *testing.T) {
|
||||
const (
|
||||
domain = "app.example.com"
|
||||
accountID = "account1"
|
||||
)
|
||||
|
||||
activityManager := &mockActivityManager{}
|
||||
server := &ProxyServiceServer{
|
||||
activityManager: activityManager,
|
||||
serviceManager: &mockReverseProxyManager{
|
||||
proxiesByAccount: map[string][]*service.Service{
|
||||
accountID: {{Domain: domain, AccountID: accountID}},
|
||||
},
|
||||
},
|
||||
peersManager: &mockTunnelPeersManager{
|
||||
peer: &peer.Peer{ID: "peer1", Name: "agent", UserID: "user1", Status: &peer.PeerStatus{LastSeen: time.Now().Add(-3 * time.Hour)}},
|
||||
},
|
||||
// The owner is blocked, so the tunnel gate denies before the mint.
|
||||
usersManager: &mockUsersManager{users: map[string]*types.User{
|
||||
"user1": {Id: "user1", AccountID: accountID, Blocked: true},
|
||||
}},
|
||||
}
|
||||
|
||||
resp, err := server.ValidateTunnelPeer(context.Background(), &proto.ValidateTunnelPeerRequest{
|
||||
Domain: domain,
|
||||
TunnelIp: "100.64.0.1",
|
||||
})
|
||||
|
||||
require.NoError(t, err)
|
||||
require.False(t, resp.GetValid(), "blocked owner should be denied")
|
||||
assert.Empty(t, activityManager.seenMarks, "a denied peer must not be marked seen")
|
||||
}
|
||||
|
||||
func TestGetAccountProxyByDomain(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
|
||||
@@ -281,6 +281,9 @@ const (
|
||||
// AccountMetricsPushDisabled indicates that a user disabled metrics push for the account
|
||||
AccountMetricsPushDisabled Activity = 141
|
||||
|
||||
// AgentNetworkSettingsDeleted indicates that a user deleted the Agent Network account settings, releasing the endpoint
|
||||
AgentNetworkSettingsDeleted Activity = 142
|
||||
|
||||
AccountDeleted Activity = 99999
|
||||
)
|
||||
|
||||
@@ -453,6 +456,7 @@ var activityMap = map[Activity]Code{
|
||||
AgentNetworkBudgetRuleDeleted: {"Agent Network budget rule deleted", "agent_network.budget_rule.delete"},
|
||||
|
||||
AgentNetworkSettingsUpdated: {"Agent Network settings updated", "agent_network.settings.update"},
|
||||
AgentNetworkSettingsDeleted: {"Agent Network settings deleted", "agent_network.settings.delete"},
|
||||
|
||||
AccountMetricsPushEnabled: {"Account metrics push enabled", "account.setting.metrics.push.enable"},
|
||||
AccountMetricsPushDisabled: {"Account metrics push disabled", "account.setting.metrics.push.disable"},
|
||||
|
||||
@@ -68,7 +68,10 @@ func TestAgentNetwork_BudgetRuleCRUD_RealManager(t *testing.T) {
|
||||
|
||||
// TestAgentNetwork_UpdateSettings_PreservesImmutableAndTogglesCollection is the
|
||||
// GC-1 guard for UpdateSettings: it must apply the collection toggles while
|
||||
// preserving the immutable Cluster/Subdomain pinned at bootstrap.
|
||||
// preserving the immutable Domain/ProxyAddress assigned at bootstrap. The
|
||||
// request echoes the identity fields back — the PUT convention every other
|
||||
// endpoint follows — and a request echoing anything else is rejected outright
|
||||
// rather than quietly ignored.
|
||||
func TestAgentNetwork_UpdateSettings_PreservesImmutableAndTogglesCollection(t *testing.T) {
|
||||
am, _, err := createManager(t)
|
||||
require.NoError(t, err, "createManager must succeed")
|
||||
@@ -84,7 +87,14 @@ func TestAgentNetwork_UpdateSettings_PreservesImmutableAndTogglesCollection(t *t
|
||||
|
||||
mgr := agentnetwork.NewManager(am.Store, permissions.NewManager(am.Store), am, nil)
|
||||
|
||||
// Creating a provider bootstraps the settings row (cluster + subdomain).
|
||||
// Bootstrap is an explicit settings create; providers have no settings
|
||||
// side effects anymore.
|
||||
before, err := mgr.CreateSettings(ctx, adminUserID, agenttypes.DefaultSettings(accountID), clusterAddr, "")
|
||||
require.NoError(t, err, "CreateSettings must bootstrap the row")
|
||||
require.Equal(t, clusterAddr, before.ProxyAddress, "proxy address pinned at bootstrap")
|
||||
require.NotEmpty(t, before.Domain, "endpoint allocated at bootstrap")
|
||||
assert.False(t, before.EnablePromptCollection, "prompt collection defaults off")
|
||||
|
||||
_, err = mgr.CreateProvider(ctx, adminUserID, &agenttypes.Provider{
|
||||
AccountID: accountID,
|
||||
ProviderID: "openai_api",
|
||||
@@ -93,43 +103,64 @@ func TestAgentNetwork_UpdateSettings_PreservesImmutableAndTogglesCollection(t *t
|
||||
APIKey: "sk-test",
|
||||
Enabled: true,
|
||||
Models: []agenttypes.ProviderModel{{ID: "gpt-5.4"}},
|
||||
}, clusterAddr)
|
||||
require.NoError(t, err, "CreateProvider must bootstrap settings")
|
||||
|
||||
before, err := mgr.GetSettings(ctx, accountID, adminUserID)
|
||||
require.NoError(t, err, "GetSettings must succeed after bootstrap")
|
||||
require.Equal(t, clusterAddr, before.Cluster, "cluster pinned at bootstrap")
|
||||
require.NotEmpty(t, before.Subdomain, "subdomain pinned at bootstrap")
|
||||
assert.False(t, before.EnablePromptCollection, "prompt collection defaults off")
|
||||
|
||||
// A cluster different from the one pinned at bootstrap must be rejected
|
||||
// outright — never silently swapped or ignored.
|
||||
_, err = mgr.UpdateSettings(ctx, adminUserID, &agenttypes.Settings{
|
||||
AccountID: accountID,
|
||||
Cluster: "attacker.cluster",
|
||||
EnableLogCollection: true,
|
||||
})
|
||||
require.Error(t, err, "UpdateSettings with a mismatched cluster must fail")
|
||||
require.NoError(t, err, "CreateProvider must succeed")
|
||||
|
||||
// Flipping the toggles works with the pinned cluster echoed back (and
|
||||
// with it omitted); the subdomain is never taken from the request.
|
||||
// Flipping the toggles works when the request echoes the assigned
|
||||
// identity. Retention is echoed too: UpdateSettings takes it verbatim, so
|
||||
// omitting it would zero the account's retention.
|
||||
updated, err := mgr.UpdateSettings(ctx, adminUserID, &agenttypes.Settings{
|
||||
AccountID: accountID,
|
||||
Cluster: clusterAddr,
|
||||
Subdomain: "evil",
|
||||
Domain: before.Domain,
|
||||
ProxyAddress: before.ProxyAddress,
|
||||
EnableLogCollection: true,
|
||||
EnablePromptCollection: true,
|
||||
RedactPii: true,
|
||||
AccessLogRetentionDays: before.AccessLogRetentionDays,
|
||||
})
|
||||
require.NoError(t, err, "UpdateSettings must succeed")
|
||||
assert.Equal(t, before.Cluster, updated.Cluster, "cluster is immutable and must be preserved")
|
||||
assert.Equal(t, before.Subdomain, updated.Subdomain, "subdomain is immutable and must be preserved")
|
||||
assert.Equal(t, before.Domain, updated.Domain, "domain is immutable and must be preserved")
|
||||
assert.Equal(t, before.ProxyAddress, updated.ProxyAddress, "proxy address is immutable and must be preserved")
|
||||
assert.True(t, updated.EnableLogCollection, "log collection toggle must apply")
|
||||
assert.True(t, updated.EnablePromptCollection, "prompt collection toggle must apply")
|
||||
assert.True(t, updated.RedactPii, "redact toggle must apply")
|
||||
assert.Equal(t, before.AccessLogRetentionDays, updated.AccessLogRetentionDays, "echoed retention must survive")
|
||||
|
||||
// Neither identity field can be smuggled into the row: a hand-rolled
|
||||
// Settings value carrying a different endpoint or proxy address is
|
||||
// rejected, not silently ignored.
|
||||
for _, tc := range []struct {
|
||||
name string
|
||||
domain string
|
||||
proxyAddress string
|
||||
}{
|
||||
{name: "foreign endpoint", domain: "evil.example.com", proxyAddress: before.ProxyAddress},
|
||||
{name: "foreign proxy address", domain: before.Domain, proxyAddress: "attacker.cluster"},
|
||||
{name: "empty identity echo", domain: "", proxyAddress: ""},
|
||||
} {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
_, err := mgr.UpdateSettings(ctx, adminUserID, &agenttypes.Settings{
|
||||
AccountID: accountID,
|
||||
Domain: tc.domain,
|
||||
ProxyAddress: tc.proxyAddress,
|
||||
EnableLogCollection: false,
|
||||
EnablePromptCollection: false,
|
||||
RedactPii: false,
|
||||
AccessLogRetentionDays: before.AccessLogRetentionDays,
|
||||
})
|
||||
assert.Error(t, err, "a mismatched identity echo must be rejected")
|
||||
assert.ErrorContains(t, err, "immutable", "the rejection must name the immutability rule")
|
||||
})
|
||||
}
|
||||
|
||||
// The rejected updates left the row exactly as the accepted one wrote it.
|
||||
afterRejects, err := mgr.GetSettings(ctx, accountID, adminUserID)
|
||||
require.NoError(t, err, "GetSettings must succeed")
|
||||
assert.True(t, afterRejects.EnablePromptCollection, "a rejected update must not roll back the accepted toggles")
|
||||
|
||||
reloaded, err := am.Store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID)
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, before.Cluster, reloaded.Cluster, "persisted cluster unchanged")
|
||||
assert.Equal(t, before.Domain, reloaded.Domain, "persisted domain unchanged")
|
||||
assert.Equal(t, before.ProxyAddress, reloaded.ProxyAddress, "persisted proxy address unchanged")
|
||||
assert.True(t, reloaded.EnablePromptCollection, "persisted prompt collection toggled on")
|
||||
}
|
||||
|
||||
@@ -92,6 +92,14 @@ func TestAgentNetwork_ProviderCRUD_FansOutToProxyAndClientPeers(t *testing.T) {
|
||||
// UpdateAccountPeers, which is the path under test.
|
||||
agentMgr := agentnetwork.NewManager(am.Store, permissions.NewManager(am.Store), am, nil)
|
||||
|
||||
_, err = agentMgr.CreateSettings(ctx, adminUserID, agenttypes.DefaultSettings(accountID), clusterAddr, "")
|
||||
require.NoError(t, err, "CreateSettings must bootstrap the endpoint")
|
||||
// The bootstrap itself reconciles and queues updates on both channels;
|
||||
// drain them so the fan-out assertions below can only be satisfied by the
|
||||
// operation under test, not by this leftover.
|
||||
drain(clientCh)
|
||||
drain(proxyCh)
|
||||
|
||||
provider, err := agentMgr.CreateProvider(ctx, adminUserID, &agenttypes.Provider{
|
||||
AccountID: accountID,
|
||||
ProviderID: "openai_api",
|
||||
@@ -100,7 +108,7 @@ func TestAgentNetwork_ProviderCRUD_FansOutToProxyAndClientPeers(t *testing.T) {
|
||||
APIKey: "sk-test-key",
|
||||
Enabled: true,
|
||||
Models: []agenttypes.ProviderModel{{ID: "gpt-5.4"}},
|
||||
}, clusterAddr)
|
||||
})
|
||||
require.NoError(t, err, "CreateProvider must succeed")
|
||||
|
||||
policy, err := agentMgr.CreatePolicy(ctx, adminUserID, &agenttypes.Policy{
|
||||
|
||||
@@ -19,6 +19,7 @@ import (
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
|
||||
activitymanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/activity/manager"
|
||||
nbproxy "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/proxy"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
|
||||
@@ -221,6 +222,7 @@ func setupAuthCallbackTest(t *testing.T) *testSetup {
|
||||
)
|
||||
|
||||
proxyService.SetServiceManager(&testServiceManager{store: testStore})
|
||||
proxyService.SetActivityManager(activitymanager.NewManager(testStore))
|
||||
|
||||
handler := NewAuthCallbackHandler(proxyService, nil)
|
||||
|
||||
@@ -538,6 +540,55 @@ func TestAuthCallback_UserAllowedToLogin(t *testing.T) {
|
||||
// TestAuthCallback_UserDeniedByAccountStatus asserts that a user whose account
|
||||
// is pending approval or blocked never receives a session token from the OIDC
|
||||
// callback, and that the redirect carries a description the proxy can render.
|
||||
// TestAuthCallback_RecordsUserLogin drives the real OIDC callback and asserts
|
||||
// the login lands on the user row. That timestamp is what activity accounting
|
||||
// reads, and it is the only signal that can ever count someone who reaches
|
||||
// proxy-protected services from a browser and never opens the dashboard.
|
||||
func TestAuthCallback_RecordsUserLogin(t *testing.T) {
|
||||
setup := setupAuthCallbackTest(t)
|
||||
defer setup.cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
before, err := setup.store.GetUserByUserID(ctx, store.LockingStrengthNone, "allowedUserId")
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, before.LastLogin, "fixture user starts with no login on record")
|
||||
|
||||
setup.oidcServer.tokenSubject = "allowedUserId"
|
||||
state := createTestState(t, setup.proxyService, "https://test-proxy.example.com/dashboard")
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/reverse-proxy/callback?code=test-auth-code&state="+url.QueryEscape(state), nil)
|
||||
rec := httptest.NewRecorder()
|
||||
setup.router.ServeHTTP(rec, req)
|
||||
require.Equal(t, http.StatusFound, rec.Code)
|
||||
|
||||
after, err := setup.store.GetUserByUserID(ctx, store.LockingStrengthNone, "allowedUserId")
|
||||
require.NoError(t, err)
|
||||
require.NotNil(t, after.LastLogin, "a completed proxy SSO login must be recorded on the user")
|
||||
require.WithinDuration(t, time.Now().UTC(), after.LastLogin.UTC(), time.Minute, "login should be stamped at sign-in time")
|
||||
}
|
||||
|
||||
// TestAuthCallback_DeniedUserLoginNotRecorded keeps the write on the granted
|
||||
// path: a refused sign-in is not a login.
|
||||
func TestAuthCallback_DeniedUserLoginNotRecorded(t *testing.T) {
|
||||
setup := setupAuthCallbackTest(t)
|
||||
defer setup.cleanup()
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
setup.oidcServer.tokenSubject = "blockedUserId"
|
||||
state := createTestState(t, setup.proxyService, "https://test-proxy.example.com/dashboard")
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/reverse-proxy/callback?code=test-auth-code&state="+url.QueryEscape(state), nil)
|
||||
rec := httptest.NewRecorder()
|
||||
setup.router.ServeHTTP(rec, req)
|
||||
require.Equal(t, http.StatusFound, rec.Code)
|
||||
|
||||
after, err := setup.store.GetUserByUserID(ctx, store.LockingStrengthNone, "blockedUserId")
|
||||
require.NoError(t, err)
|
||||
require.Nil(t, after.LastLogin, "a denied user must not be recorded as having logged in")
|
||||
}
|
||||
|
||||
func TestAuthCallback_UserDeniedByAccountStatus(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
|
||||
112
management/server/migration/migration_agentnetwork.go
Normal file
112
management/server/migration/migration_agentnetwork.go
Normal file
@@ -0,0 +1,112 @@
|
||||
package migration
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"gorm.io/gorm"
|
||||
)
|
||||
|
||||
// agentNetworkSettingsMigration is a local view of the agent_network_settings
|
||||
// table spanning both the legacy identity columns (cluster, subdomain) and
|
||||
// their replacement (domain, proxy_address), so the migrator can address all
|
||||
// four during the reshape without importing the current model.
|
||||
type agentNetworkSettingsMigration struct {
|
||||
AccountID string `gorm:"primaryKey"`
|
||||
Cluster string
|
||||
Subdomain string
|
||||
Domain string `gorm:"type:varchar(255)"`
|
||||
ProxyAddress string `gorm:"type:varchar(255)"`
|
||||
}
|
||||
|
||||
func (agentNetworkSettingsMigration) TableName() string { return "agent_network_settings" }
|
||||
|
||||
// MigrateAgentNetworkSettingsToDomain reshapes agent_network_settings from the
|
||||
// legacy (cluster, subdomain) identity columns to (domain, proxy_address):
|
||||
// domain becomes `<subdomain>.<cluster>` — the endpoint hostname the old
|
||||
// columns derived — and proxy_address becomes the cluster address, preserving
|
||||
// which proxy serves the account. Runs before AutoMigrate, which then creates
|
||||
// the unique index on the freshly backfilled domain column.
|
||||
//
|
||||
// A legacy row missing either half cannot be given an endpoint; the old
|
||||
// bootstrap always wrote both, so such a row indicates corruption and the
|
||||
// migration fails loudly rather than leaving an empty domain to collide with
|
||||
// the unique index confusingly.
|
||||
//
|
||||
// The transaction is real only on sqlite and postgres, where DDL is
|
||||
// transactional. MySQL implicitly commits around every ALTER TABLE, so there
|
||||
// each step stands alone; what makes an interrupted run resumable on MySQL is
|
||||
// that every step is guarded by the schema state it changes — the entry check
|
||||
// fires while either legacy column remains, the adds skip existing columns,
|
||||
// the backfill and its loud-failure check run only while the legacy cluster
|
||||
// column exists (they provably completed before any drop), and each drop
|
||||
// skips what is already gone.
|
||||
func MigrateAgentNetworkSettingsToDomain(ctx context.Context, db *gorm.DB) error {
|
||||
model := &agentNetworkSettingsMigration{}
|
||||
migrator := db.Migrator()
|
||||
|
||||
if !migrator.HasTable(model) {
|
||||
return nil
|
||||
}
|
||||
hasCluster := migrator.HasColumn(model, "cluster")
|
||||
if !hasCluster && !migrator.HasColumn(model, "subdomain") {
|
||||
// Fresh schema or already migrated — nothing to reshape.
|
||||
return nil
|
||||
}
|
||||
|
||||
return db.Transaction(func(tx *gorm.DB) error {
|
||||
txMigrator := tx.Migrator()
|
||||
for _, field := range []string{"Domain", "ProxyAddress"} {
|
||||
if !txMigrator.HasColumn(model, field) {
|
||||
if err := txMigrator.AddColumn(model, field); err != nil {
|
||||
return fmt.Errorf("add %s column to agent_network_settings: %w", field, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
if hasCluster {
|
||||
concat := "subdomain || '.' || cluster"
|
||||
if tx.Name() == "mysql" {
|
||||
concat = "CONCAT(subdomain, '.', cluster)"
|
||||
}
|
||||
res := tx.Exec(fmt.Sprintf(
|
||||
"UPDATE agent_network_settings SET domain = %s, proxy_address = cluster WHERE (domain IS NULL OR domain = '') AND cluster <> '' AND subdomain <> ''",
|
||||
concat,
|
||||
))
|
||||
if res.Error != nil {
|
||||
return fmt.Errorf("backfill agent_network_settings domain: %w", res.Error)
|
||||
}
|
||||
|
||||
var unmigratable int64
|
||||
if err := tx.Model(model).Where("domain IS NULL OR domain = ''").Count(&unmigratable).Error; err != nil {
|
||||
return fmt.Errorf("count unmigratable agent_network_settings rows: %w", err)
|
||||
}
|
||||
if unmigratable > 0 {
|
||||
return fmt.Errorf(
|
||||
"%d agent_network_settings row(s) have no cluster/subdomain to derive an endpoint from; resolve them manually before upgrading",
|
||||
unmigratable,
|
||||
)
|
||||
}
|
||||
|
||||
if res.RowsAffected > 0 {
|
||||
log.WithContext(ctx).Infof("migrated %d agent_network_settings row(s) to domain/proxy_address", res.RowsAffected)
|
||||
}
|
||||
}
|
||||
|
||||
if txMigrator.HasIndex(model, "idx_agent_network_settings_cluster_subdomain") {
|
||||
if err := txMigrator.DropIndex(model, "idx_agent_network_settings_cluster_subdomain"); err != nil {
|
||||
return fmt.Errorf("drop legacy agent_network_settings index: %w", err)
|
||||
}
|
||||
}
|
||||
for _, field := range []string{"Cluster", "Subdomain"} {
|
||||
if txMigrator.HasColumn(model, field) {
|
||||
if err := txMigrator.DropColumn(model, field); err != nil {
|
||||
return fmt.Errorf("drop legacy agent_network_settings column %s: %w", field, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
})
|
||||
}
|
||||
@@ -736,3 +736,125 @@ func TestFoldCostAggregatesIntoBuckets_SkipsAlreadyMigrated(t *testing.T) {
|
||||
assert.InDelta(t, 0.004, row.OutputCostUSD, 1e-9)
|
||||
assert.InDelta(t, 0.01, row.TotalCostUSD(), 1e-9, "derived total sums the four buckets")
|
||||
}
|
||||
|
||||
// legacyAgentNetworkSettings is the pre-reshape schema: identity carried as
|
||||
// (cluster, subdomain) instead of (domain, proxy_address).
|
||||
type legacyAgentNetworkSettings struct {
|
||||
AccountID string `gorm:"primaryKey"`
|
||||
Cluster string
|
||||
Subdomain string
|
||||
EnableLogCollection bool
|
||||
}
|
||||
|
||||
func (legacyAgentNetworkSettings) TableName() string { return "agent_network_settings" }
|
||||
|
||||
// TestMigrateAgentNetworkSettingsToDomain_BackfillsAndDropsLegacyColumns pins
|
||||
// the reshape: domain becomes `<subdomain>.<cluster>`, proxy_address becomes
|
||||
// the cluster, the legacy columns are dropped, and non-identity fields ride
|
||||
// through untouched.
|
||||
func TestMigrateAgentNetworkSettingsToDomain_BackfillsAndDropsLegacyColumns(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
db := setupDatabase(t)
|
||||
require.NoError(t, db.Migrator().DropTable(&legacyAgentNetworkSettings{}))
|
||||
require.NoError(t, db.AutoMigrate(&legacyAgentNetworkSettings{}))
|
||||
require.NoError(t, db.Create(&legacyAgentNetworkSettings{
|
||||
AccountID: "acct-1", Cluster: "eu.proxy.netbird.io", Subdomain: "violet", EnableLogCollection: true,
|
||||
}).Error)
|
||||
require.NoError(t, db.Create(&legacyAgentNetworkSettings{
|
||||
AccountID: "acct-2", Cluster: "us.proxy.netbird.io", Subdomain: "violet",
|
||||
}).Error)
|
||||
|
||||
require.NoError(t, migration.MigrateAgentNetworkSettingsToDomain(ctx, db))
|
||||
require.NoError(t, db.AutoMigrate(&agentNetworkTypes.Settings{}),
|
||||
"AutoMigrate must create the domain unique index over the backfilled values")
|
||||
|
||||
var one, two agentNetworkTypes.Settings
|
||||
require.NoError(t, db.First(&one, "account_id = ?", "acct-1").Error)
|
||||
assert.Equal(t, "violet.eu.proxy.netbird.io", one.Domain, "domain must combine subdomain and cluster")
|
||||
assert.Equal(t, "eu.proxy.netbird.io", one.ProxyAddress, "proxy address must carry the cluster")
|
||||
assert.True(t, one.EnableLogCollection, "non-identity fields must ride through")
|
||||
require.NoError(t, db.First(&two, "account_id = ?", "acct-2").Error)
|
||||
assert.Equal(t, "violet.us.proxy.netbird.io", two.Domain,
|
||||
"duplicate labels on different clusters are distinct hostnames and must both survive")
|
||||
|
||||
migrator := db.Migrator()
|
||||
assert.False(t, migrator.HasColumn(&legacyAgentNetworkSettings{}, "cluster"), "legacy cluster column must be dropped")
|
||||
assert.False(t, migrator.HasColumn(&legacyAgentNetworkSettings{}, "subdomain"), "legacy subdomain column must be dropped")
|
||||
}
|
||||
|
||||
// TestMigrateAgentNetworkSettingsToDomain_SkipsAlreadyMigrated proves the
|
||||
// migration is safe to re-run: with no legacy column present it is a no-op.
|
||||
func TestMigrateAgentNetworkSettingsToDomain_SkipsAlreadyMigrated(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
db := setupDatabase(t)
|
||||
require.NoError(t, db.Migrator().DropTable(&agentNetworkTypes.Settings{}))
|
||||
require.NoError(t, db.AutoMigrate(&agentNetworkTypes.Settings{}))
|
||||
require.NoError(t, db.Create(&agentNetworkTypes.Settings{
|
||||
AccountID: "acct-1", Domain: "gw.example.com", ProxyAddress: "gw.example.com",
|
||||
}).Error)
|
||||
|
||||
require.NoError(t, migration.MigrateAgentNetworkSettingsToDomain(ctx, db),
|
||||
"running against an already-migrated table must be a no-op, not an error")
|
||||
|
||||
var row agentNetworkTypes.Settings
|
||||
require.NoError(t, db.First(&row, "account_id = ?", "acct-1").Error)
|
||||
assert.Equal(t, "gw.example.com", row.Domain, "migrated rows must be untouched")
|
||||
}
|
||||
|
||||
// TestMigrateAgentNetworkSettingsToDomain_FailsOnUnmigratableRow pins the
|
||||
// loud-failure contract: a legacy row missing its identity halves cannot be
|
||||
// given an endpoint, and silently leaving an empty domain would collide with
|
||||
// the unique index confusingly later.
|
||||
func TestMigrateAgentNetworkSettingsToDomain_FailsOnUnmigratableRow(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
db := setupDatabase(t)
|
||||
require.NoError(t, db.Migrator().DropTable(&legacyAgentNetworkSettings{}))
|
||||
require.NoError(t, db.AutoMigrate(&legacyAgentNetworkSettings{}))
|
||||
require.NoError(t, db.Create(&legacyAgentNetworkSettings{
|
||||
AccountID: "acct-broken", Cluster: "", Subdomain: "",
|
||||
}).Error)
|
||||
|
||||
err := migration.MigrateAgentNetworkSettingsToDomain(ctx, db)
|
||||
require.Error(t, err, "a row with no identity to derive an endpoint from must fail the migration")
|
||||
assert.Contains(t, err.Error(), "resolve them manually", "the error must tell the operator what to do")
|
||||
}
|
||||
|
||||
// partialAgentNetworkSettings models the one non-atomic state a MySQL run can
|
||||
// be interrupted in: DDL auto-commits there, so a crash between the two legacy
|
||||
// column drops leaves subdomain behind while cluster (and the completed
|
||||
// backfill) are already committed.
|
||||
type partialAgentNetworkSettings struct {
|
||||
AccountID string `gorm:"primaryKey"`
|
||||
Subdomain string
|
||||
Domain string `gorm:"type:varchar(255)"`
|
||||
ProxyAddress string `gorm:"type:varchar(255)"`
|
||||
}
|
||||
|
||||
func (partialAgentNetworkSettings) TableName() string { return "agent_network_settings" }
|
||||
|
||||
// TestMigrateAgentNetworkSettingsToDomain_ResumesAfterPartialDrop pins MySQL
|
||||
// resumability: a rerun over the interrupted state must remove the leftover
|
||||
// subdomain column without re-running the backfill (the cluster column that
|
||||
// feeds it is gone) and without touching the migrated values.
|
||||
func TestMigrateAgentNetworkSettingsToDomain_ResumesAfterPartialDrop(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
db := setupDatabase(t)
|
||||
require.NoError(t, db.Migrator().DropTable(&partialAgentNetworkSettings{}))
|
||||
require.NoError(t, db.AutoMigrate(&partialAgentNetworkSettings{}))
|
||||
require.NoError(t, db.Create(&partialAgentNetworkSettings{
|
||||
AccountID: "acct-1", Subdomain: "violet",
|
||||
Domain: "violet.eu.proxy.netbird.io", ProxyAddress: "eu.proxy.netbird.io",
|
||||
}).Error)
|
||||
|
||||
require.NoError(t, migration.MigrateAgentNetworkSettingsToDomain(ctx, db),
|
||||
"a rerun over a partially-dropped schema must resume, not error")
|
||||
|
||||
migrator := db.Migrator()
|
||||
assert.False(t, migrator.HasColumn(&partialAgentNetworkSettings{}, "subdomain"),
|
||||
"the leftover legacy column must be dropped on resume")
|
||||
|
||||
var row agentNetworkTypes.Settings
|
||||
require.NoError(t, db.First(&row, "account_id = ?", "acct-1").Error)
|
||||
assert.Equal(t, "violet.eu.proxy.netbird.io", row.Domain, "migrated values must be untouched")
|
||||
assert.Equal(t, "eu.proxy.netbird.io", row.ProxyAddress, "migrated values must be untouched")
|
||||
}
|
||||
|
||||
@@ -599,6 +599,34 @@ func (s *SqlStore) ApproveAccountPeers(ctx context.Context, accountID string) (i
|
||||
return int(result.RowsAffected), nil
|
||||
}
|
||||
|
||||
// RefreshPeerLastSeen updates only peer_status_last_seen. Every other status
|
||||
// column is left untouched: peer_status_connected and
|
||||
// peer_status_session_started_at belong to the sync stream that owns the
|
||||
// session, and a blind write here would corrupt the fencing
|
||||
// MarkPeerConnectedIfNewerSession relies on.
|
||||
//
|
||||
// LastSeen comes from the database clock for the same reason it does there: a
|
||||
// Go-side timestamp is taken before the write and can land after a connect that
|
||||
// used CURRENT_TIMESTAMP, dragging the column backwards.
|
||||
//
|
||||
// staleBefore carries the caller's throttle into the same statement, so
|
||||
// concurrent requests for one peer collapse into a single write instead of
|
||||
// each racing on its own stale read. The column is nullable — Status is an
|
||||
// embedded pointer, so a peer stored without one leaves it NULL — and NULL
|
||||
// loses every comparison, hence the explicit branch for a peer never seen.
|
||||
func (s *SqlStore) RefreshPeerLastSeen(ctx context.Context, accountID, peerID string, staleBefore time.Time) (bool, error) {
|
||||
result := s.db.WithContext(ctx).
|
||||
Model(&nbpeer.Peer{}).
|
||||
Where(accountAndIDQueryCondition, accountID, peerID).
|
||||
Where("(peer_status_last_seen IS NULL OR peer_status_last_seen < ?)", staleBefore).
|
||||
Update("peer_status_last_seen", gorm.Expr("CURRENT_TIMESTAMP"))
|
||||
if result.Error != nil {
|
||||
return false, status.Errorf(status.Internal, "refresh peer last seen: %v", result.Error)
|
||||
}
|
||||
|
||||
return result.RowsAffected > 0, nil
|
||||
}
|
||||
|
||||
// SaveUsers saves the given list of users to the database.
|
||||
func (s *SqlStore) SaveUsers(ctx context.Context, users []*types.User) error {
|
||||
if len(users) == 0 {
|
||||
@@ -6340,6 +6368,30 @@ func (s *SqlStore) CountProxiesByAccountID(ctx context.Context, accountID string
|
||||
return count, nil
|
||||
}
|
||||
|
||||
// HasActiveProxyAtClusterAddress reports whether any proxy — shared or
|
||||
// account-scoped — is currently active at the given cluster address, using
|
||||
// the same connected-within-threshold window as the other active-proxy
|
||||
// queries. Backs the agent-network settings delete guard: settings cannot be
|
||||
// deleted while a proxy declares the endpoint hostname as its address.
|
||||
//
|
||||
// The comparison folds case on both sides: the caller passes a normalized
|
||||
// (lowercase) hostname, but proxies declare their cluster address verbatim
|
||||
// and Connect stores it unchanged, so on case-sensitive collations a proxy
|
||||
// declaring "GW.Example.com" would otherwise slip past the guard. Hostnames
|
||||
// are case-insensitive per RFC 4343; the guard must be too.
|
||||
func (s *SqlStore) HasActiveProxyAtClusterAddress(ctx context.Context, clusterAddress string) (bool, error) {
|
||||
var count int64
|
||||
result := s.db.
|
||||
Model(&proxy.Proxy{}).
|
||||
Where("LOWER(cluster_address) = LOWER(?) AND status = ? AND last_seen > ?", clusterAddress, proxy.StatusConnected, time.Now().Add(-proxyActiveThreshold)).
|
||||
Count(&count)
|
||||
if result.Error != nil {
|
||||
log.WithContext(ctx).Errorf("failed to count active proxies at cluster address: %v", result.Error)
|
||||
return false, status.Errorf(status.Internal, "failed to count active proxies at cluster address")
|
||||
}
|
||||
return count > 0, nil
|
||||
}
|
||||
|
||||
func (s *SqlStore) IsClusterAddressConflicting(ctx context.Context, clusterAddress, accountID string) (bool, error) {
|
||||
var count int64
|
||||
result := s.db.
|
||||
|
||||
122
management/server/store/sql_store_activity_test.go
Normal file
122
management/server/store/sql_store_activity_test.go
Normal file
@@ -0,0 +1,122 @@
|
||||
package store
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/netip"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
nbpeer "github.com/netbirdio/netbird/management/server/peer"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
)
|
||||
|
||||
const activityAccountID = "activityAccountId"
|
||||
|
||||
func newActivityTestStore(t *testing.T) Store {
|
||||
t.Helper()
|
||||
|
||||
store, cleanUp, err := NewTestStoreFromSQL(context.Background(), "", t.TempDir())
|
||||
require.NoError(t, err)
|
||||
t.Cleanup(cleanUp)
|
||||
|
||||
require.NoError(t, store.SaveAccount(context.Background(), &types.Account{
|
||||
Id: activityAccountID,
|
||||
Domain: "activity.example.com",
|
||||
CreatedAt: time.Now().UTC(),
|
||||
}))
|
||||
|
||||
return store
|
||||
}
|
||||
|
||||
func TestRefreshPeerLastSeen(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
store := newActivityTestStore(t)
|
||||
stored := time.Now().UTC().Add(-3 * time.Hour)
|
||||
require.NoError(t, store.AddPeerToAccount(ctx, activityPeer(stored)))
|
||||
|
||||
refreshed, err := store.RefreshPeerLastSeen(ctx, activityAccountID, "activityPeer", time.Now().UTC().Add(-time.Hour))
|
||||
require.NoError(t, err)
|
||||
assert.True(t, refreshed, "a peer seen three hours ago is stale enough to refresh")
|
||||
|
||||
peer, err := store.GetPeerByID(ctx, LockingStrengthNone, activityAccountID, "activityPeer")
|
||||
require.NoError(t, err)
|
||||
assert.WithinDuration(t, time.Now().UTC(), peer.Status.LastSeen.UTC(), time.Minute, "last seen should be stamped at write time")
|
||||
assert.True(t, peer.Status.LastSeen.After(stored), "last seen must move forward")
|
||||
}
|
||||
|
||||
// TestRefreshPeerLastSeenHonoursCutoff covers the throttle the caller relies on:
|
||||
// two concurrent requests both read the same stale peer, but only the statement
|
||||
// that still finds LastSeen behind the cutoff writes.
|
||||
func TestRefreshPeerLastSeenHonoursCutoff(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
store := newActivityTestStore(t)
|
||||
stored := time.Now().UTC().Add(-10 * time.Minute)
|
||||
require.NoError(t, store.AddPeerToAccount(ctx, activityPeer(stored)))
|
||||
|
||||
refreshed, err := store.RefreshPeerLastSeen(ctx, activityAccountID, "activityPeer", time.Now().UTC().Add(-time.Hour))
|
||||
require.NoError(t, err)
|
||||
assert.False(t, refreshed, "a peer seen inside the interval must not be written")
|
||||
|
||||
peer, err := store.GetPeerByID(ctx, LockingStrengthNone, activityAccountID, "activityPeer")
|
||||
require.NoError(t, err)
|
||||
assert.WithinDuration(t, stored, peer.Status.LastSeen.UTC(), time.Second, "last seen must be left where it was")
|
||||
}
|
||||
|
||||
// TestRefreshPeerLastSeenRecordsNeverSeenPeer covers the nullable column. Status
|
||||
// is an embedded pointer, so a peer stored without one leaves last seen NULL,
|
||||
// and NULL loses the cutoff comparison — such a peer would never record its
|
||||
// first activity.
|
||||
func TestRefreshPeerLastSeenRecordsNeverSeenPeer(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
store := newActivityTestStore(t)
|
||||
stored := activityPeer(time.Time{})
|
||||
stored.Status = nil
|
||||
require.NoError(t, store.AddPeerToAccount(ctx, stored))
|
||||
|
||||
refreshed, err := store.RefreshPeerLastSeen(ctx, activityAccountID, "activityPeer", time.Now().UTC().Add(-time.Hour))
|
||||
require.NoError(t, err)
|
||||
assert.True(t, refreshed, "a peer that was never seen must record its first activity")
|
||||
|
||||
peer, err := store.GetPeerByID(ctx, LockingStrengthNone, activityAccountID, "activityPeer")
|
||||
require.NoError(t, err)
|
||||
assert.WithinDuration(t, time.Now().UTC(), peer.Status.LastSeen.UTC(), time.Minute, "last seen should be stamped at write time")
|
||||
}
|
||||
|
||||
// TestRefreshPeerLastSeenLeavesSessionStateAlone pins the column boundary: the
|
||||
// connected flag and the session token belong to the sync stream that owns the
|
||||
// peer's session, and a blind write here would corrupt its fencing. This is why
|
||||
// SavePeerStatus is not reused for an activity bump.
|
||||
func TestRefreshPeerLastSeenLeavesSessionStateAlone(t *testing.T) {
|
||||
ctx := context.Background()
|
||||
store := newActivityTestStore(t)
|
||||
|
||||
stored := activityPeer(time.Date(2026, 3, 1, 9, 0, 0, 0, time.UTC))
|
||||
stored.Status.Connected = true
|
||||
stored.Status.SessionStartedAt = 1234567890
|
||||
require.NoError(t, store.AddPeerToAccount(ctx, stored))
|
||||
|
||||
refreshed, err := store.RefreshPeerLastSeen(ctx, activityAccountID, "activityPeer", time.Now().UTC().Add(-time.Hour))
|
||||
require.NoError(t, err)
|
||||
require.True(t, refreshed, "the peer is stale enough to refresh")
|
||||
|
||||
peer, err := store.GetPeerByID(ctx, LockingStrengthNone, activityAccountID, "activityPeer")
|
||||
require.NoError(t, err)
|
||||
assert.WithinDuration(t, time.Now().UTC(), peer.Status.LastSeen.UTC(), time.Minute, "last seen should move forward")
|
||||
assert.True(t, peer.Status.Connected, "connected flag must survive an activity write")
|
||||
assert.Equal(t, int64(1234567890), peer.Status.SessionStartedAt, "session token must survive an activity write")
|
||||
}
|
||||
|
||||
func activityPeer(lastSeen time.Time) *nbpeer.Peer {
|
||||
return &nbpeer.Peer{
|
||||
ID: "activityPeer",
|
||||
AccountID: activityAccountID,
|
||||
Key: "activityPeerKey",
|
||||
IP: netip.MustParseAddr("100.64.0.9"),
|
||||
Name: "activity-peer",
|
||||
DNSLabel: "activity-peer",
|
||||
Status: &nbpeer.PeerStatus{LastSeen: lastSeen},
|
||||
}
|
||||
}
|
||||
@@ -315,25 +315,65 @@ func (s *SqlStore) GetAllAgentNetworkSettings(ctx context.Context, lockStrength
|
||||
return settings, nil
|
||||
}
|
||||
|
||||
// GetAgentNetworkSettingsByCluster returns every Settings row pinned to
|
||||
// the given proxy cluster. Used by the bootstrap label generator to
|
||||
// build the set of subdomains already taken on a cluster.
|
||||
func (s *SqlStore) GetAgentNetworkSettingsByCluster(ctx context.Context, lockStrength LockingStrength, cluster string) ([]*agentNetworkTypes.Settings, error) {
|
||||
// GetAgentNetworkSettingsByProxyAddress returns every Settings row whose
|
||||
// gateway is served by the proxy declaring the given cluster address. Used by
|
||||
// cluster-scoped synthesis to find the accounts a shared proxy serves.
|
||||
func (s *SqlStore) GetAgentNetworkSettingsByProxyAddress(ctx context.Context, lockStrength LockingStrength, proxyAddress string) ([]*agentNetworkTypes.Settings, error) {
|
||||
tx := s.db
|
||||
if lockStrength != LockingStrengthNone {
|
||||
tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)})
|
||||
}
|
||||
|
||||
var settings []*agentNetworkTypes.Settings
|
||||
result := tx.Find(&settings, "cluster = ?", cluster)
|
||||
result := tx.Find(&settings, "proxy_address = ?", proxyAddress)
|
||||
if result.Error != nil {
|
||||
log.WithContext(ctx).Errorf("failed to get agent network settings by cluster from store: %v", result.Error)
|
||||
return nil, status.Errorf(status.Internal, "failed to get agent network settings by cluster from store")
|
||||
log.WithContext(ctx).Errorf("failed to get agent network settings by proxy address from store: %v", result.Error)
|
||||
return nil, status.Errorf(status.Internal, "failed to get agent network settings by proxy address from store")
|
||||
}
|
||||
|
||||
return settings, nil
|
||||
}
|
||||
|
||||
// GetAgentNetworkSettingsByDomain resolves the single Settings row holding the
|
||||
// given endpoint hostname — a point query on the domain unique index. Returns
|
||||
// status.NotFound when no account owns the domain.
|
||||
func (s *SqlStore) GetAgentNetworkSettingsByDomain(ctx context.Context, lockStrength LockingStrength, domain string) (*agentNetworkTypes.Settings, error) {
|
||||
tx := s.db
|
||||
if lockStrength != LockingStrengthNone {
|
||||
tx = tx.Clauses(clause.Locking{Strength: string(lockStrength)})
|
||||
}
|
||||
|
||||
var settings agentNetworkTypes.Settings
|
||||
result := tx.Take(&settings, "domain = ?", domain)
|
||||
if result.Error != nil {
|
||||
if errors.Is(result.Error, gorm.ErrRecordNotFound) {
|
||||
return nil, status.Errorf(status.NotFound, "agent network settings for domain %s not found", domain)
|
||||
}
|
||||
|
||||
log.WithContext(ctx).Errorf("failed to get agent network settings by domain from store: %v", result.Error)
|
||||
return nil, status.Errorf(status.Internal, "failed to get agent network settings by domain from store")
|
||||
}
|
||||
|
||||
return &settings, nil
|
||||
}
|
||||
|
||||
// CreateAgentNetworkSettings inserts a new settings row.
|
||||
//
|
||||
// Unlike SaveAgentNetworkSettings (an upsert) this is a plain INSERT, and it
|
||||
// returns the driver error unwrapped. Both properties are required by the
|
||||
// bootstrap allocator: an upsert would overwrite whichever row it collided
|
||||
// with, and the allocator classifies the rejection by matching the driver's
|
||||
// message — a unique violation on the account primary key means a concurrent
|
||||
// bootstrap for the same account won, and one on the domain index means the
|
||||
// hostname is taken.
|
||||
func (s *SqlStore) CreateAgentNetworkSettings(ctx context.Context, settings *agentNetworkTypes.Settings) error {
|
||||
if err := s.db.Create(settings).Error; err != nil {
|
||||
log.WithContext(ctx).Debugf("failed to create agent network settings: %v", err)
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// SaveAgentNetworkSettings upserts the per-account Agent Network
|
||||
// settings row.
|
||||
func (s *SqlStore) SaveAgentNetworkSettings(ctx context.Context, settings *agentNetworkTypes.Settings) error {
|
||||
@@ -346,6 +386,25 @@ func (s *SqlStore) SaveAgentNetworkSettings(ctx context.Context, settings *agent
|
||||
return nil
|
||||
}
|
||||
|
||||
// DeleteAgentNetworkSettings removes the per-account Agent Network settings
|
||||
// row, releasing the account's endpoint. Returns status.NotFound when no row
|
||||
// exists. The guards on the delete (no providers, no proxy actively serving
|
||||
// the endpoint) live in the manager, which runs this inside a transaction
|
||||
// after re-checking them under a row lock.
|
||||
func (s *SqlStore) DeleteAgentNetworkSettings(ctx context.Context, accountID string) error {
|
||||
result := s.db.Delete(&agentNetworkTypes.Settings{}, "account_id = ?", accountID)
|
||||
if result.Error != nil {
|
||||
log.WithContext(ctx).Errorf("failed to delete agent network settings from store: %v", result.Error)
|
||||
return status.Errorf(status.Internal, "failed to delete agent network settings from store")
|
||||
}
|
||||
|
||||
if result.RowsAffected == 0 {
|
||||
return status.Errorf(status.NotFound, "agent network settings for account %s not found", accountID)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// IncrementAgentNetworkConsumption atomically upserts the consumption
|
||||
// row keyed on (account, dim_kind, dim_id, window_seconds, window_start)
|
||||
// and adds the supplied deltas. Concurrent calls from multiple proxy
|
||||
|
||||
@@ -88,9 +88,9 @@ func TestAgentNetworkSettings_RealStore_CollectionTogglesRoundTrip(t *testing.T)
|
||||
|
||||
const accountID = "acc-settings-toggles"
|
||||
require.NoError(t, s.SaveAgentNetworkSettings(ctx, &agentNetworkTypes.Settings{
|
||||
AccountID: accountID,
|
||||
Cluster: "eu.proxy.netbird.io",
|
||||
Subdomain: "violet",
|
||||
AccountID: accountID,
|
||||
Domain: "violet.eu.proxy.netbird.io",
|
||||
ProxyAddress: "eu.proxy.netbird.io",
|
||||
}))
|
||||
|
||||
got, err := s.GetAgentNetworkSettings(ctx, LockingStrengthNone, accountID)
|
||||
|
||||
@@ -180,6 +180,14 @@ type Store interface {
|
||||
// Returns true when the update happened, false when this stream lost
|
||||
// the race against a newer session.
|
||||
MarkPeerConnectedIfNewerSession(ctx context.Context, accountID, peerID string, newSessionStartedAt int64) (bool, error)
|
||||
// RefreshPeerLastSeen records that a peer was just seen, stamping the
|
||||
// database clock like the other status writers. Connected and
|
||||
// SessionStartedAt are left alone, so this never interferes with the
|
||||
// session-ownership protocol MarkPeerConnectedIfNewerSession implements.
|
||||
// The write only lands when the stored LastSeen is older than
|
||||
// staleBefore, which keeps a caller's throttle atomic under concurrent
|
||||
// requests for the same peer. Returns true when the update happened.
|
||||
RefreshPeerLastSeen(ctx context.Context, accountID, peerID string, staleBefore time.Time) (bool, error)
|
||||
// MarkPeerDisconnectedIfSameSession sets the peer to disconnected and
|
||||
// resets SessionStartedAt to zero, but only when the stored
|
||||
// SessionStartedAt equals the given sessionStartedAt. LastSeen is
|
||||
@@ -328,6 +336,7 @@ type Store interface {
|
||||
GetProxyByAccountID(ctx context.Context, accountID string) (*proxy.Proxy, error)
|
||||
CountProxiesByAccountID(ctx context.Context, accountID string) (int64, error)
|
||||
IsClusterAddressConflicting(ctx context.Context, clusterAddress, accountID string) (bool, error)
|
||||
HasActiveProxyAtClusterAddress(ctx context.Context, clusterAddress string) (bool, error)
|
||||
DeleteAccountCluster(ctx context.Context, clusterAddress, accountID string) error
|
||||
|
||||
GetCustomDomainsCounts(ctx context.Context) (total int64, validated int64, err error)
|
||||
@@ -360,8 +369,11 @@ type Store interface {
|
||||
DeleteAgentNetworkGuardrail(ctx context.Context, accountID, guardrailID string) error
|
||||
GetAgentNetworkSettings(ctx context.Context, lockStrength LockingStrength, accountID string) (*agentNetworkTypes.Settings, error)
|
||||
GetAllAgentNetworkSettings(ctx context.Context, lockStrength LockingStrength) ([]*agentNetworkTypes.Settings, error)
|
||||
GetAgentNetworkSettingsByCluster(ctx context.Context, lockStrength LockingStrength, cluster string) ([]*agentNetworkTypes.Settings, error)
|
||||
GetAgentNetworkSettingsByProxyAddress(ctx context.Context, lockStrength LockingStrength, proxyAddress string) ([]*agentNetworkTypes.Settings, error)
|
||||
GetAgentNetworkSettingsByDomain(ctx context.Context, lockStrength LockingStrength, domain string) (*agentNetworkTypes.Settings, error)
|
||||
CreateAgentNetworkSettings(ctx context.Context, settings *agentNetworkTypes.Settings) error
|
||||
SaveAgentNetworkSettings(ctx context.Context, settings *agentNetworkTypes.Settings) error
|
||||
DeleteAgentNetworkSettings(ctx context.Context, accountID string) error
|
||||
IncrementAgentNetworkConsumption(ctx context.Context, accountID string, kind agentNetworkTypes.ConsumptionDimension, dimID string, windowSeconds int64, windowStart time.Time, tokensIn, tokensOut int64, costUSD float64) error
|
||||
IncrementAgentNetworkConsumptionBatch(ctx context.Context, accountID string, keys []agentNetworkTypes.ConsumptionKey, tokensIn, tokensOut int64, costUSD float64) error
|
||||
GetAgentNetworkConsumption(ctx context.Context, lockStrength LockingStrength, accountID string, kind agentNetworkTypes.ConsumptionDimension, dimID string, windowSeconds int64, windowStart time.Time) (*agentNetworkTypes.Consumption, error)
|
||||
@@ -608,6 +620,9 @@ func getMigrationsPreAuto(ctx context.Context) []migrationFunc {
|
||||
func(db *gorm.DB) error {
|
||||
return migration.BackfillPublicIDs[posture.Checks](ctx, db)
|
||||
},
|
||||
func(db *gorm.DB) error {
|
||||
return migration.MigrateAgentNetworkSettingsToDomain(ctx, db)
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -268,6 +268,20 @@ func (mr *MockStoreMockRecorder) CreateAgentNetworkAccessLog(ctx, entry, groups
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CreateAgentNetworkAccessLog", reflect.TypeOf((*MockStore)(nil).CreateAgentNetworkAccessLog), ctx, entry, groups)
|
||||
}
|
||||
|
||||
// CreateAgentNetworkSettings mocks base method.
|
||||
func (m *MockStore) CreateAgentNetworkSettings(ctx context.Context, settings *types.Settings) error {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "CreateAgentNetworkSettings", ctx, settings)
|
||||
ret0, _ := ret[0].(error)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// CreateAgentNetworkSettings indicates an expected call of CreateAgentNetworkSettings.
|
||||
func (mr *MockStoreMockRecorder) CreateAgentNetworkSettings(ctx, settings interface{}) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "CreateAgentNetworkSettings", reflect.TypeOf((*MockStore)(nil).CreateAgentNetworkSettings), ctx, settings)
|
||||
}
|
||||
|
||||
// CreateAgentNetworkUsage mocks base method.
|
||||
func (m *MockStore) CreateAgentNetworkUsage(ctx context.Context, usage *types.AgentNetworkUsage, groups []types.AgentNetworkUsageGroup) error {
|
||||
m.ctrl.T.Helper()
|
||||
@@ -493,6 +507,20 @@ func (mr *MockStoreMockRecorder) DeleteAgentNetworkProvider(ctx, accountID, prov
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAgentNetworkProvider", reflect.TypeOf((*MockStore)(nil).DeleteAgentNetworkProvider), ctx, accountID, providerID)
|
||||
}
|
||||
|
||||
// DeleteAgentNetworkSettings mocks base method.
|
||||
func (m *MockStore) DeleteAgentNetworkSettings(ctx context.Context, accountID string) error {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "DeleteAgentNetworkSettings", ctx, accountID)
|
||||
ret0, _ := ret[0].(error)
|
||||
return ret0
|
||||
}
|
||||
|
||||
// DeleteAgentNetworkSettings indicates an expected call of DeleteAgentNetworkSettings.
|
||||
func (mr *MockStoreMockRecorder) DeleteAgentNetworkSettings(ctx, accountID interface{}) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "DeleteAgentNetworkSettings", reflect.TypeOf((*MockStore)(nil).DeleteAgentNetworkSettings), ctx, accountID)
|
||||
}
|
||||
|
||||
// DeleteCustomDomain mocks base method.
|
||||
func (m *MockStore) DeleteCustomDomain(ctx context.Context, accountID, domainID string) error {
|
||||
m.ctrl.T.Helper()
|
||||
@@ -1687,19 +1715,34 @@ func (mr *MockStoreMockRecorder) GetAgentNetworkSettings(ctx, lockStrength, acco
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAgentNetworkSettings", reflect.TypeOf((*MockStore)(nil).GetAgentNetworkSettings), ctx, lockStrength, accountID)
|
||||
}
|
||||
|
||||
// GetAgentNetworkSettingsByCluster mocks base method.
|
||||
func (m *MockStore) GetAgentNetworkSettingsByCluster(ctx context.Context, lockStrength LockingStrength, cluster string) ([]*types.Settings, error) {
|
||||
// GetAgentNetworkSettingsByDomain mocks base method.
|
||||
func (m *MockStore) GetAgentNetworkSettingsByDomain(ctx context.Context, lockStrength LockingStrength, domain string) (*types.Settings, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetAgentNetworkSettingsByCluster", ctx, lockStrength, cluster)
|
||||
ret := m.ctrl.Call(m, "GetAgentNetworkSettingsByDomain", ctx, lockStrength, domain)
|
||||
ret0, _ := ret[0].(*types.Settings)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetAgentNetworkSettingsByDomain indicates an expected call of GetAgentNetworkSettingsByDomain.
|
||||
func (mr *MockStoreMockRecorder) GetAgentNetworkSettingsByDomain(ctx, lockStrength, domain interface{}) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAgentNetworkSettingsByDomain", reflect.TypeOf((*MockStore)(nil).GetAgentNetworkSettingsByDomain), ctx, lockStrength, domain)
|
||||
}
|
||||
|
||||
// GetAgentNetworkSettingsByProxyAddress mocks base method.
|
||||
func (m *MockStore) GetAgentNetworkSettingsByProxyAddress(ctx context.Context, lockStrength LockingStrength, proxyAddress string) ([]*types.Settings, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "GetAgentNetworkSettingsByProxyAddress", ctx, lockStrength, proxyAddress)
|
||||
ret0, _ := ret[0].([]*types.Settings)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// GetAgentNetworkSettingsByCluster indicates an expected call of GetAgentNetworkSettingsByCluster.
|
||||
func (mr *MockStoreMockRecorder) GetAgentNetworkSettingsByCluster(ctx, lockStrength, cluster interface{}) *gomock.Call {
|
||||
// GetAgentNetworkSettingsByProxyAddress indicates an expected call of GetAgentNetworkSettingsByProxyAddress.
|
||||
func (mr *MockStoreMockRecorder) GetAgentNetworkSettingsByProxyAddress(ctx, lockStrength, proxyAddress interface{}) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAgentNetworkSettingsByCluster", reflect.TypeOf((*MockStore)(nil).GetAgentNetworkSettingsByCluster), ctx, lockStrength, cluster)
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetAgentNetworkSettingsByProxyAddress", reflect.TypeOf((*MockStore)(nil).GetAgentNetworkSettingsByProxyAddress), ctx, lockStrength, proxyAddress)
|
||||
}
|
||||
|
||||
// GetAgentNetworkUsageRows mocks base method.
|
||||
@@ -2956,6 +2999,21 @@ func (mr *MockStoreMockRecorder) GetZoneDNSRecordsByName(ctx, lockStrength, acco
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "GetZoneDNSRecordsByName", reflect.TypeOf((*MockStore)(nil).GetZoneDNSRecordsByName), ctx, lockStrength, accountID, zoneID, name)
|
||||
}
|
||||
|
||||
// HasActiveProxyAtClusterAddress mocks base method.
|
||||
func (m *MockStore) HasActiveProxyAtClusterAddress(ctx context.Context, clusterAddress string) (bool, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "HasActiveProxyAtClusterAddress", ctx, clusterAddress)
|
||||
ret0, _ := ret[0].(bool)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// HasActiveProxyAtClusterAddress indicates an expected call of HasActiveProxyAtClusterAddress.
|
||||
func (mr *MockStoreMockRecorder) HasActiveProxyAtClusterAddress(ctx, clusterAddress interface{}) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "HasActiveProxyAtClusterAddress", reflect.TypeOf((*MockStore)(nil).HasActiveProxyAtClusterAddress), ctx, clusterAddress)
|
||||
}
|
||||
|
||||
// IncrementAgentNetworkConsumption mocks base method.
|
||||
func (m *MockStore) IncrementAgentNetworkConsumption(ctx context.Context, accountID string, kind types.ConsumptionDimension, dimID string, windowSeconds int64, windowStart time.Time, tokensIn, tokensOut int64, costUSD float64) error {
|
||||
m.ctrl.T.Helper()
|
||||
@@ -3203,6 +3261,21 @@ func (mr *MockStoreMockRecorder) MarkProxyAccessTokenUsed(ctx, tokenID interface
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "MarkProxyAccessTokenUsed", reflect.TypeOf((*MockStore)(nil).MarkProxyAccessTokenUsed), ctx, tokenID)
|
||||
}
|
||||
|
||||
// RefreshPeerLastSeen mocks base method.
|
||||
func (m *MockStore) RefreshPeerLastSeen(ctx context.Context, accountID, peerID string, staleBefore time.Time) (bool, error) {
|
||||
m.ctrl.T.Helper()
|
||||
ret := m.ctrl.Call(m, "RefreshPeerLastSeen", ctx, accountID, peerID, staleBefore)
|
||||
ret0, _ := ret[0].(bool)
|
||||
ret1, _ := ret[1].(error)
|
||||
return ret0, ret1
|
||||
}
|
||||
|
||||
// RefreshPeerLastSeen indicates an expected call of RefreshPeerLastSeen.
|
||||
func (mr *MockStoreMockRecorder) RefreshPeerLastSeen(ctx, accountID, peerID, staleBefore interface{}) *gomock.Call {
|
||||
mr.mock.ctrl.T.Helper()
|
||||
return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "RefreshPeerLastSeen", reflect.TypeOf((*MockStore)(nil).RefreshPeerLastSeen), ctx, accountID, peerID, staleBefore)
|
||||
}
|
||||
|
||||
// RemovePeerFromAllGroups mocks base method.
|
||||
func (m *MockStore) RemovePeerFromAllGroups(ctx context.Context, peerID string) error {
|
||||
m.ctrl.T.Helper()
|
||||
|
||||
@@ -30,8 +30,6 @@ const (
|
||||
SessionJWTIssuer = "netbird-management"
|
||||
)
|
||||
|
||||
const ClientIPMetadataKey = "nb-client-ip"
|
||||
|
||||
// ResolveProto determines the protocol scheme based on the forwarded proto
|
||||
// configuration. When set to "http" or "https" the value is used directly.
|
||||
// Otherwise TLS state is used: if conn is non-nil "https" is returned, else "http".
|
||||
|
||||
@@ -1,188 +0,0 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"crypto/hmac"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"os"
|
||||
"strings"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"golang.org/x/sync/singleflight"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/internal/types"
|
||||
)
|
||||
|
||||
const headerAuthCacheTTL = 60 * time.Second
|
||||
|
||||
const envHeaderAuthCacheTTL = "NB_PROXY_HEADER_AUTH_CACHE_TTL"
|
||||
|
||||
const headerAuthCachePerService = 1024
|
||||
|
||||
const headerAuthCacheSkew = 30 * time.Second
|
||||
|
||||
const headerAuthRPCTimeout = 10 * time.Second
|
||||
|
||||
type headerCacheKey struct {
|
||||
serviceID types.ServiceID
|
||||
headerName string
|
||||
credential [sha256.Size]byte
|
||||
}
|
||||
|
||||
type headerCacheEntry struct {
|
||||
token string
|
||||
expiresAt time.Time
|
||||
}
|
||||
|
||||
type headerAuthCache struct {
|
||||
mu sync.Mutex
|
||||
entries map[types.ServiceID]*headerServiceBucket
|
||||
flight singleflight.Group
|
||||
ttl time.Duration
|
||||
maxSize int
|
||||
macKey []byte
|
||||
now func() time.Time
|
||||
}
|
||||
|
||||
type headerServiceBucket struct {
|
||||
items map[headerCacheKey]headerCacheEntry
|
||||
order []headerCacheKey
|
||||
}
|
||||
|
||||
func newHeaderAuthCache() *headerAuthCache {
|
||||
macKey := make([]byte, sha256.Size)
|
||||
_, _ = rand.Read(macKey)
|
||||
|
||||
return &headerAuthCache{
|
||||
entries: make(map[types.ServiceID]*headerServiceBucket),
|
||||
ttl: headerAuthCacheTTLFromEnv(),
|
||||
maxSize: headerAuthCachePerService,
|
||||
macKey: macKey,
|
||||
now: time.Now,
|
||||
}
|
||||
}
|
||||
|
||||
func headerAuthCacheTTLFromEnv() time.Duration {
|
||||
raw := strings.TrimSpace(os.Getenv(envHeaderAuthCacheTTL))
|
||||
if raw == "" {
|
||||
return headerAuthCacheTTL
|
||||
}
|
||||
d, err := time.ParseDuration(raw)
|
||||
if err != nil || d <= 0 {
|
||||
log.Warnf("ignoring invalid %s=%q (want a positive Go duration like 30s or 2m); using default %s",
|
||||
envHeaderAuthCacheTTL, raw, headerAuthCacheTTL)
|
||||
return headerAuthCacheTTL
|
||||
}
|
||||
return d
|
||||
}
|
||||
|
||||
func (c *headerAuthCache) key(serviceID types.ServiceID, headerName, credential string) headerCacheKey {
|
||||
mac := hmac.New(sha256.New, c.macKey)
|
||||
mac.Write([]byte(credential))
|
||||
|
||||
key := headerCacheKey{serviceID: serviceID, headerName: headerName}
|
||||
copy(key.credential[:], mac.Sum(nil))
|
||||
return key
|
||||
}
|
||||
|
||||
func (c *headerAuthCache) get(key headerCacheKey) string {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
bucket, ok := c.entries[key.serviceID]
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
entry, ok := bucket.items[key]
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
if !c.now().Before(entry.expiresAt) {
|
||||
delete(bucket.items, key)
|
||||
bucket.order = removeKey(bucket.order, key)
|
||||
return ""
|
||||
}
|
||||
return entry.token
|
||||
}
|
||||
|
||||
func (c *headerAuthCache) put(key headerCacheKey, token string, sessionExpiration time.Duration) {
|
||||
lifetime := c.ttl
|
||||
if sessionExpiration > 0 && sessionExpiration-headerAuthCacheSkew < lifetime {
|
||||
lifetime = sessionExpiration - headerAuthCacheSkew
|
||||
}
|
||||
if lifetime <= 0 {
|
||||
return
|
||||
}
|
||||
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
bucket, ok := c.entries[key.serviceID]
|
||||
if !ok {
|
||||
bucket = &headerServiceBucket{items: make(map[headerCacheKey]headerCacheEntry)}
|
||||
c.entries[key.serviceID] = bucket
|
||||
}
|
||||
if _, exists := bucket.items[key]; !exists {
|
||||
bucket.order = append(bucket.order, key)
|
||||
}
|
||||
bucket.items[key] = headerCacheEntry{token: token, expiresAt: c.now().Add(lifetime)}
|
||||
|
||||
for len(bucket.order) > c.maxSize {
|
||||
oldest := bucket.order[0]
|
||||
bucket.order = bucket.order[1:]
|
||||
delete(bucket.items, oldest)
|
||||
}
|
||||
}
|
||||
|
||||
func (c *headerAuthCache) invalidate(key headerCacheKey) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
bucket, ok := c.entries[key.serviceID]
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
delete(bucket.items, key)
|
||||
bucket.order = removeKey(bucket.order, key)
|
||||
}
|
||||
|
||||
func (c *headerAuthCache) invalidateService(serviceID types.ServiceID) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
|
||||
delete(c.entries, serviceID)
|
||||
}
|
||||
|
||||
type authenticateHeaderFn func() (string, error)
|
||||
|
||||
func (c *headerAuthCache) fetch(key headerCacheKey, sessionExpiration time.Duration, authenticate authenticateHeaderFn) (string, bool, error) {
|
||||
if token := c.get(key); token != "" {
|
||||
return token, true, nil
|
||||
}
|
||||
|
||||
res, err, _ := c.flight.Do(headerFlightKey(key), func() (any, error) {
|
||||
if token := c.get(key); token != "" {
|
||||
return token, nil
|
||||
}
|
||||
token, err := authenticate()
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
if token != "" {
|
||||
c.put(key, token, sessionExpiration)
|
||||
}
|
||||
return token, nil
|
||||
})
|
||||
if err != nil {
|
||||
return "", false, err
|
||||
}
|
||||
|
||||
token, _ := res.(string)
|
||||
return token, false, nil
|
||||
}
|
||||
|
||||
func headerFlightKey(key headerCacheKey) string {
|
||||
return string(key.serviceID) + "|" + key.headerName + "|" + string(key.credential[:])
|
||||
}
|
||||
@@ -1,244 +0,0 @@
|
||||
package auth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"sync"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/sessionkey"
|
||||
"github.com/netbirdio/netbird/proxy/auth"
|
||||
"github.com/netbirdio/netbird/proxy/internal/proxy"
|
||||
"github.com/netbirdio/netbird/shared/management/proto"
|
||||
)
|
||||
|
||||
func newCountingHeaderScheme(t *testing.T, kp *sessionkey.KeyPair, headerName, expectedValue string, calls *atomic.Int32) Header {
|
||||
t.Helper()
|
||||
token, err := sessionkey.SignToken(kp.PrivateKey, "header-user", "", "example.com", auth.MethodHeader, nil, nil, time.Hour)
|
||||
require.NoError(t, err)
|
||||
|
||||
mock := &mockAuthenticator{fn: func(_ context.Context, req *proto.AuthenticateRequest) (*proto.AuthenticateResponse, error) {
|
||||
calls.Add(1)
|
||||
ha := req.GetHeaderAuth()
|
||||
if ha != nil && ha.GetHeaderValue() == expectedValue {
|
||||
return &proto.AuthenticateResponse{Success: true, SessionToken: token}, nil
|
||||
}
|
||||
return &proto.AuthenticateResponse{Success: false}, nil
|
||||
}}
|
||||
return NewHeader(mock, "svc1", "acc1", headerName)
|
||||
}
|
||||
|
||||
func doHeaderRequest(t *testing.T, mw *Middleware, credential string) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
handler := mw.Protect(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) {
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}))
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "http://example.com/path", nil)
|
||||
req.Header.Set("X-API-Key", credential)
|
||||
req = req.WithContext(proxy.WithCapturedData(req.Context(), proxy.NewCapturedData("")))
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
handler.ServeHTTP(rec, req)
|
||||
return rec
|
||||
}
|
||||
|
||||
func TestProtect_HeaderAuth_ReusesSessionTokenAcrossRequests(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
hdr := newCountingHeaderScheme(t, kp, "X-API-Key", "secret-key", &calls)
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false))
|
||||
|
||||
for i := 0; i < 25; i++ {
|
||||
rec := doHeaderRequest(t, mw, "secret-key")
|
||||
require.Equal(t, http.StatusOK, rec.Code)
|
||||
}
|
||||
|
||||
assert.Equal(t, int32(1), calls.Load(), "a repeated credential must be verified once")
|
||||
}
|
||||
|
||||
func TestProtect_HeaderAuth_DoesNotCacheFailures(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
hdr := newCountingHeaderScheme(t, kp, "X-API-Key", "secret-key", &calls)
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false))
|
||||
|
||||
for i := 0; i < 3; i++ {
|
||||
rec := doHeaderRequest(t, mw, "wrong-key")
|
||||
require.Equal(t, http.StatusUnauthorized, rec.Code)
|
||||
}
|
||||
|
||||
assert.Equal(t, int32(3), calls.Load(), "rejected credentials must not be cached")
|
||||
}
|
||||
|
||||
func TestProtect_HeaderAuth_MissingHeaderSkipsRPC(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
hdr := newCountingHeaderScheme(t, kp, "X-API-Key", "secret-key", &calls)
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false))
|
||||
|
||||
rec := doHeaderRequest(t, mw, "")
|
||||
assert.NotEqual(t, http.StatusOK, rec.Code)
|
||||
assert.Zero(t, calls.Load(), "an absent header must not reach management")
|
||||
}
|
||||
|
||||
func TestHeaderAuthCache_EvictsExpiredEntries(t *testing.T) {
|
||||
c := newHeaderAuthCache()
|
||||
now := time.Now()
|
||||
c.now = func() time.Time { return now }
|
||||
|
||||
key := c.key("svc1", "X-API-Key", "secret")
|
||||
c.put(key, "token", time.Hour)
|
||||
require.Equal(t, "token", c.get(key))
|
||||
|
||||
now = now.Add(c.ttl + time.Second)
|
||||
assert.Empty(t, c.get(key))
|
||||
}
|
||||
|
||||
func TestHeaderAuthCache_SkipsCacheWhenSessionExpiresWithinSkew(t *testing.T) {
|
||||
c := newHeaderAuthCache()
|
||||
|
||||
key := c.key("svc1", "X-API-Key", "secret")
|
||||
c.put(key, "token", headerAuthCacheSkew)
|
||||
|
||||
assert.Empty(t, c.get(key), "a token must never outlive the session it was minted for")
|
||||
}
|
||||
|
||||
func TestHeaderAuthCache_SessionExpirationShortensTTL(t *testing.T) {
|
||||
c := newHeaderAuthCache()
|
||||
now := time.Now()
|
||||
c.now = func() time.Time { return now }
|
||||
|
||||
key := c.key("svc1", "X-API-Key", "secret")
|
||||
c.put(key, "token", headerAuthCacheSkew+10*time.Second)
|
||||
require.Equal(t, "token", c.get(key))
|
||||
|
||||
now = now.Add(11 * time.Second)
|
||||
assert.Empty(t, c.get(key))
|
||||
}
|
||||
|
||||
func TestHeaderAuthCache_BoundsEntriesPerService(t *testing.T) {
|
||||
c := newHeaderAuthCache()
|
||||
c.maxSize = 4
|
||||
|
||||
var first headerCacheKey
|
||||
for i := 0; i < 10; i++ {
|
||||
key := c.key("svc1", "X-API-Key", fmt.Sprintf("secret-%d", i))
|
||||
if i == 0 {
|
||||
first = key
|
||||
}
|
||||
c.put(key, "token", time.Hour)
|
||||
}
|
||||
|
||||
assert.Len(t, c.entries["svc1"].items, 4)
|
||||
assert.Empty(t, c.get(first), "the oldest entry must be evicted")
|
||||
}
|
||||
|
||||
func TestHeaderAuthCache_DistinguishesCredentials(t *testing.T) {
|
||||
c := newHeaderAuthCache()
|
||||
|
||||
good := c.key("svc1", "X-API-Key", "good")
|
||||
other := c.key("svc1", "X-API-Key", "other")
|
||||
c.put(good, "token", time.Hour)
|
||||
|
||||
assert.Equal(t, "token", c.get(good))
|
||||
assert.Empty(t, c.get(other))
|
||||
}
|
||||
|
||||
func TestProtect_HeaderAuth_MappingUpdateInvalidatesCache(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
hdr := newCountingHeaderScheme(t, kp, "X-API-Key", "secret-key", &calls)
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false))
|
||||
|
||||
require.Equal(t, http.StatusOK, doHeaderRequest(t, mw, "secret-key").Code)
|
||||
require.Equal(t, int32(1), calls.Load())
|
||||
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false))
|
||||
|
||||
require.Equal(t, http.StatusOK, doHeaderRequest(t, mw, "secret-key").Code)
|
||||
assert.Equal(t, int32(2), calls.Load(), "a mapping update must drop the service's cached credentials")
|
||||
}
|
||||
|
||||
func TestHeaderAuthCache_InvalidateService(t *testing.T) {
|
||||
c := newHeaderAuthCache()
|
||||
|
||||
key := c.key("svc1", "X-API-Key", "secret")
|
||||
other := c.key("svc2", "X-API-Key", "secret")
|
||||
c.put(key, "token", time.Hour)
|
||||
c.put(other, "token", time.Hour)
|
||||
|
||||
c.invalidateService("svc1")
|
||||
|
||||
assert.Empty(t, c.get(key))
|
||||
assert.Equal(t, "token", c.get(other), "other services must be untouched")
|
||||
}
|
||||
|
||||
func TestHeaderAuthCache_Invalidate(t *testing.T) {
|
||||
c := newHeaderAuthCache()
|
||||
|
||||
key := c.key("svc1", "X-API-Key", "secret")
|
||||
other := c.key("svc1", "X-API-Key", "second")
|
||||
c.put(key, "token", time.Hour)
|
||||
c.put(other, "token", time.Hour)
|
||||
|
||||
c.invalidate(key)
|
||||
|
||||
assert.Empty(t, c.get(key))
|
||||
assert.Equal(t, "token", c.get(other))
|
||||
}
|
||||
|
||||
func TestProtect_HeaderAuth_RevalidatesWhenCachedTokenRejected(t *testing.T) {
|
||||
var calls atomic.Int32
|
||||
mw := NewMiddleware(log.StandardLogger(), nil, nil)
|
||||
kp := generateTestKeyPair(t)
|
||||
hdr := newCountingHeaderScheme(t, kp, "X-API-Key", "secret-key", &calls)
|
||||
require.NoError(t, mw.AddDomain("example.com", []Scheme{hdr}, kp.PublicKey, time.Hour, "acc1", "svc1", nil, false))
|
||||
|
||||
key := mw.headerCache.key("svc1", "X-API-Key", "secret-key")
|
||||
mw.headerCache.put(key, "not-a-valid-token", time.Hour)
|
||||
|
||||
rec := doHeaderRequest(t, mw, "secret-key")
|
||||
|
||||
assert.Equal(t, http.StatusOK, rec.Code, "an unusable cached token must not fail the request")
|
||||
assert.Equal(t, int32(1), calls.Load(), "the credential must be re-verified once")
|
||||
}
|
||||
|
||||
func TestHeaderAuthCache_CollapsesConcurrentMisses(t *testing.T) {
|
||||
c := newHeaderAuthCache()
|
||||
key := c.key("svc1", "X-API-Key", "secret")
|
||||
|
||||
var calls atomic.Int32
|
||||
release := make(chan struct{})
|
||||
var wg sync.WaitGroup
|
||||
|
||||
for i := 0; i < 20; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
_, _, _ = c.fetch(key, time.Hour, func() (string, error) {
|
||||
calls.Add(1)
|
||||
<-release
|
||||
return "token", nil
|
||||
})
|
||||
}()
|
||||
}
|
||||
|
||||
time.Sleep(50 * time.Millisecond)
|
||||
close(release)
|
||||
wg.Wait()
|
||||
|
||||
assert.Equal(t, int32(1), calls.Load(), "a burst of cold requests must collapse into one RPC")
|
||||
}
|
||||
@@ -16,7 +16,6 @@ import (
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/metadata"
|
||||
|
||||
"github.com/netbirdio/netbird/proxy/auth"
|
||||
"github.com/netbirdio/netbird/proxy/internal/proxy"
|
||||
@@ -83,7 +82,6 @@ type Middleware struct {
|
||||
sessionValidator SessionValidator
|
||||
geo restrict.GeoResolver
|
||||
tunnelCache *tunnelValidationCache
|
||||
headerCache *headerAuthCache
|
||||
}
|
||||
|
||||
// NewMiddleware creates a new authentication middleware. The sessionValidator is
|
||||
@@ -98,7 +96,6 @@ func NewMiddleware(logger *log.Logger, sessionValidator SessionValidator, geo re
|
||||
sessionValidator: sessionValidator,
|
||||
geo: geo,
|
||||
tunnelCache: newTunnelValidationCache(),
|
||||
headerCache: newHeaderAuthCache(),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -455,23 +452,7 @@ func (mw *Middleware) forwardWithHeaderAuth(w http.ResponseWriter, r *http.Reque
|
||||
}
|
||||
|
||||
func (mw *Middleware) tryHeaderScheme(w http.ResponseWriter, r *http.Request, host string, config DomainConfig, hdr Header, next http.Handler) bool {
|
||||
credential := r.Header.Get(hdr.headerName)
|
||||
if credential == "" {
|
||||
return false
|
||||
}
|
||||
|
||||
key := mw.headerCache.key(hdr.id, hdr.headerName, credential)
|
||||
authenticate := func() (string, error) {
|
||||
ctx, cancel := context.WithTimeout(context.WithoutCancel(r.Context()), headerAuthRPCTimeout)
|
||||
defer cancel()
|
||||
if clientIP := mw.resolveClientIP(r); clientIP.IsValid() {
|
||||
ctx = metadata.AppendToOutgoingContext(ctx, auth.ClientIPMetadataKey, clientIP.String())
|
||||
}
|
||||
token, _, err := hdr.Authenticate(r.WithContext(ctx))
|
||||
return token, err
|
||||
}
|
||||
|
||||
token, cached, err := mw.headerCache.fetch(key, config.SessionExpiration, authenticate)
|
||||
token, _, err := hdr.Authenticate(r)
|
||||
if err != nil {
|
||||
return mw.handleHeaderAuthError(w, r, err)
|
||||
}
|
||||
@@ -480,17 +461,6 @@ func (mw *Middleware) tryHeaderScheme(w http.ResponseWriter, r *http.Request, ho
|
||||
}
|
||||
|
||||
result, err := mw.validateSessionToken(r.Context(), host, token, config.SessionPublicKey, auth.MethodHeader)
|
||||
if err != nil && cached {
|
||||
mw.headerCache.invalidate(key)
|
||||
if token, err = authenticate(); err != nil {
|
||||
return mw.handleHeaderAuthError(w, r, err)
|
||||
}
|
||||
if token == "" {
|
||||
return false
|
||||
}
|
||||
mw.headerCache.put(key, token, config.SessionExpiration)
|
||||
result, err = mw.validateSessionToken(r.Context(), host, token, config.SessionPublicKey, auth.MethodHeader)
|
||||
}
|
||||
if err != nil {
|
||||
setHeaderCapturedData(r.Context(), "", "", nil, nil)
|
||||
status := http.StatusBadRequest
|
||||
@@ -675,8 +645,6 @@ func wasCredentialSubmitted(r *http.Request, method auth.Method) bool {
|
||||
// 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 {
|
||||
mw.headerCache.invalidateService(serviceID)
|
||||
|
||||
if len(schemes) == 0 {
|
||||
mw.domainsMux.Lock()
|
||||
defer mw.domainsMux.Unlock()
|
||||
@@ -713,10 +681,6 @@ func (mw *Middleware) AddDomain(domain string, schemes []Scheme, publicKeyB64 st
|
||||
|
||||
// RemoveDomain unregisters authentication for the given domain.
|
||||
func (mw *Middleware) RemoveDomain(domain string) {
|
||||
if config, exists := mw.getDomainConfig(domain); exists {
|
||||
mw.headerCache.invalidateService(config.ServiceID)
|
||||
}
|
||||
|
||||
mw.domainsMux.Lock()
|
||||
defer mw.domainsMux.Unlock()
|
||||
delete(mw.domains, domain)
|
||||
|
||||
@@ -146,7 +146,7 @@ func (c *tunnelValidationCache) put(key tunnelCacheKey, resp *proto.ValidateTunn
|
||||
|
||||
// removeKey drops the first occurrence of needle from order. The cache
|
||||
// uses small slices so a linear scan is cheaper than a map+slice combo.
|
||||
func removeKey[T comparable](order []T, needle T) []T {
|
||||
func removeKey(order []tunnelCacheKey, needle tunnelCacheKey) []tunnelCacheKey {
|
||||
for i, k := range order {
|
||||
if k == needle {
|
||||
return append(order[:i], order[i+1:]...)
|
||||
|
||||
@@ -73,8 +73,8 @@ func TestReverseProxy_AgentNetworkRequest_FullChain(t *testing.T) {
|
||||
testAdminUser = "user-admin-1"
|
||||
adminGroupID = "grp-admins"
|
||||
providerID = "prov-openai-test"
|
||||
cluster = "test.proxy.local"
|
||||
subdomain = "fullchain"
|
||||
domain = "fullchain.test.proxy.local"
|
||||
proxyAddress = "test.proxy.local"
|
||||
)
|
||||
testLogger := log.New()
|
||||
testLogger.SetLevel(log.PanicLevel) // keep test output clean
|
||||
@@ -127,8 +127,8 @@ func TestReverseProxy_AgentNetworkRequest_FullChain(t *testing.T) {
|
||||
// increments on the response leg.
|
||||
require.NoError(t, st.SaveAgentNetworkSettings(ctx, &agentNetworkTypes.Settings{
|
||||
AccountID: testAccountID,
|
||||
Cluster: cluster,
|
||||
Subdomain: subdomain,
|
||||
Domain: domain,
|
||||
ProxyAddress: proxyAddress,
|
||||
EnablePromptCollection: true,
|
||||
EnableLogCollection: true,
|
||||
RedactPii: true,
|
||||
|
||||
@@ -57,10 +57,9 @@ func (a *AgentNetworkAPI) GetProvider(ctx context.Context, providerID string) (*
|
||||
return &ret, err
|
||||
}
|
||||
|
||||
// CreateProvider creates a new Agent Network provider. Set
|
||||
// request.BootstrapCluster on the account's first provider to bootstrap the
|
||||
// per-account gateway endpoint (alternatively bootstrap via UpdateSettings
|
||||
// with a cluster).
|
||||
// CreateProvider creates a new Agent Network provider. Providers have no
|
||||
// settings side effects — bootstrap the account's gateway endpoint separately
|
||||
// via CreateSettings.
|
||||
func (a *AgentNetworkAPI) CreateProvider(ctx context.Context, request api.PostApiAgentNetworkProvidersJSONRequestBody) (*api.AgentNetworkProvider, error) {
|
||||
requestBytes, err := json.Marshal(request)
|
||||
if err != nil {
|
||||
@@ -329,14 +328,13 @@ func (a *AgentNetworkAPI) DeleteBudgetRule(ctx context.Context, ruleID string) e
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetSettings gets the account's Agent Network gateway settings (cluster,
|
||||
// subdomain, endpoint, collection toggles). An account that has not been
|
||||
// bootstrapped yet — via UpdateSettings with a cluster, or by creating the
|
||||
// first provider with bootstrap_cluster set — reads as the defaults with an
|
||||
// empty Cluster, Subdomain and Endpoint. Management servers prior to that
|
||||
// contract answered 200 with a JSON null body instead; that legacy shape is
|
||||
// translated to an APIError matchable via IsNotFound rather than fabricating
|
||||
// defaults the server never stated.
|
||||
// GetSettings gets the account's Agent Network gateway settings (endpoint,
|
||||
// proxy address, collection toggles). An account that has not been
|
||||
// bootstrapped yet — via CreateSettings — reads as the defaults with an empty
|
||||
// Endpoint and ProxyAddress. Management servers prior to that contract
|
||||
// answered 200 with a JSON null body instead; that legacy shape is translated
|
||||
// to an APIError matchable via IsNotFound rather than fabricating defaults
|
||||
// the server never stated.
|
||||
func (a *AgentNetworkAPI) GetSettings(ctx context.Context) (*api.AgentNetworkSettings, error) {
|
||||
resp, err := a.c.NewRequest(ctx, "GET", "/api/agent-network/settings", nil, nil)
|
||||
if err != nil {
|
||||
@@ -359,11 +357,33 @@ func (a *AgentNetworkAPI) GetSettings(ctx context.Context) (*api.AgentNetworkSet
|
||||
return &ret, nil
|
||||
}
|
||||
|
||||
// CreateSettings bootstraps the account's Agent Network settings row,
|
||||
// assigning the immutable endpoint. Exactly one of request.ProxyAddress
|
||||
// (labeled endpoint beneath that cluster; the server allocates the label) and
|
||||
// request.Endpoint (self-addressed dedicated endpoint, claimed verbatim) must
|
||||
// be set. Returns a conflict when the account already has a settings row.
|
||||
func (a *AgentNetworkAPI) CreateSettings(ctx context.Context, request api.PostApiAgentNetworkSettingsJSONRequestBody) (*api.AgentNetworkSettings, error) {
|
||||
requestBytes, err := json.Marshal(request)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
resp, err := a.c.NewRequest(ctx, "POST", "/api/agent-network/settings", bytes.NewReader(requestBytes), nil)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if resp.Body != nil {
|
||||
defer resp.Body.Close()
|
||||
}
|
||||
ret, err := parseResponse[api.AgentNetworkSettings](resp)
|
||||
return &ret, err
|
||||
}
|
||||
|
||||
// UpdateSettings updates the account's Agent Network settings; the request
|
||||
// replaces every mutable field (collection toggles and retention). Setting
|
||||
// request.Cluster bootstraps the settings row when the account does not have
|
||||
// one yet; on a bootstrapped account it must match the assigned cluster (or
|
||||
// be nil) and any other value is rejected — the cluster is immutable.
|
||||
// carries every field, replacing the mutable ones (collection toggles and
|
||||
// retention). The endpoint and proxy address are assigned at bootstrap
|
||||
// (CreateSettings) and immutable — the request must echo them unchanged, and
|
||||
// a request carrying different values is rejected. Returns not-found until
|
||||
// the account is bootstrapped.
|
||||
func (a *AgentNetworkAPI) UpdateSettings(ctx context.Context, request api.PutApiAgentNetworkSettingsJSONRequestBody) (*api.AgentNetworkSettings, error) {
|
||||
requestBytes, err := json.Marshal(request)
|
||||
if err != nil {
|
||||
@@ -379,3 +399,19 @@ func (a *AgentNetworkAPI) UpdateSettings(ctx context.Context, request api.PutApi
|
||||
ret, err := parseResponse[api.AgentNetworkSettings](resp)
|
||||
return &ret, err
|
||||
}
|
||||
|
||||
// DeleteSettings deletes the account's Agent Network settings row, releasing
|
||||
// the endpoint. The server refuses (precondition failed) while any provider
|
||||
// exists for the account or while a proxy is actively serving the endpoint.
|
||||
// Bootstrapping again afterwards allocates a new endpoint.
|
||||
func (a *AgentNetworkAPI) DeleteSettings(ctx context.Context) error {
|
||||
resp, err := a.c.NewRequest(ctx, "DELETE", "/api/agent-network/settings", nil, nil)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if resp.Body != nil {
|
||||
defer resp.Body.Close()
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
@@ -47,9 +47,9 @@ var (
|
||||
}
|
||||
|
||||
testAgentNetworkSettings = api.AgentNetworkSettings{
|
||||
Cluster: "eu.proxy.netbird.io",
|
||||
Subdomain: "violet",
|
||||
Endpoint: "violet.eu.proxy.netbird.io",
|
||||
ProxyAddress: "eu.proxy.netbird.io",
|
||||
Dedicated: false,
|
||||
EnableLogCollection: true,
|
||||
AccessLogRetentionDays: ptr(30),
|
||||
}
|
||||
@@ -120,18 +120,15 @@ func TestAgentNetwork_CreateProvider_200(t *testing.T) {
|
||||
var req api.PostApiAgentNetworkProvidersJSONRequestBody
|
||||
require.NoError(t, json.Unmarshal(reqBytes, &req))
|
||||
assert.Equal(t, "OpenAI", req.Name)
|
||||
require.NotNil(t, req.BootstrapCluster)
|
||||
assert.Equal(t, "eu.proxy.netbird.io", *req.BootstrapCluster)
|
||||
retBytes, _ := json.Marshal(testAgentNetworkProvider)
|
||||
_, err = w.Write(retBytes)
|
||||
require.NoError(t, err)
|
||||
})
|
||||
ret, err := c.AgentNetwork.CreateProvider(context.Background(), api.PostApiAgentNetworkProvidersJSONRequestBody{
|
||||
ProviderId: "openai_api",
|
||||
Name: "OpenAI",
|
||||
UpstreamUrl: "https://api.openai.com",
|
||||
ApiKey: ptr("sk-test"),
|
||||
BootstrapCluster: ptr("eu.proxy.netbird.io"),
|
||||
ProviderId: "openai_api",
|
||||
Name: "OpenAI",
|
||||
UpstreamUrl: "https://api.openai.com",
|
||||
ApiKey: ptr("sk-test"),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, testAgentNetworkProvider, *ret)
|
||||
@@ -456,6 +453,45 @@ func TestAgentNetwork_GetSettings_LegacyNullBody(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestAgentNetwork_CreateSettings_200(t *testing.T) {
|
||||
withMockClient(func(c *rest.Client, mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) {
|
||||
assert.Equal(t, "POST", r.Method)
|
||||
reqBytes, err := io.ReadAll(r.Body)
|
||||
require.NoError(t, err)
|
||||
var req api.PostApiAgentNetworkSettingsJSONRequestBody
|
||||
require.NoError(t, json.Unmarshal(reqBytes, &req))
|
||||
require.NotNil(t, req.ProxyAddress, "proxy address must be on the wire")
|
||||
assert.Equal(t, "eu.proxy.netbird.io", *req.ProxyAddress)
|
||||
assert.Nil(t, req.Endpoint, "endpoint must stay off the wire for a labeled bootstrap")
|
||||
retBytes, _ := json.Marshal(testAgentNetworkSettings)
|
||||
_, err = w.Write(retBytes)
|
||||
require.NoError(t, err)
|
||||
})
|
||||
ret, err := c.AgentNetwork.CreateSettings(context.Background(), api.PostApiAgentNetworkSettingsJSONRequestBody{
|
||||
ProxyAddress: ptr("eu.proxy.netbird.io"),
|
||||
})
|
||||
require.NoError(t, err)
|
||||
assert.Equal(t, testAgentNetworkSettings, *ret)
|
||||
})
|
||||
}
|
||||
|
||||
func TestAgentNetwork_CreateSettings_Conflict(t *testing.T) {
|
||||
withMockClient(func(c *rest.Client, mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) {
|
||||
retBytes, _ := json.Marshal(util.ErrorResponse{Message: "agent network settings already bootstrapped for account acct1", Code: 409})
|
||||
w.WriteHeader(409)
|
||||
_, err := w.Write(retBytes)
|
||||
require.NoError(t, err)
|
||||
})
|
||||
_, err := c.AgentNetwork.CreateSettings(context.Background(), api.PostApiAgentNetworkSettingsJSONRequestBody{
|
||||
Endpoint: ptr("gw.example.com"),
|
||||
})
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "already bootstrapped")
|
||||
})
|
||||
}
|
||||
|
||||
func TestAgentNetwork_UpdateSettings_200(t *testing.T) {
|
||||
withMockClient(func(c *rest.Client, mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) {
|
||||
@@ -464,15 +500,18 @@ func TestAgentNetwork_UpdateSettings_200(t *testing.T) {
|
||||
require.NoError(t, err)
|
||||
var req api.PutApiAgentNetworkSettingsJSONRequestBody
|
||||
require.NoError(t, json.Unmarshal(reqBytes, &req))
|
||||
require.NotNil(t, req.Cluster, "bootstrap cluster must be on the wire")
|
||||
assert.Equal(t, "eu.proxy.netbird.io", *req.Cluster)
|
||||
assert.True(t, req.EnableLogCollection)
|
||||
assert.Equal(t, "brave-otter.eu.proxy.netbird.io", req.Endpoint,
|
||||
"the identity echo must be on the wire")
|
||||
assert.Equal(t, "eu.proxy.netbird.io", req.ProxyAddress,
|
||||
"the identity echo must be on the wire")
|
||||
retBytes, _ := json.Marshal(testAgentNetworkSettings)
|
||||
_, err = w.Write(retBytes)
|
||||
require.NoError(t, err)
|
||||
})
|
||||
ret, err := c.AgentNetwork.UpdateSettings(context.Background(), api.PutApiAgentNetworkSettingsJSONRequestBody{
|
||||
Cluster: ptr("eu.proxy.netbird.io"),
|
||||
Endpoint: "brave-otter.eu.proxy.netbird.io",
|
||||
ProxyAddress: "eu.proxy.netbird.io",
|
||||
EnableLogCollection: true,
|
||||
})
|
||||
require.NoError(t, err)
|
||||
@@ -483,15 +522,40 @@ func TestAgentNetwork_UpdateSettings_200(t *testing.T) {
|
||||
func TestAgentNetwork_UpdateSettings_Err(t *testing.T) {
|
||||
withMockClient(func(c *rest.Client, mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) {
|
||||
retBytes, _ := json.Marshal(util.ErrorResponse{Message: "cluster is immutable once assigned (current: eu.proxy.netbird.io)", Code: 422})
|
||||
w.WriteHeader(422)
|
||||
retBytes, _ := json.Marshal(util.ErrorResponse{Message: "agent network settings have not been bootstrapped yet; POST /api/agent-network/settings to bootstrap them", Code: 404})
|
||||
w.WriteHeader(404)
|
||||
_, err := w.Write(retBytes)
|
||||
require.NoError(t, err)
|
||||
})
|
||||
_, err := c.AgentNetwork.UpdateSettings(context.Background(), api.PutApiAgentNetworkSettingsJSONRequestBody{
|
||||
Cluster: ptr("us.proxy.netbird.io"),
|
||||
EnableLogCollection: true,
|
||||
})
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "immutable")
|
||||
assert.True(t, rest.IsNotFound(err), "an unbootstrapped account must surface as IsNotFound")
|
||||
})
|
||||
}
|
||||
|
||||
func TestAgentNetwork_DeleteSettings_200(t *testing.T) {
|
||||
withMockClient(func(c *rest.Client, mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) {
|
||||
assert.Equal(t, "DELETE", r.Method)
|
||||
_, err := w.Write([]byte("{}"))
|
||||
require.NoError(t, err)
|
||||
})
|
||||
require.NoError(t, c.AgentNetwork.DeleteSettings(context.Background()))
|
||||
})
|
||||
}
|
||||
|
||||
func TestAgentNetwork_DeleteSettings_Guarded(t *testing.T) {
|
||||
withMockClient(func(c *rest.Client, mux *http.ServeMux) {
|
||||
mux.HandleFunc("/api/agent-network/settings", func(w http.ResponseWriter, r *http.Request) {
|
||||
retBytes, _ := json.Marshal(util.ErrorResponse{Message: "agent network settings cannot be deleted while 2 provider(s) exist; delete the providers first", Code: 412})
|
||||
w.WriteHeader(412)
|
||||
_, err := w.Write(retBytes)
|
||||
require.NoError(t, err)
|
||||
})
|
||||
err := c.AgentNetwork.DeleteSettings(context.Background())
|
||||
require.Error(t, err)
|
||||
assert.Contains(t, err.Error(), "cannot be deleted")
|
||||
})
|
||||
}
|
||||
|
||||
@@ -5208,10 +5208,6 @@ components:
|
||||
type: string
|
||||
description: Full upstream URL (with scheme) that NetBird forwards traffic to.
|
||||
example: "https://api.openai.com"
|
||||
bootstrap_cluster:
|
||||
type: string
|
||||
description: Proxy cluster used to bootstrap the per-account agent-network endpoint when the first provider is created. Ignored on subsequent creates and on updates because the cluster is pinned on the account-level Settings row.
|
||||
example: "eu.proxy.netbird.io"
|
||||
api_key:
|
||||
type: string
|
||||
description: Upstream provider API key. Sealed at rest on the management server and never returned in responses. Required on create; optional on update (omit to keep the existing key).
|
||||
@@ -6193,20 +6189,20 @@ components:
|
||||
- cache_cost_usd
|
||||
AgentNetworkSettings:
|
||||
type: object
|
||||
description: Per-account Agent Network gateway settings. One row per account; cluster and subdomain are assigned at bootstrap and immutable thereafter. Before bootstrap the account reads as the default values with empty cluster, subdomain and endpoint.
|
||||
description: Per-account Agent Network gateway settings. One row per account; endpoint and proxy_address are assigned at bootstrap (POST) and immutable thereafter. Before bootstrap the account reads as the default values with empty endpoint and proxy_address.
|
||||
properties:
|
||||
cluster:
|
||||
type: string
|
||||
description: Address of the NetBird proxy cluster fronting this account's agent-network endpoint. Empty until the account is bootstrapped.
|
||||
example: "eu.proxy.netbird.io"
|
||||
subdomain:
|
||||
type: string
|
||||
description: Auto-generated DNS-safe label that prefixes the cluster to form the agent-network endpoint. Empty until the account is bootstrapped.
|
||||
example: "violet"
|
||||
endpoint:
|
||||
type: string
|
||||
description: Bare hostname agents call for this account, computed as `<subdomain>.<cluster>`. Empty until the account is bootstrapped.
|
||||
example: "violet.eu.proxy.netbird.io"
|
||||
description: Bare hostname agents call for this account. Empty until the account is bootstrapped.
|
||||
example: "brave-otter.eu.proxy.netbird.io"
|
||||
proxy_address:
|
||||
type: string
|
||||
description: Declared cluster address of the proxy serving this account's gateway. Equal to `endpoint` when a dedicated proxy serves the account; otherwise the endpoint's immediate parent (a shared cluster the endpoint hangs one label beneath). Empty until the account is bootstrapped.
|
||||
example: "eu.proxy.netbird.io"
|
||||
dedicated:
|
||||
type: boolean
|
||||
description: Whether the account's gateway is served by a proxy dedicated to it (endpoint equals proxy_address).
|
||||
example: false
|
||||
enable_log_collection:
|
||||
type: boolean
|
||||
description: Whether per-request access-log entries are collected for this account's agent-network traffic.
|
||||
@@ -6236,19 +6232,51 @@ components:
|
||||
readOnly: true
|
||||
example: "2026-04-26T10:30:00Z"
|
||||
required:
|
||||
- cluster
|
||||
- subdomain
|
||||
- endpoint
|
||||
- proxy_address
|
||||
- dedicated
|
||||
- enable_log_collection
|
||||
- enable_prompt_collection
|
||||
- redact_pii
|
||||
AgentNetworkSettingsCreateRequest:
|
||||
type: object
|
||||
description: Bootstraps the per-account Agent Network settings row, assigning the account's immutable endpoint. Exactly one of `proxy_address` and `endpoint` must be provided. `proxy_address` requests a labeled endpoint — the server allocates a label and the endpoint becomes `<label>.<proxy_address>`, served by whichever proxy declares that parent address. `endpoint` claims the given hostname itself as a self-addressed (dedicated) endpoint, served only by a proxy declaring exactly that address — the claim is legitimate before the proxy exists (address-first). Collection toggles may ride along; omitted toggles take their defaults.
|
||||
properties:
|
||||
proxy_address:
|
||||
type: string
|
||||
description: Cluster address to allocate a labeled endpoint beneath. Mutually exclusive with `endpoint`.
|
||||
example: "eu.proxy.netbird.io"
|
||||
endpoint:
|
||||
type: string
|
||||
description: Hostname to claim as the account's self-addressed (dedicated) endpoint. Mutually exclusive with `proxy_address`. Rejected when another account already holds it.
|
||||
example: "brave-otter.gateway.example.com"
|
||||
enable_log_collection:
|
||||
type: boolean
|
||||
description: Whether per-request access-log entries are collected for this account's agent-network traffic. Defaults to true.
|
||||
example: true
|
||||
enable_prompt_collection:
|
||||
type: boolean
|
||||
description: Master switch for request/response prompt capture. Defaults to false.
|
||||
example: false
|
||||
redact_pii:
|
||||
type: boolean
|
||||
description: Whether captured prompts have PII redacted. Defaults to false.
|
||||
example: false
|
||||
access_log_retention_days:
|
||||
type: integer
|
||||
description: Days to retain full access-log rows; older rows are swept. 0 or less means keep indefinitely. Defaults to 30.
|
||||
example: 30
|
||||
AgentNetworkSettingsRequest:
|
||||
type: object
|
||||
description: Account-level Agent Network settings update. The request replaces every mutable field. `cluster` additionally bootstraps the per-account settings row when the account does not have one yet; the subdomain is always server-assigned.
|
||||
description: Account-level Agent Network settings update. Every field is required, matching the PUT convention of the other endpoints. The endpoint and proxy address are assigned at bootstrap (POST) and are immutable — the request must carry them unchanged, and a request carrying different values is rejected. To change them, delete the settings (DELETE, guarded) and bootstrap again; re-creating allocates a new endpoint.
|
||||
properties:
|
||||
cluster:
|
||||
endpoint:
|
||||
type: string
|
||||
description: Address of the NetBird proxy cluster fronting this account's agent-network endpoint. When the account has no settings row yet, providing it bootstraps the row (assigning the subdomain that forms the agent endpoint). The cluster is immutable once assigned — later updates must omit it or send the assigned value; any other value is rejected.
|
||||
description: The account's gateway endpoint hostname. Immutable — must match the assigned value; a different value is rejected.
|
||||
example: "brave-otter.eu.proxy.netbird.io"
|
||||
proxy_address:
|
||||
type: string
|
||||
description: Declared cluster address of the proxy serving this account's gateway. Immutable — must match the assigned value; a different value is rejected.
|
||||
example: "eu.proxy.netbird.io"
|
||||
enable_log_collection:
|
||||
type: boolean
|
||||
@@ -6267,9 +6295,12 @@ components:
|
||||
description: Days to retain full access-log rows; older rows are swept. 0 or less means keep indefinitely.
|
||||
example: 30
|
||||
required:
|
||||
- endpoint
|
||||
- proxy_address
|
||||
- enable_log_collection
|
||||
- enable_prompt_collection
|
||||
- redact_pii
|
||||
- access_log_retention_days
|
||||
AgentNetworkBudgetRule:
|
||||
type: object
|
||||
description: Account-level budget rule. A limit-only rule bound to groups and/or users that applies across all policies as a min-wins ceiling. Empty targets means it applies to every caller.
|
||||
@@ -13694,7 +13725,7 @@ paths:
|
||||
/api/agent-network/settings:
|
||||
get:
|
||||
summary: Retrieve Agent Network settings
|
||||
description: Returns the per-account Agent Network gateway settings (cluster, subdomain, endpoint). Before the account is bootstrapped — on first provider create (`bootstrap_cluster`) or via PUT with `cluster` — the response carries the default values with empty cluster, subdomain and endpoint.
|
||||
description: Returns the per-account Agent Network gateway settings (endpoint, proxy address, collection toggles). Before the account is bootstrapped via POST, the response carries the default values with an empty endpoint and proxy address.
|
||||
tags: [ Agent Network ]
|
||||
security:
|
||||
- BearerAuth: [ ]
|
||||
@@ -13712,9 +13743,42 @@ paths:
|
||||
"$ref": "#/components/responses/forbidden"
|
||||
'500':
|
||||
"$ref": "#/components/responses/internal_error"
|
||||
post:
|
||||
summary: Bootstrap Agent Network settings
|
||||
description: Creates the per-account Agent Network settings row and allocates the account's endpoint. Exactly one of `proxy_address` (labeled endpoint under that cluster; the server allocates the label) and `endpoint` (self-addressed dedicated endpoint, claimed verbatim) must be provided. The endpoint and proxy address are immutable once assigned. Returns 409 when the account already has a settings row.
|
||||
tags: [ Agent Network ]
|
||||
security:
|
||||
- BearerAuth: [ ]
|
||||
- TokenAuth: [ ]
|
||||
requestBody:
|
||||
required: true
|
||||
description: Settings bootstrap request
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/AgentNetworkSettingsCreateRequest'
|
||||
responses:
|
||||
'200':
|
||||
description: The freshly bootstrapped Agent Network settings
|
||||
content:
|
||||
application/json:
|
||||
schema:
|
||||
$ref: '#/components/schemas/AgentNetworkSettings'
|
||||
'400':
|
||||
"$ref": "#/components/responses/bad_request"
|
||||
'401':
|
||||
"$ref": "#/components/responses/requires_authentication"
|
||||
'403':
|
||||
"$ref": "#/components/responses/forbidden"
|
||||
'409':
|
||||
"$ref": "#/components/responses/conflict"
|
||||
'422':
|
||||
"$ref": "#/components/responses/validation_failed"
|
||||
'500':
|
||||
"$ref": "#/components/responses/internal_error"
|
||||
put:
|
||||
summary: Update Agent Network settings
|
||||
description: Updates the account-level Agent Network settings; the request replaces every mutable field (collection toggles and retention). When the account has no settings row yet, providing `cluster` bootstraps it (assigning the subdomain that forms the agent endpoint); without `cluster` the request returns 404. Sending a `cluster` different from the assigned one is rejected (the cluster is immutable once assigned). The subdomain is always server-assigned and immutable.
|
||||
description: Updates the account-level Agent Network settings; the request carries every field, replacing the mutable ones (collection toggles and retention). Returns 404 when the account has no settings row yet — bootstrap it with POST first. The endpoint and proxy address are assigned at bootstrap and immutable; the request must carry them unchanged, and a request carrying different values is rejected.
|
||||
tags: [ Agent Network ]
|
||||
security:
|
||||
- BearerAuth: [ ]
|
||||
@@ -13744,6 +13808,27 @@ paths:
|
||||
"$ref": "#/components/responses/validation_failed"
|
||||
'500':
|
||||
"$ref": "#/components/responses/internal_error"
|
||||
delete:
|
||||
summary: Delete Agent Network settings
|
||||
description: Deletes the account's Agent Network settings row, releasing the endpoint. Guarded — the delete is refused with 412 while any Agent Network provider exists for the account or while a proxy is actively serving the endpoint. Bootstrapping again after a delete allocates a new endpoint; the released hostname is not reserved.
|
||||
tags: [ Agent Network ]
|
||||
security:
|
||||
- BearerAuth: [ ]
|
||||
- TokenAuth: [ ]
|
||||
responses:
|
||||
'200':
|
||||
description: Settings deleted
|
||||
'401':
|
||||
"$ref": "#/components/responses/requires_authentication"
|
||||
'403':
|
||||
"$ref": "#/components/responses/forbidden"
|
||||
'404':
|
||||
"$ref": "#/components/responses/not_found"
|
||||
'412':
|
||||
description: Delete refused — Agent Network providers still exist for the account, or a proxy is actively serving the endpoint
|
||||
content: { }
|
||||
'500':
|
||||
"$ref": "#/components/responses/internal_error"
|
||||
/api/agent-network/budget-rules:
|
||||
get:
|
||||
summary: List all Agent Network budget rules
|
||||
|
||||
@@ -2329,9 +2329,6 @@ type AgentNetworkProviderRequest struct {
|
||||
// ApiKey Upstream provider API key. Sealed at rest on the management server and never returned in responses. Required on create; optional on update (omit to keep the existing key).
|
||||
ApiKey *string `json:"api_key,omitempty"`
|
||||
|
||||
// BootstrapCluster Proxy cluster used to bootstrap the per-account agent-network endpoint when the first provider is created. Ignored on subsequent creates and on updates because the cluster is pinned on the account-level Settings row.
|
||||
BootstrapCluster *string `json:"bootstrap_cluster,omitempty"`
|
||||
|
||||
// Enabled Whether the provider is enabled. Defaults to true on create.
|
||||
Enabled *bool `json:"enabled,omitempty"`
|
||||
|
||||
@@ -2363,43 +2360,61 @@ type AgentNetworkProviderRequest struct {
|
||||
UpstreamUrl string `json:"upstream_url"`
|
||||
}
|
||||
|
||||
// AgentNetworkSettings Per-account Agent Network gateway settings. One row per account; cluster and subdomain are assigned at bootstrap and immutable thereafter. Before bootstrap the account reads as the default values with empty cluster, subdomain and endpoint.
|
||||
// AgentNetworkSettings Per-account Agent Network gateway settings. One row per account; endpoint and proxy_address are assigned at bootstrap (POST) and immutable thereafter. Before bootstrap the account reads as the default values with empty endpoint and proxy_address.
|
||||
type AgentNetworkSettings struct {
|
||||
// AccessLogRetentionDays Days to retain full access-log rows; older rows are swept. 0 or less means keep indefinitely. Usage records are retained independently.
|
||||
AccessLogRetentionDays *int `json:"access_log_retention_days,omitempty"`
|
||||
|
||||
// Cluster Address of the NetBird proxy cluster fronting this account's agent-network endpoint. Empty until the account is bootstrapped.
|
||||
Cluster string `json:"cluster"`
|
||||
|
||||
// CreatedAt Timestamp when the settings row was created. Absent until the account is bootstrapped.
|
||||
CreatedAt *time.Time `json:"created_at,omitempty"`
|
||||
|
||||
// Dedicated Whether the account's gateway is served by a proxy dedicated to it (endpoint equals proxy_address).
|
||||
Dedicated bool `json:"dedicated"`
|
||||
|
||||
// EnableLogCollection Whether per-request access-log entries are collected for this account's agent-network traffic.
|
||||
EnableLogCollection bool `json:"enable_log_collection"`
|
||||
|
||||
// EnablePromptCollection Master switch for request/response prompt capture. Capture runs only when this is on AND a policy guardrail also enables it.
|
||||
EnablePromptCollection bool `json:"enable_prompt_collection"`
|
||||
|
||||
// Endpoint Bare hostname agents call for this account, computed as `<subdomain>.<cluster>`. Empty until the account is bootstrapped.
|
||||
// Endpoint Bare hostname agents call for this account. Empty until the account is bootstrapped.
|
||||
Endpoint string `json:"endpoint"`
|
||||
|
||||
// ProxyAddress Declared cluster address of the proxy serving this account's gateway. Equal to `endpoint` when a dedicated proxy serves the account; otherwise the endpoint's immediate parent (a shared cluster the endpoint hangs one label beneath). Empty until the account is bootstrapped.
|
||||
ProxyAddress string `json:"proxy_address"`
|
||||
|
||||
// RedactPii Whether captured prompts have PII redacted. Effective redaction is the OR of this and any policy guardrail's redact setting.
|
||||
RedactPii bool `json:"redact_pii"`
|
||||
|
||||
// Subdomain Auto-generated DNS-safe label that prefixes the cluster to form the agent-network endpoint. Empty until the account is bootstrapped.
|
||||
Subdomain string `json:"subdomain"`
|
||||
|
||||
// UpdatedAt Timestamp when the settings row was last updated. Absent until the account is bootstrapped.
|
||||
UpdatedAt *time.Time `json:"updated_at,omitempty"`
|
||||
}
|
||||
|
||||
// AgentNetworkSettingsRequest Account-level Agent Network settings update. The request replaces every mutable field. `cluster` additionally bootstraps the per-account settings row when the account does not have one yet; the subdomain is always server-assigned.
|
||||
type AgentNetworkSettingsRequest struct {
|
||||
// AccessLogRetentionDays Days to retain full access-log rows; older rows are swept. 0 or less means keep indefinitely.
|
||||
// AgentNetworkSettingsCreateRequest Bootstraps the per-account Agent Network settings row, assigning the account's immutable endpoint. Exactly one of `proxy_address` and `endpoint` must be provided. `proxy_address` requests a labeled endpoint — the server allocates a label and the endpoint becomes `<label>.<proxy_address>`, served by whichever proxy declares that parent address. `endpoint` claims the given hostname itself as a self-addressed (dedicated) endpoint, served only by a proxy declaring exactly that address — the claim is legitimate before the proxy exists (address-first). Collection toggles may ride along; omitted toggles take their defaults.
|
||||
type AgentNetworkSettingsCreateRequest struct {
|
||||
// AccessLogRetentionDays Days to retain full access-log rows; older rows are swept. 0 or less means keep indefinitely. Defaults to 30.
|
||||
AccessLogRetentionDays *int `json:"access_log_retention_days,omitempty"`
|
||||
|
||||
// Cluster Address of the NetBird proxy cluster fronting this account's agent-network endpoint. When the account has no settings row yet, providing it bootstraps the row (assigning the subdomain that forms the agent endpoint). The cluster is immutable once assigned — later updates must omit it or send the assigned value; any other value is rejected.
|
||||
Cluster *string `json:"cluster,omitempty"`
|
||||
// EnableLogCollection Whether per-request access-log entries are collected for this account's agent-network traffic. Defaults to true.
|
||||
EnableLogCollection *bool `json:"enable_log_collection,omitempty"`
|
||||
|
||||
// EnablePromptCollection Master switch for request/response prompt capture. Defaults to false.
|
||||
EnablePromptCollection *bool `json:"enable_prompt_collection,omitempty"`
|
||||
|
||||
// Endpoint Hostname to claim as the account's self-addressed (dedicated) endpoint. Mutually exclusive with `proxy_address`. Rejected when another account already holds it.
|
||||
Endpoint *string `json:"endpoint,omitempty"`
|
||||
|
||||
// ProxyAddress Cluster address to allocate a labeled endpoint beneath. Mutually exclusive with `endpoint`.
|
||||
ProxyAddress *string `json:"proxy_address,omitempty"`
|
||||
|
||||
// RedactPii Whether captured prompts have PII redacted. Defaults to false.
|
||||
RedactPii *bool `json:"redact_pii,omitempty"`
|
||||
}
|
||||
|
||||
// AgentNetworkSettingsRequest Account-level Agent Network settings update. Every field is required, matching the PUT convention of the other endpoints. The endpoint and proxy address are assigned at bootstrap (POST) and are immutable — the request must carry them unchanged, and a request carrying different values is rejected. To change them, delete the settings (DELETE, guarded) and bootstrap again; re-creating allocates a new endpoint.
|
||||
type AgentNetworkSettingsRequest struct {
|
||||
// AccessLogRetentionDays Days to retain full access-log rows; older rows are swept. 0 or less means keep indefinitely.
|
||||
AccessLogRetentionDays int `json:"access_log_retention_days"`
|
||||
|
||||
// EnableLogCollection Whether per-request access-log entries are collected for this account's agent-network traffic.
|
||||
EnableLogCollection bool `json:"enable_log_collection"`
|
||||
@@ -2407,6 +2422,12 @@ type AgentNetworkSettingsRequest struct {
|
||||
// EnablePromptCollection Master switch for request/response prompt capture.
|
||||
EnablePromptCollection bool `json:"enable_prompt_collection"`
|
||||
|
||||
// Endpoint The account's gateway endpoint hostname. Immutable — must match the assigned value; a different value is rejected.
|
||||
Endpoint string `json:"endpoint"`
|
||||
|
||||
// ProxyAddress Declared cluster address of the proxy serving this account's gateway. Immutable — must match the assigned value; a different value is rejected.
|
||||
ProxyAddress string `json:"proxy_address"`
|
||||
|
||||
// RedactPii Whether captured prompts have PII redacted.
|
||||
RedactPii bool `json:"redact_pii"`
|
||||
}
|
||||
@@ -6176,6 +6197,9 @@ type PostApiAgentNetworkProvidersJSONRequestBody = AgentNetworkProviderRequest
|
||||
// PutApiAgentNetworkProvidersProviderIdJSONRequestBody defines body for PutApiAgentNetworkProvidersProviderId for application/json ContentType.
|
||||
type PutApiAgentNetworkProvidersProviderIdJSONRequestBody = AgentNetworkProviderRequest
|
||||
|
||||
// PostApiAgentNetworkSettingsJSONRequestBody defines body for PostApiAgentNetworkSettings for application/json ContentType.
|
||||
type PostApiAgentNetworkSettingsJSONRequestBody = AgentNetworkSettingsCreateRequest
|
||||
|
||||
// PutApiAgentNetworkSettingsJSONRequestBody defines body for PutApiAgentNetworkSettings for application/json ContentType.
|
||||
type PutApiAgentNetworkSettingsJSONRequestBody = AgentNetworkSettingsRequest
|
||||
|
||||
|
||||
@@ -9,7 +9,6 @@ import (
|
||||
"time"
|
||||
|
||||
"github.com/quic-go/quic-go"
|
||||
"github.com/quic-go/quic-go/logging"
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
nbnet "github.com/netbirdio/netbird/client/net"
|
||||
@@ -80,28 +79,6 @@ func (d Dialer) Dial(ctx context.Context, address, serverName string) (net.Conn,
|
||||
return conn, nil
|
||||
}
|
||||
|
||||
// connectionTracer returns a QUIC tracer that logs the DPLPMTUD result and the
|
||||
// reason a relay connection closed, so the path MTU settled on and teardown
|
||||
// cause are visible in logs. Lines carry the relay address as a structured
|
||||
// field, matching the rest of the relay client logging.
|
||||
func connectionTracer(addr string) func(context.Context, logging.Perspective, quic.ConnectionID) *logging.ConnectionTracer {
|
||||
relayLog := log.WithField("relay", addr)
|
||||
return func(context.Context, logging.Perspective, quic.ConnectionID) *logging.ConnectionTracer {
|
||||
return &logging.ConnectionTracer{
|
||||
UpdatedMTU: func(mtu logging.ByteCount, done bool) {
|
||||
if done {
|
||||
relayLog.Infof("QUIC path MTU settled at %d", mtu)
|
||||
return
|
||||
}
|
||||
relayLog.Debugf("QUIC path MTU probing at %d", mtu)
|
||||
},
|
||||
ClosedConnection: func(err error) {
|
||||
relayLog.Debugf("QUIC connection closed: %v", err)
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func prepareURL(address string) (string, error) {
|
||||
var host string
|
||||
var defaultPort string
|
||||
|
||||
145
shared/relay/client/dialer/quic/quic_test.go
Normal file
145
shared/relay/client/dialer/quic/quic_test.go
Normal file
@@ -0,0 +1,145 @@
|
||||
package quic
|
||||
|
||||
import (
|
||||
"testing"
|
||||
|
||||
"github.com/quic-go/quic-go/qlog"
|
||||
"github.com/quic-go/quic-go/qlogwriter"
|
||||
log "github.com/sirupsen/logrus"
|
||||
"github.com/sirupsen/logrus/hooks/test"
|
||||
)
|
||||
|
||||
func TestCloseReason(t *testing.T) {
|
||||
transportErr := qlog.TransportErrorCode(0x2) // CONNECTION_REFUSED
|
||||
appErr := qlog.ApplicationErrorCode(42)
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
event qlog.ConnectionClosed
|
||||
want string
|
||||
}{
|
||||
{
|
||||
// A close carrying nothing but an initiator still reads sensibly.
|
||||
name: "initiator only",
|
||||
event: qlog.ConnectionClosed{Initiator: qlog.InitiatorLocal},
|
||||
want: "closed by local",
|
||||
},
|
||||
{
|
||||
name: "transport error with trigger",
|
||||
event: qlog.ConnectionClosed{
|
||||
Initiator: qlog.InitiatorRemote,
|
||||
ConnectionError: &transportErr,
|
||||
Trigger: qlog.ConnectionCloseTriggerIdleTimeout,
|
||||
},
|
||||
want: "closed by remote, transport error: CONNECTION_REFUSED, trigger: idle_timeout",
|
||||
},
|
||||
{
|
||||
name: "application error with reason",
|
||||
event: qlog.ConnectionClosed{
|
||||
Initiator: qlog.InitiatorLocal,
|
||||
ApplicationError: &appErr,
|
||||
Reason: "bye",
|
||||
},
|
||||
want: "closed by local, application error: 42, reason: bye",
|
||||
},
|
||||
{
|
||||
// Transport and application errors are mutually exclusive in
|
||||
// practice; if both are set the transport code wins.
|
||||
name: "transport error takes precedence over application error",
|
||||
event: qlog.ConnectionClosed{
|
||||
Initiator: qlog.InitiatorLocal,
|
||||
ConnectionError: &transportErr,
|
||||
ApplicationError: &appErr,
|
||||
},
|
||||
want: "closed by local, transport error: CONNECTION_REFUSED",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
if got := closeReason(tt.event); got != tt.want {
|
||||
t.Errorf("closeReason() = %q, want %q", got, tt.want)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogSinkRecordEvent(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
event qlogwriter.Event
|
||||
wantLevel log.Level
|
||||
wantMsg string
|
||||
}{
|
||||
{
|
||||
name: "settled MTU is logged at info",
|
||||
event: qlog.MTUUpdated{Value: 1400, Done: true},
|
||||
wantLevel: log.InfoLevel,
|
||||
wantMsg: "QUIC path MTU settled at 1400",
|
||||
},
|
||||
{
|
||||
// Probing fires repeatedly during discovery, so it stays at debug.
|
||||
name: "MTU probe is logged at debug",
|
||||
event: qlog.MTUUpdated{Value: 1300, Done: false},
|
||||
wantLevel: log.DebugLevel,
|
||||
wantMsg: "QUIC path MTU probing at 1300",
|
||||
},
|
||||
{
|
||||
name: "connection closed is logged at debug",
|
||||
event: qlog.ConnectionClosed{Initiator: qlog.InitiatorRemote},
|
||||
wantLevel: log.DebugLevel,
|
||||
wantMsg: "QUIC connection closed: closed by remote",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
logger, hook := test.NewNullLogger()
|
||||
logger.SetLevel(log.DebugLevel)
|
||||
recorder := logSink{log: logger.WithField("relay", "relay.example.com:443")}
|
||||
|
||||
recorder.RecordEvent(tt.event)
|
||||
|
||||
entries := hook.AllEntries()
|
||||
if len(entries) != 1 {
|
||||
t.Fatalf("got %d log entries, want 1", len(entries))
|
||||
}
|
||||
if entries[0].Level != tt.wantLevel {
|
||||
t.Errorf("level = %v, want %v", entries[0].Level, tt.wantLevel)
|
||||
}
|
||||
if entries[0].Message != tt.wantMsg {
|
||||
t.Errorf("message = %q, want %q", entries[0].Message, tt.wantMsg)
|
||||
}
|
||||
if relay := entries[0].Data["relay"]; relay != "relay.example.com:443" {
|
||||
t.Errorf("relay field = %v, want relay.example.com:443", relay)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Events the relay client does not care about must not produce log lines.
|
||||
func TestLogSinkIgnoresUnhandledEvents(t *testing.T) {
|
||||
logger, hook := test.NewNullLogger()
|
||||
logger.SetLevel(log.DebugLevel)
|
||||
recorder := logSink{log: logger.WithField("relay", "relay.example.com:443")}
|
||||
|
||||
recorder.RecordEvent(qlog.PacketLost{})
|
||||
|
||||
if entries := hook.AllEntries(); len(entries) != 0 {
|
||||
t.Errorf("got %d log entries, want 0", len(entries))
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogSinkSupportsSchemas(t *testing.T) {
|
||||
trace := logSink{log: log.WithField("relay", "relay.example.com:443")}
|
||||
|
||||
if !trace.SupportsSchemas(qlog.EventSchema) {
|
||||
t.Errorf("SupportsSchemas(%q) = false, want true", qlog.EventSchema)
|
||||
}
|
||||
if trace.SupportsSchemas("urn:ietf:params:qlog:events:http3-12") {
|
||||
t.Error("SupportsSchemas() = true for an unrelated schema, want false")
|
||||
}
|
||||
if trace.AddProducer() == nil {
|
||||
t.Error("AddProducer() = nil, want a recorder")
|
||||
}
|
||||
}
|
||||
70
shared/relay/client/dialer/quic/tracer.go
Normal file
70
shared/relay/client/dialer/quic/tracer.go
Normal file
@@ -0,0 +1,70 @@
|
||||
package quic
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/quic-go/quic-go"
|
||||
"github.com/quic-go/quic-go/qlog"
|
||||
"github.com/quic-go/quic-go/qlogwriter"
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
// logSink implements both qlogwriter.Trace and qlogwriter.Recorder, forwarding
|
||||
// the few qlog events the relay client cares about to logrus instead of
|
||||
// writing a qlog file. It holds no mutable state and logrus entries are safe
|
||||
// to share, so one value can serve every producer on the connection.
|
||||
type logSink struct {
|
||||
log *log.Entry
|
||||
}
|
||||
|
||||
func (s logSink) AddProducer() qlogwriter.Recorder { return s }
|
||||
|
||||
func (s logSink) SupportsSchemas(schema string) bool { return schema == qlog.EventSchema }
|
||||
|
||||
func (s logSink) RecordEvent(event qlogwriter.Event) {
|
||||
switch e := event.(type) {
|
||||
case qlog.MTUUpdated:
|
||||
if e.Done {
|
||||
s.log.Infof("QUIC path MTU settled at %d", e.Value)
|
||||
return
|
||||
}
|
||||
s.log.Debugf("QUIC path MTU probing at %d", e.Value)
|
||||
case qlog.ConnectionClosed:
|
||||
s.log.Debugf("QUIC connection closed: %s", closeReason(e))
|
||||
}
|
||||
}
|
||||
|
||||
func (s logSink) Close() error { return nil }
|
||||
|
||||
// connectionTracer returns a QUIC tracer that logs the DPLPMTUD result and the
|
||||
// reason a relay connection closed, so the path MTU settled on and teardown
|
||||
// cause are visible in logs. Lines carry the relay address as a structured
|
||||
// field, matching the rest of the relay client logging.
|
||||
func connectionTracer(addr string) func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace {
|
||||
relayLog := log.WithField("relay", addr)
|
||||
return func(context.Context, bool, quic.ConnectionID) qlogwriter.Trace {
|
||||
return logSink{log: relayLog}
|
||||
}
|
||||
}
|
||||
|
||||
// closeReason renders a ConnectionClosed event as a single line. The event
|
||||
// carries the error as separate initiator, code, trigger and reason fields,
|
||||
// any of which may be unset.
|
||||
func closeReason(e qlog.ConnectionClosed) string {
|
||||
parts := []string{fmt.Sprintf("closed by %s", e.Initiator)}
|
||||
switch {
|
||||
case e.ConnectionError != nil:
|
||||
parts = append(parts, fmt.Sprintf("transport error: %s", *e.ConnectionError))
|
||||
case e.ApplicationError != nil:
|
||||
parts = append(parts, fmt.Sprintf("application error: %d", *e.ApplicationError))
|
||||
}
|
||||
if e.Trigger != "" {
|
||||
parts = append(parts, fmt.Sprintf("trigger: %s", e.Trigger))
|
||||
}
|
||||
if e.Reason != "" {
|
||||
parts = append(parts, fmt.Sprintf("reason: %s", e.Reason))
|
||||
}
|
||||
return strings.Join(parts, ", ")
|
||||
}
|
||||
Reference in New Issue
Block a user