Merge remote-tracking branch 'origin/revert/component-types' into revert/component-types

Signed-off-by: Dmitri Dolguikh <dmitri.external@netbird.io>
This commit is contained in:
Dmitri Dolguikh
2026-08-11 14:52:14 +02:00
130 changed files with 16832 additions and 2206 deletions
@@ -187,6 +187,7 @@ func (c *Controller) sendUpdateAccountPeers(ctx context.Context, accountID strin
semaphore := make(chan struct{}, 10)
c.injectAllProxyPolicies(ctx, account)
account.PrecomputePostureValidation(ctx)
dnsCache := &cache.DNSConfigCache{}
dnsDomain := c.GetDNSDomain(account.Settings)
peersCustomZone := account.GetPeersCustomZone(ctx, dnsDomain)
@@ -627,6 +628,7 @@ func (c *Controller) sendUpdateForAffectedPeers(ctx context.Context, accountID s
// network map that omitted the synth DNS zone, and the agent kept
// resolving against the stale or absent record.
c.injectAllProxyPolicies(ctx, account)
account.PrecomputePostureValidation(ctx)
dnsCache := &cache.DNSConfigCache{}
dnsDomain := c.GetDNSDomain(account.Settings)
peersCustomZone := account.GetPeersCustomZone(ctx, dnsDomain)
@@ -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
}
@@ -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")
}
@@ -33,6 +33,9 @@ func (sc *SqliteStoreConn) GetAccountSettings(ctx context.Context, accountId str
}
a, err := CollectOneRowForSqlite[networkmapdb.Account](rows)
if err != nil {
return nmdata.AccountSettingsInfo{}, err
}
settingsInfo := nmdata.AccountSettingsInfo{}
err = networkmapdb.FromSqlTypesToSharedTypes(reflect.ValueOf(&a), reflect.ValueOf(&settingsInfo))
+272 -77
View File
@@ -61,6 +61,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
@@ -99,6 +110,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
@@ -262,6 +276,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
@@ -446,8 +477,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")
@@ -467,6 +499,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
}
@@ -1579,9 +1624,62 @@ func (s *ProxyServiceServer) ValidateState(state string) (verifier, redirectURL
return verifier, redirectURL, nil
}
// Denied reasons reported to the proxy when access is refused because of the
// account status of the user behind the request.
const (
deniedReasonPendingApproval = "pending_approval"
deniedReasonUserBlocked = "user_blocked"
deniedReasonUserNotFound = "user_not_found"
)
var (
// ErrUserPendingApproval reports a user whose account still awaits approval
// by an administrator and may therefore not hold a proxy session.
ErrUserPendingApproval = errors.New("user pending approval")
// ErrUserBlocked reports a blocked user, who may not hold a proxy session.
ErrUserBlocked = errors.New("user blocked")
errUserUnresolved = errors.New("user could not be resolved")
)
// checkUserStatus reports whether the user's account status permits reverse
// proxy access, returning the denied reason for the proxy access log together
// with the sentinel error callers match on. A user awaiting approval is stored
// as both pending and blocked, so the pending state is reported first: it is
// the one an administrator can act on.
func checkUserStatus(user *types.User) (string, error) {
switch {
case user == nil:
return deniedReasonUserNotFound, errUserUnresolved
case user.PendingApproval:
return deniedReasonPendingApproval, ErrUserPendingApproval
case user.IsBlocked():
return deniedReasonUserBlocked, ErrUserBlocked
default:
return "", nil
}
}
// userStatusDeniedReason returns the denied reason for callers that report a
// decision rather than an error, and an empty string when the user may proceed.
func userStatusDeniedReason(user *types.User) string {
reason, _ := checkUserStatus(user)
return reason
}
// sameAccount reports whether a user belongs to a service's account. An empty
// identifier on either side never matches: two unset accounts must not compare
// equal into a grant.
func sameAccount(userAccountID, serviceAccountID string) bool {
return userAccountID != "" && serviceAccountID != "" && userAccountID == serviceAccountID
}
// GenerateSessionToken creates a signed session JWT for the given domain and
// user. The user's group memberships are embedded in the token so policy-aware
// middlewares on the proxy can authorise without an extra management round-trip.
// A user the store cannot resolve, or whose account is pending approval or
// blocked, gets no token at all, so the browser never receives a session cookie.
func (s *ProxyServiceServer) GenerateSessionToken(ctx context.Context, domain, userID string, method proxyauth.Method) (string, error) {
service, err := s.getServiceByDomain(ctx, domain)
if err != nil {
@@ -1592,25 +1690,37 @@ func (s *ProxyServiceServer) GenerateSessionToken(ctx context.Context, domain, u
return "", fmt.Errorf("no session key configured for domain: %s", domain)
}
var (
email string
groupIDs []string
groupNames []string
)
if s.usersManager != nil {
user, userGroups, uerr := s.usersManager.GetUserWithGroups(ctx, userID)
if uerr != nil {
log.WithContext(ctx).Debugf("session token mint: lookup user %s: %v", userID, uerr)
} else if user != nil {
email = user.Email
groupIDs, groupNames = pairGroupIDsAndNames(userGroups)
}
if s.usersManager == nil {
return "", errors.New("users manager not configured")
}
user, userGroups, err := s.usersManager.GetUserWithGroups(ctx, userID)
if err != nil {
return "", fmt.Errorf("get user %s: %w", userID, err)
}
if user == nil {
return "", fmt.Errorf("get user %s: %w", userID, errUserUnresolved)
}
// Bind the OIDC identity to the service's account before signing anything
// with that service's session key. The proxy validates an installed cookie
// locally against the service public key, so a token minted for a user of
// another account would be honoured without a management round-trip.
if !sameAccount(user.AccountID, service.AccountID) {
return "", fmt.Errorf("user %s does not belong to the service account", userID)
}
if _, err := checkUserStatus(user); err != nil {
return "", fmt.Errorf("session token for user %s: %w", userID, err)
}
groupIDs, groupNames := pairGroupIDsAndNames(userGroups)
return sessionkey.SignToken(
service.SessionPrivateKey,
userID,
email,
user.Email,
domain,
method,
groupIDs,
@@ -1628,6 +1738,10 @@ func (s *ProxyServiceServer) ValidateUserGroupAccess(ctx context.Context, domain
return fmt.Errorf("user not found: %s", userID)
}
if _, err := checkUserStatus(user); err != nil {
return fmt.Errorf("user %s denied access to domain %s: %w", userID, domain, err)
}
service, err := s.getAccountServiceByDomain(ctx, user.AccountID, domain)
if err != nil {
return err
@@ -1682,10 +1796,7 @@ func (s *ProxyServiceServer) ValidateSession(ctx context.Context, req *proto.Val
sessionToken := req.GetSessionToken()
if domain == "" || sessionToken == "" {
return &proto.ValidateSessionResponse{
Valid: false,
DeniedReason: "missing domain or session_token",
}, nil
return deniedSessionResponse("missing domain or session_token"), nil
}
service, err := s.getServiceByDomain(ctx, domain)
@@ -1695,83 +1806,49 @@ func (s *ProxyServiceServer) ValidateSession(ctx context.Context, req *proto.Val
"error": err.Error(),
}).Debug("ValidateSession: service not found")
//nolint:nilerr
return &proto.ValidateSessionResponse{
Valid: false,
DeniedReason: "service_not_found",
}, nil
return deniedSessionResponse("service_not_found"), nil
}
if err := enforceAccountScope(ctx, service.AccountID); err != nil {
return nil, err
}
pubKeyBytes, err := base64.StdEncoding.DecodeString(service.SessionPublicKey)
if err != nil {
log.WithFields(log.Fields{
"domain": domain,
"error": err.Error(),
}).Error("ValidateSession: decode public key")
//nolint:nilerr
return &proto.ValidateSessionResponse{
Valid: false,
DeniedReason: "invalid_service_config",
}, nil
}
userID, _, _, _, _, err := proxyauth.ValidateSessionJWT(sessionToken, domain, pubKeyBytes)
if err != nil {
log.WithFields(log.Fields{
"domain": domain,
"error": err.Error(),
}).Debug("ValidateSession: invalid session token")
//nolint:nilerr
return &proto.ValidateSessionResponse{
Valid: false,
DeniedReason: "invalid_token",
}, nil
userID, reason := sessionTokenSubject(domain, service, sessionToken)
if reason != "" {
return deniedSessionResponse(reason), nil
}
user, userGroups, err := s.usersManager.GetUserWithGroups(ctx, userID)
if err != nil {
if err != nil || user == nil {
log.WithFields(log.Fields{
"domain": domain,
"user_id": userID,
"error": err.Error(),
"error": err,
}).Debug("ValidateSession: user not found")
//nolint:nilerr
return &proto.ValidateSessionResponse{
Valid: false,
DeniedReason: "user_not_found",
}, nil
return deniedSessionResponse(deniedReasonUserNotFound), nil
}
if user.AccountID != service.AccountID {
// A user from another account gets a bare response: none of their identity
// belongs in an answer to a proxy serving a different account.
if !sameAccount(user.AccountID, service.AccountID) {
log.WithFields(log.Fields{
"domain": domain,
"user_id": userID,
"user_account": user.AccountID,
"service_account": service.AccountID,
}).Debug("ValidateSession: user account mismatch")
//nolint:nilerr
return &proto.ValidateSessionResponse{
Valid: false,
DeniedReason: "account_mismatch",
}, nil
return deniedSessionResponse("account_mismatch"), nil
}
if err := s.checkGroupAccess(service, user); err != nil {
log.WithFields(log.Fields{
"domain": domain,
"user_id": userID,
"error": err.Error(),
}).Debug("ValidateSession: access denied")
groupIDs, groupNames := pairGroupIDsAndNames(userGroups)
//nolint:nilerr
groupIDs, groupNames := pairGroupIDsAndNames(userGroups)
if reason := s.accountUserDeniedReason(domain, service, user); reason != "" {
return &proto.ValidateSessionResponse{
Valid: false,
UserId: user.Id,
UserEmail: user.Email,
DeniedReason: "not_in_group",
DeniedReason: reason,
PeerGroupIds: groupIDs,
PeerGroupNames: groupNames,
}, nil
@@ -1783,7 +1860,6 @@ func (s *ProxyServiceServer) ValidateSession(ctx context.Context, req *proto.Val
"email": user.Email,
}).Debug("ValidateSession: access granted")
groupIDs, groupNames := pairGroupIDsAndNames(userGroups)
return &proto.ValidateSessionResponse{
Valid: true,
UserId: user.Id,
@@ -1793,6 +1869,66 @@ func (s *ProxyServiceServer) ValidateSession(ctx context.Context, req *proto.Val
}, nil
}
// deniedSessionResponse builds a denial that carries no identity, for the
// checks that run before a user of this service's account is resolved.
func deniedSessionResponse(reason string) *proto.ValidateSessionResponse {
return &proto.ValidateSessionResponse{
Valid: false,
DeniedReason: reason,
}
}
// sessionTokenSubject verifies the session token against the service's session
// key and returns the user it was minted for, or the reason it cannot be
// trusted.
func sessionTokenSubject(domain string, service *rpservice.Service, sessionToken string) (userID, deniedReason string) {
pubKeyBytes, err := base64.StdEncoding.DecodeString(service.SessionPublicKey)
if err != nil {
log.WithFields(log.Fields{
"domain": domain,
"error": err.Error(),
}).Error("ValidateSession: decode public key")
return "", "invalid_service_config"
}
userID, _, _, _, _, err = proxyauth.ValidateSessionJWT(sessionToken, domain, pubKeyBytes)
if err != nil {
log.WithFields(log.Fields{
"domain": domain,
"error": err.Error(),
}).Debug("ValidateSession: invalid session token")
return "", "invalid_token"
}
return userID, ""
}
// accountUserDeniedReason gates a user of the service's own account, returning
// an empty string when access is granted. Account status comes before group
// membership: a user awaiting approval or blocked has no access regardless of
// the groups they were auto-assigned.
func (s *ProxyServiceServer) accountUserDeniedReason(domain string, service *rpservice.Service, user *types.User) string {
if reason := userStatusDeniedReason(user); reason != "" {
log.WithFields(log.Fields{
"domain": domain,
"user_id": user.Id,
"reason": reason,
}).Debug("ValidateSession: user status denies access")
return reason
}
if err := s.checkGroupAccess(service, user); err != nil {
log.WithFields(log.Fields{
"domain": domain,
"user_id": user.Id,
"error": err.Error(),
}).Debug("ValidateSession: access denied")
return "not_in_group"
}
return ""
}
func (s *ProxyServiceServer) getServiceByDomain(ctx context.Context, domain string) (*rpservice.Service, error) {
service, err := s.serviceManager.GetServiceByDomain(ctx, domain)
if err == nil {
@@ -1907,7 +2043,20 @@ func (s *ProxyServiceServer) ValidateTunnelPeer(ctx context.Context, req *proto.
}
groupIDs, groupNames := pairGroupIDsAndNames(peerGroups)
principalID, displayIdentity := s.getTunnelPeerInfo(ctx, domain, service, peer)
owner := s.resolvePeerOwner(ctx, peer, service.AccountID)
principalID, displayIdentity := s.getTunnelPeerInfo(ctx, domain, service, peer, owner)
if reason := peerOwnerDeniedReason(peer, owner); reason != "" {
log.WithFields(log.Fields{"domain": domain, "peer_id": peer.ID, "user_id": peer.UserID, "reason": reason}).Debug("ValidateTunnelPeer: owner status denies access")
return &proto.ValidateTunnelPeerResponse{
Valid: false,
UserId: principalID,
UserEmail: displayIdentity,
DeniedReason: reason,
PeerGroupIds: groupIDs,
PeerGroupNames: groupNames,
}, nil
}
if err := checkPeerGroupAccess(service, groupIDs); err != nil {
log.WithFields(log.Fields{"domain": domain, "peer_id": peer.ID, "error": err.Error()}).Debug("ValidateTunnelPeer: access denied")
@@ -1944,9 +2093,55 @@ func (s *ProxyServiceServer) ValidateTunnelPeer(ctx context.Context, req *proto.
}, nil
}
// 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
// nil rather than an error: both callers treat an unresolved owner the same
// way, and neither may trust one it could not read.
func (s *ProxyServiceServer) resolvePeerOwner(ctx context.Context, peer *peer.Peer, accountID string) *types.User {
if peer.UserID == "" {
return nil
}
user, err := s.usersManager.GetUser(ctx, peer.UserID)
if err != nil {
log.WithContext(ctx).Debugf("ValidateTunnelPeer: look up owner %s of peer %s: %v", peer.UserID, peer.ID, err)
return nil
}
// The lookup is by user ID alone, so a peer row pointing outside the
// service's account would otherwise resolve a foreign user. Leave the owner
// unresolved instead: the gate denies it, and neither the response nor the
// minted token carries an identity from another account.
if !sameAccount(user.AccountID, accountID) {
log.WithContext(ctx).Debugf("ValidateTunnelPeer: owner %s of peer %s belongs to another account", peer.UserID, peer.ID)
return nil
}
return user
}
// peerOwnerDeniedReason gates the mesh fast-path on the account status of the
// peer's owning user, so a user blocked after registering a peer loses
// mesh-origin access too. Unlinked peers (machine agents) have no owner to gate
// on and stay first-class callers. An owner the store cannot resolve denies:
// an unavailable lookup must not grant access.
func peerOwnerDeniedReason(peer *peer.Peer, owner *types.User) string {
if peer.UserID == "" {
return ""
}
if owner == nil {
return deniedReasonUserNotFound
}
return userStatusDeniedReason(owner)
}
// getTunnelPeerInfo returns the principal ID and display name for a peer, e.g. a
// user or peer ID, and peer name or user email.
func (s *ProxyServiceServer) getTunnelPeerInfo(ctx context.Context, domain string, service *rpservice.Service, peer *peer.Peer) (string, string) {
// user or peer ID, and peer name or user email. owner is the already-resolved
// user the peer is linked to, or nil.
func (s *ProxyServiceServer) getTunnelPeerInfo(ctx context.Context, domain string, service *rpservice.Service, peer *peer.Peer, owner *types.User) (string, string) {
// Resolve the principal: when the peer is linked to a user, the human is the
// principal so multiple peers owned by the same user share a single
// identity. Unlinked peers (machine agents) are their own principal keyed on
@@ -1963,10 +2158,10 @@ func (s *ProxyServiceServer) getTunnelPeerInfo(ctx context.Context, domain strin
principalID := peer.UserID
displayIdentity := peer.Name
// Stored column first (cheap, but often empty for OIDC-provisioned users).
if user, uerr := s.usersManager.GetUser(ctx, peer.UserID); uerr == nil && user != nil {
principalID = user.Id
if user.Email != "" {
displayIdentity = user.Email
if owner != nil {
principalID = owner.Id
if owner.Email != "" {
displayIdentity = owner.Email
}
}
// IdP enrichment wins when available — the stored email column is a
@@ -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")
}
@@ -119,11 +119,13 @@ func (m *mockReverseProxyManager) GetClusters(_ context.Context, _, _ string) ([
}
type mockUsersManager struct {
users map[string]*types.User
err error
users map[string]*types.User
err error
getUserCalls int
}
func (m *mockUsersManager) GetUser(ctx context.Context, userID string) (*types.User, error) {
m.getUserCalls++
if m.err != nil {
return nil, m.err
}
@@ -350,6 +352,64 @@ func TestValidateUserGroupAccess(t *testing.T) {
},
expectErr: false,
},
{
name: "user pending approval denied despite group membership",
domain: "app.example.com",
userID: "user1",
proxiesByAccount: map[string][]*service.Service{
"account1": {{
Domain: "app.example.com",
AccountID: "account1",
Auth: service.AuthConfig{
BearerAuth: &service.BearerAuthConfig{
Enabled: true,
DistributionGroups: []string{"group1"},
},
},
}},
},
users: map[string]*types.User{
// The approval flow stores a pending user as blocked as well.
"user1": {Id: "user1", AccountID: "account1", AutoGroups: []string{"group1"}, Blocked: true, PendingApproval: true},
},
expectErr: true,
expectErrMsg: "user pending approval",
},
{
name: "blocked user denied despite group membership",
domain: "app.example.com",
userID: "user1",
proxiesByAccount: map[string][]*service.Service{
"account1": {{
Domain: "app.example.com",
AccountID: "account1",
Auth: service.AuthConfig{
BearerAuth: &service.BearerAuthConfig{
Enabled: true,
DistributionGroups: []string{"group1"},
},
},
}},
},
users: map[string]*types.User{
"user1": {Id: "user1", AccountID: "account1", AutoGroups: []string{"group1"}, Blocked: true},
},
expectErr: true,
expectErrMsg: "user blocked",
},
{
name: "blocked user denied on a service with no auth configured",
domain: "app.example.com",
userID: "user1",
proxiesByAccount: map[string][]*service.Service{
"account1": {{Domain: "app.example.com", AccountID: "account1", Auth: service.AuthConfig{}}},
},
users: map[string]*types.User{
"user1": {Id: "user1", AccountID: "account1", Blocked: true},
},
expectErr: true,
expectErrMsg: "user blocked",
},
{
name: "proxy manager error",
domain: "app.example.com",
@@ -421,17 +481,18 @@ func TestValidateTunnelPeerUserEmailEnrichment(t *testing.T) {
storedUserNoEmail := map[string]*types.User{userID: {Id: userID, AccountID: accountID, Email: ""}}
tests := []struct {
name string
peerUserID string
storedUsers map[string]*types.User
storedErr error
noIdP bool
idpEmail string
idpHasData bool
idpErr error
expectEmail string
expectUserID string
expectIdPHit bool
name string
peerUserID string
storedUsers map[string]*types.User
storedErr error
noIdP bool
idpEmail string
idpHasData bool
idpErr error
expectEmail string
expectUserID string
expectIdPHit bool
expectDeniedReason string
}{
{
name: "idp email wins over stored email",
@@ -490,14 +551,17 @@ func TestValidateTunnelPeerUserEmailEnrichment(t *testing.T) {
expectIdPHit: true,
},
{
name: "idp email when stored user missing keeps peer.UserID as principal",
peerUserID: userID,
storedUsers: map[string]*types.User{},
idpEmail: "idp@example.com",
idpHasData: true,
expectEmail: "idp@example.com",
expectUserID: userID,
expectIdPHit: true,
// The identity still resolves from the IdP, but an owner the store
// cannot resolve denies the fast-path rather than granting it.
name: "idp email when stored user missing keeps peer.UserID as principal",
peerUserID: userID,
storedUsers: map[string]*types.User{},
idpEmail: "idp@example.com",
idpHasData: true,
expectEmail: "idp@example.com",
expectUserID: userID,
expectIdPHit: true,
expectDeniedReason: deniedReasonUserNotFound,
},
{
name: "unlinked peer uses peer name and never consults idp",
@@ -545,9 +609,13 @@ func TestValidateTunnelPeerUserEmailEnrichment(t *testing.T) {
require.NoError(t, err)
require.NotNil(t, resp)
assert.True(t, resp.GetValid(), "expected access granted")
assert.Equal(t, tt.expectDeniedReason == "", resp.GetValid(), "unexpected access decision")
assert.Equal(t, tt.expectDeniedReason, resp.GetDeniedReason(), "unexpected denied reason")
assert.Equal(t, tt.expectEmail, resp.GetUserEmail())
assert.Equal(t, tt.expectUserID, resp.GetUserId())
if tt.expectDeniedReason != "" {
assert.Empty(t, resp.GetSessionToken(), "a denied peer must not receive a session token")
}
if idpMock != nil {
if tt.expectIdPHit {
@@ -562,6 +630,121 @@ func TestValidateTunnelPeerUserEmailEnrichment(t *testing.T) {
}
}
// TestDeniedReasonValues pins the wire values of the account status denied
// reasons. The proxy logs them and operators filter access logs on them, so a
// rename is a breaking change rather than an internal detail.
// TestSameAccount pins the fail-closed behaviour of the account binding: an
// unset account on either side must never compare equal into a grant.
func TestSameAccount(t *testing.T) {
assert.True(t, sameAccount("account1", "account1"), "matching accounts should bind")
assert.False(t, sameAccount("account1", "account2"), "different accounts must not bind")
assert.False(t, sameAccount("", ""), "two unset accounts must not bind")
assert.False(t, sameAccount("account1", ""), "an unset service account must not bind")
assert.False(t, sameAccount("", "account1"), "an unset user account must not bind")
}
func TestDeniedReasonValues(t *testing.T) {
assert.Equal(t, "pending_approval", deniedReasonPendingApproval, "pending approval denied reason wire value")
assert.Equal(t, "user_blocked", deniedReasonUserBlocked, "blocked user denied reason wire value")
assert.Equal(t, "user_not_found", deniedReasonUserNotFound, "unresolved user denied reason wire value")
}
// TestValidateTunnelPeerOwnerStatus verifies that the mesh fast-path gates on
// the account status of the peer's owning user. A peer whose owner was blocked
// after the peer registered must lose access, while an unlinked machine peer
// keeps it.
func TestValidateTunnelPeerOwnerStatus(t *testing.T) {
const (
domain = "app.example.com"
accountID = "account1"
peerID = "peer1"
peerName = "peer-display-name"
userID = "user1"
)
tests := []struct {
name string
peerUserID string
owner *types.User
expectDeniedReason string
expectEmail string
}{
{
name: "active owner allowed",
peerUserID: userID,
owner: &types.User{Id: userID, AccountID: accountID, Email: "user@example.com"},
},
{
name: "owner pending approval denied",
peerUserID: userID,
owner: &types.User{Id: userID, AccountID: accountID, Email: "user@example.com", Blocked: true, PendingApproval: true},
expectDeniedReason: deniedReasonPendingApproval,
},
{
name: "owner blocked after registering the peer denied",
peerUserID: userID,
owner: &types.User{Id: userID, AccountID: accountID, Email: "user@example.com", Blocked: true},
expectDeniedReason: deniedReasonUserBlocked,
},
{
name: "unlinked machine peer stays allowed",
peerUserID: "",
owner: &types.User{Id: userID, AccountID: accountID, Blocked: true},
},
{
// The user lookup is not account-scoped, so a peer row pointing at
// another account's user must not resolve into an owner: the peer is
// denied and the foreign email never reaches the response.
name: "owner in another account denied and not disclosed",
peerUserID: userID,
owner: &types.User{Id: userID, AccountID: "otherAccount", Email: "foreign@example.com"},
expectDeniedReason: deniedReasonUserNotFound,
expectEmail: peerName,
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
svc := &service.Service{Domain: domain, AccountID: accountID}
usersManager := &mockUsersManager{users: map[string]*types.User{userID: tt.owner}}
server := &ProxyServiceServer{
serviceManager: &mockReverseProxyManager{
proxiesByAccount: map[string][]*service.Service{accountID: {svc}},
},
peersManager: &mockTunnelPeersManager{
peer: &peer.Peer{ID: peerID, Name: peerName, UserID: tt.peerUserID},
},
usersManager: usersManager,
}
resp, err := server.ValidateTunnelPeer(context.Background(), &proto.ValidateTunnelPeerRequest{
Domain: domain,
TunnelIp: "100.64.0.1",
})
require.NoError(t, err)
require.NotNil(t, resp)
assert.Equal(t, tt.expectDeniedReason, resp.GetDeniedReason(), "unexpected denied reason")
assert.Equal(t, tt.expectDeniedReason == "", resp.GetValid(), "unexpected access decision")
if tt.expectDeniedReason != "" {
assert.Empty(t, resp.GetSessionToken(), "a denied peer must not receive a session token")
}
if tt.expectEmail != "" {
assert.Equal(t, tt.expectEmail, resp.GetUserEmail(), "unexpected identity on the response")
}
// The status gate and the identity resolution share one lookup;
// an unlinked peer has no owner to look up at all.
wantLookups := 1
if tt.peerUserID == "" {
wantLookups = 0
}
assert.Equal(t, wantLookups, usersManager.getUserCalls, "owner must be resolved exactly once per request")
})
}
}
func TestGetAccountProxyByDomain(t *testing.T) {
tests := []struct {
name string
@@ -46,6 +46,7 @@ func setupValidateSessionTest(t *testing.T) *validateSessionTestSetup {
proxyService.SetServiceManager(serviceManager)
createTestProxies(t, ctx, testStore)
createStatusTestUsers(t, ctx, testStore)
return &validateSessionTestSetup{
proxyService: proxyService,
@@ -91,6 +92,82 @@ func createTestProxies(t *testing.T, ctx context.Context, testStore store.Store)
},
}
require.NoError(t, testStore.CreateService(ctx, restrictedProxy))
// Distributed to the account's "All" group, the configuration that hands a
// service to every user in the account.
allUsersProxy := &service.Service{
ID: "allUsersProxyId",
AccountID: "testAccountId",
Name: "All Users Proxy",
Domain: "all-users-proxy.example.com",
Enabled: true,
SessionPrivateKey: privKey,
SessionPublicKey: pubKey,
Auth: service.AuthConfig{
BearerAuth: &service.BearerAuthConfig{
Enabled: true,
DistributionGroups: []string{allUsersGroupID},
},
},
}
require.NoError(t, testStore.CreateService(ctx, allUsersProxy))
}
const (
allUsersGroupID = "allUsersGroupId"
pendingUserID = "pendingUserId"
blockedUserID = "blockedUserId"
pendingAllUsersID = "pendingAllUsersUserId"
)
// createStatusTestUsers adds the users whose account status must keep them out
// of a proxy session. A user awaiting approval is persisted as both blocked and
// pending approval, the way the approval flow stores one.
func createStatusTestUsers(t *testing.T, ctx context.Context, testStore store.Store) {
t.Helper()
require.NoError(t, testStore.CreateGroup(ctx, &types.Group{
ID: allUsersGroupID,
AccountID: "testAccountId",
Name: "All",
Issued: types.GroupIssuedAPI,
}))
users := []*types.User{
{
Id: pendingUserID,
AccountID: "testAccountId",
Role: types.UserRoleUser,
AutoGroups: []string{"allowedGroupId"},
Blocked: true,
PendingApproval: true,
Issued: "api",
CreatedAt: time.Now(),
},
{
Id: pendingAllUsersID,
AccountID: "testAccountId",
Role: types.UserRoleUser,
AutoGroups: []string{allUsersGroupID},
Blocked: true,
PendingApproval: true,
Issued: "api",
CreatedAt: time.Now(),
},
{
Id: blockedUserID,
AccountID: "testAccountId",
Role: types.UserRoleUser,
AutoGroups: []string{"allowedGroupId"},
Blocked: true,
PendingApproval: false,
Issued: "api",
CreatedAt: time.Now(),
},
}
for _, user := range users {
require.NoError(t, testStore.SaveUser(ctx, user))
}
}
func generateSessionKeyPair(t *testing.T) (string, string) {
@@ -149,6 +226,114 @@ func TestValidateSession_UserNotInAllowedGroup(t *testing.T) {
assert.Empty(t, resp.GetPeerGroupIds(), "PeerGroupIds must mirror the resolved user's actual (empty) memberships on denial")
}
// TestValidateSession_PendingApprovalUserDenied covers a user who is a member of
// the service's distribution group but is still waiting for an administrator to
// approve the account. Group membership alone must not open the service.
func TestValidateSession_PendingApprovalUserDenied(t *testing.T) {
setup := setupValidateSessionTest(t)
defer setup.cleanup()
proxy, err := setup.store.GetServiceByID(context.Background(), store.LockingStrengthNone, "testAccountId", "restrictedProxyId")
require.NoError(t, err)
token := createSessionToken(t, proxy.SessionPrivateKey, pendingUserID, "restricted-proxy.example.com")
resp, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
Domain: "restricted-proxy.example.com",
SessionToken: token,
})
require.NoError(t, err)
assert.False(t, resp.Valid, "User pending approval should be denied")
assert.Equal(t, deniedReasonPendingApproval, resp.DeniedReason, "Denied reason should name the pending approval state")
assert.Equal(t, pendingUserID, resp.UserId, "Denial should identify the user it applies to")
assert.Equal(t, []string{"allowedGroupId"}, resp.GetPeerGroupIds(), "PeerGroupIds must mirror the resolved user's group memberships on denial")
assert.Equal(t, []string{"Allowed Group"}, resp.GetPeerGroupNames(), "PeerGroupNames must pair with PeerGroupIds on denial")
}
// TestValidateSession_PendingApprovalUserInAllUsersGroupDenied covers the same
// user against a service distributed to the account's "All" group, where every
// user of the account is a member by default.
func TestValidateSession_PendingApprovalUserInAllUsersGroupDenied(t *testing.T) {
setup := setupValidateSessionTest(t)
defer setup.cleanup()
proxy, err := setup.store.GetServiceByID(context.Background(), store.LockingStrengthNone, "testAccountId", "allUsersProxyId")
require.NoError(t, err)
token := createSessionToken(t, proxy.SessionPrivateKey, pendingAllUsersID, "all-users-proxy.example.com")
resp, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
Domain: "all-users-proxy.example.com",
SessionToken: token,
})
require.NoError(t, err)
assert.False(t, resp.Valid, "User pending approval should be denied even in the All Users group")
assert.Equal(t, deniedReasonPendingApproval, resp.DeniedReason, "Denied reason should name the pending approval state")
assert.Equal(t, pendingAllUsersID, resp.UserId, "Denial should identify the user it applies to")
assert.Equal(t, []string{allUsersGroupID}, resp.GetPeerGroupIds(), "PeerGroupIds must mirror the resolved user's group memberships on denial")
}
// TestValidateSession_BlockedUserDenied covers a user blocked after having been
// approved, so PendingApproval is false and only the blocked flag is set.
func TestValidateSession_BlockedUserDenied(t *testing.T) {
setup := setupValidateSessionTest(t)
defer setup.cleanup()
proxy, err := setup.store.GetServiceByID(context.Background(), store.LockingStrengthNone, "testAccountId", "restrictedProxyId")
require.NoError(t, err)
token := createSessionToken(t, proxy.SessionPrivateKey, blockedUserID, "restricted-proxy.example.com")
resp, err := setup.proxyService.ValidateSession(context.Background(), &proto.ValidateSessionRequest{
Domain: "restricted-proxy.example.com",
SessionToken: token,
})
require.NoError(t, err)
assert.False(t, resp.Valid, "Blocked user should be denied")
assert.Equal(t, deniedReasonUserBlocked, resp.DeniedReason, "Denied reason should name the blocked state")
assert.Equal(t, blockedUserID, resp.UserId, "Denial should identify the user it applies to")
}
// TestValidateSession_UserAllowedAfterApproval walks the same session token
// through the approval transition: denied while pending, allowed once an
// administrator clears both flags.
func TestValidateSession_UserAllowedAfterApproval(t *testing.T) {
setup := setupValidateSessionTest(t)
defer setup.cleanup()
ctx := context.Background()
proxy, err := setup.store.GetServiceByID(ctx, store.LockingStrengthNone, "testAccountId", "restrictedProxyId")
require.NoError(t, err)
token := createSessionToken(t, proxy.SessionPrivateKey, pendingUserID, "restricted-proxy.example.com")
req := &proto.ValidateSessionRequest{
Domain: "restricted-proxy.example.com",
SessionToken: token,
}
resp, err := setup.proxyService.ValidateSession(ctx, req)
require.NoError(t, err)
require.False(t, resp.Valid, "User pending approval should be denied before approval")
assert.Equal(t, deniedReasonPendingApproval, resp.DeniedReason, "Denied reason should name the pending approval state")
user, err := setup.store.GetUserByUserID(ctx, store.LockingStrengthNone, pendingUserID)
require.NoError(t, err)
user.PendingApproval = false
user.Blocked = false
require.NoError(t, setup.store.SaveUser(ctx, user))
resp, err = setup.proxyService.ValidateSession(ctx, req)
require.NoError(t, err)
assert.True(t, resp.Valid, "Approved user should be allowed access")
assert.Empty(t, resp.DeniedReason)
assert.Equal(t, pendingUserID, resp.UserId, "Approved user should be identified in the response")
assert.Equal(t, []string{"allowedGroupId"}, resp.GetPeerGroupIds(), "PeerGroupIds must mirror the approved user's group memberships")
}
func TestValidateSession_UserInDifferentAccount(t *testing.T) {
setup := setupValidateSessionTest(t)
defer setup.cleanup()
+49 -39
View File
@@ -33,6 +33,7 @@ import (
nbconfig "github.com/netbirdio/netbird/management/internals/server/config"
"github.com/netbirdio/netbird/management/server/account"
"github.com/netbirdio/netbird/management/server/activity"
"github.com/netbirdio/netbird/management/server/affectedpeers"
nbcache "github.com/netbirdio/netbird/management/server/cache"
nbcontext "github.com/netbirdio/netbird/management/server/context"
"github.com/netbirdio/netbird/management/server/geolocation"
@@ -1626,6 +1627,8 @@ func (am *DefaultAccountManager) SyncUserJWTGroups(ctx context.Context, userAuth
var removeOldGroups []string
var hasChanges bool
var user *types.User
var change affectedpeers.Change
var snap *affectedpeers.Snapshot
err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
user, err = transaction.GetUserByUserID(ctx, store.LockingStrengthNone, userAuth.UserId)
if err != nil {
@@ -1664,14 +1667,25 @@ func (am *DefaultAccountManager) SyncUserJWTGroups(ctx context.Context, userAuth
return fmt.Errorf("error saving user: %w", err)
}
allGroupChanges := slices.Concat(addNewGroups, removeOldGroups)
// The user's auto-groups changed, so the SSH rules authorizing them ship a new
// group -> user mapping even when no peer moves between groups.
change.UserGroupIDs = allGroupChanges
// The user's peers are the changed entity in every scenario the sync can
// produce — group membership, IPv6 assignment, SSH mappings — so they refresh
// together with every peer they can connect to, like on a regular peer update.
userPeers, err := transaction.GetUserPeers(ctx, store.LockingStrengthNone, userAuth.AccountId, userAuth.UserId)
if err != nil {
return fmt.Errorf("error getting user peers: %w", err)
}
for _, peer := range userPeers {
change.ChangedPeerIDs = append(change.ChangedPeerIDs, peer.ID)
}
// Propagate changes to peers if group propagation is enabled
if settings.GroupsPropagationEnabled {
peers, err := transaction.GetUserPeers(ctx, store.LockingStrengthNone, userAuth.AccountId, userAuth.UserId)
if err != nil {
return fmt.Errorf("error getting user peers: %w", err)
}
for _, peer := range peers {
for _, peer := range userPeers {
for _, g := range addNewGroups {
if err := transaction.AddPeerToGroup(ctx, userAuth.AccountId, peer.ID, g); err != nil {
return fmt.Errorf("error adding peer %s to group %s: %w", peer.ID, g, err)
@@ -1684,7 +1698,8 @@ func (am *DefaultAccountManager) SyncUserJWTGroups(ctx context.Context, userAuth
}
}
allGroupChanges := slices.Concat(addNewGroups, removeOldGroups)
change.LinkGroups = allGroupChanges
if err = am.reconcileIPv6ForGroupChanges(ctx, transaction, userAuth.AccountId, allGroupChanges); err != nil {
return fmt.Errorf("reconcile IPv6 for group changes: %w", err)
}
@@ -1694,6 +1709,10 @@ func (am *DefaultAccountManager) SyncUserJWTGroups(ctx context.Context, userAuth
}
}
if snap, err = affectedpeers.Load(ctx, transaction, userAuth.AccountId, change); err != nil {
return err
}
return nil
})
if err != nil {
@@ -1730,20 +1749,17 @@ func (am *DefaultAccountManager) SyncUserJWTGroups(ctx context.Context, userAuth
}
}
removedGroupAffectsPeers, err := areGroupChangesAffectPeers(ctx, am.Store, userAuth.AccountId, removeOldGroups)
if err != nil {
return err
}
newGroupsAffectsPeers, err := areGroupChangesAffectPeers(ctx, am.Store, userAuth.AccountId, addNewGroups)
if err != nil {
return err
}
if removedGroupAffectsPeers || newGroupsAffectsPeers {
log.WithContext(ctx).Tracef("user %s: JWT group membership changed, updating account peers", userAuth.UserId)
am.BufferUpdateAccountPeers(ctx, userAuth.AccountId, types.UpdateReason{Resource: types.UpdateResourceUser, Operation: types.UpdateOperationUpdate})
}
log.WithContext(ctx).Tracef("user %s: JWT group membership changed, updating affected peers", userAuth.UserId)
bgCtx := context.WithoutCancel(ctx)
go func() {
affectedPeerIDs := snap.Expand(bgCtx, userAuth.AccountId, change)
if len(affectedPeerIDs) == 0 {
return
}
if err := am.networkMapController.BufferUpdateAffectedPeers(bgCtx, userAuth.AccountId, affectedPeerIDs, types.UpdateReason{Resource: types.UpdateResourceUser, Operation: types.UpdateOperationUpdate}); err != nil {
log.WithContext(bgCtx).Errorf("failed to update affected peers after JWT group sync for account %s: %v", userAuth.AccountId, err)
}
}()
return nil
}
@@ -2426,30 +2442,24 @@ func (am *DefaultAccountManager) reconcileIPv6ForGroupChanges(ctx context.Contex
return fmt.Errorf("get account settings: %w", err)
}
if len(settings.IPv6EnabledGroups) == 0 {
return nil
}
enabledSet := make(map[string]struct{}, len(settings.IPv6EnabledGroups))
for _, gid := range settings.IPv6EnabledGroups {
enabledSet[gid] = struct{}{}
}
affected := false
for _, gid := range groupIDs {
if _, ok := enabledSet[gid]; ok {
affected = true
break
}
}
if !affected {
if !ipv6ReconcileNeeded(settings, groupIDs) {
return nil
}
return am.updatePeerIPv6Addresses(ctx, transaction, accountID, settings)
}
// ipv6ReconcileNeeded reports whether changes to the given groups trigger an IPv6
// reconciliation.
func ipv6ReconcileNeeded(settings *types.Settings, groupIDs []string) bool {
for _, groupID := range groupIDs {
if slices.Contains(settings.IPv6EnabledGroups, groupID) {
return true
}
}
return false
}
func (am *DefaultAccountManager) ensureIPv6Subnet(ctx context.Context, transaction store.Store, accountID string, settings *types.Settings, network *types.Network) error {
if settings.NetworkRangeV6.IsValid() {
network.NetV6 = net.IPNet{
+1
View File
@@ -1757,6 +1757,7 @@ func TestAccount_Copy(t *testing.T) {
AccountID: "account1",
},
},
PostureValidation: map[string]map[string]bool{"1": {"1": true}},
}
err := hasNilField(account)
if err != nil {
+4
View File
@@ -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"},
@@ -0,0 +1,179 @@
package server
import (
"context"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
"github.com/netbirdio/netbird/management/server/affectedpeers"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/types"
"github.com/netbirdio/netbird/shared/auth"
)
// A user's auto-group change refreshes the destinations of the SSH rules authorizing
// that group — they carry the group -> user mapping — even though no peer moved
// between groups.
func TestAffectedPeers_UserGroupChange_RefreshesSSHAuthorizedDestinations(t *testing.T) {
manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t)
ctx := context.Background()
_, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{
Enabled: true,
Rules: []*types.PolicyRule{
{
Enabled: true,
Sources: []string{groupIDs[0]},
Destinations: []string{groupIDs[1]},
Protocol: types.PolicyRuleProtocolNetbirdSSH,
Action: types.PolicyTrafficActionAccept,
AuthorizedGroups: map[string][]string{groupIDs[3]: {"root"}},
},
},
}, true)
require.NoError(t, err)
result := resolveAffected(t, s, accountID, affectedpeers.Change{UserGroupIDs: []string{groupIDs[3]}})
assert.ElementsMatch(t, []string{peerIDs[1]}, result,
"only the SSH rule's destination peers carry the changed group -> user mapping")
result = resolveAffected(t, s, accountID, affectedpeers.Change{UserGroupIDs: []string{groupIDs[4]}})
assert.Empty(t, result, "a group no SSH rule authorizes affects nobody")
}
// Creating, blocking or unblocking a user changes the account's allowed-user set, which
// reaches only the destinations of the SSH rules that ship it.
func TestAffectedPeers_AllowedUsersChange_RefreshesSSHDestinations(t *testing.T) {
manager, s, accountID, peerIDs, groupIDs := setupAffectedPeersTest(t)
ctx := context.Background()
// Ships the allowed-user set: an SSH rule naming no groups and no user.
_, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{
Enabled: true,
Rules: []*types.PolicyRule{{
Enabled: true,
Sources: []string{groupIDs[0]},
Destinations: []string{groupIDs[1]},
Protocol: types.PolicyRuleProtocolNetbirdSSH,
Action: types.PolicyTrafficActionAccept,
}},
}, true)
require.NoError(t, err)
// Does not ship it: an SSH rule that authorizes a specific group.
_, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{
Enabled: true,
Rules: []*types.PolicyRule{{
Enabled: true,
Sources: []string{groupIDs[2]},
Destinations: []string{groupIDs[3]},
Protocol: types.PolicyRuleProtocolNetbirdSSH,
Action: types.PolicyTrafficActionAccept,
AuthorizedGroups: map[string][]string{groupIDs[0]: {"root"}},
}},
}, true)
require.NoError(t, err)
result := resolveAffected(t, s, accountID, affectedpeers.Change{AllowedUsersChanged: true})
assert.ElementsMatch(t, []string{peerIDs[1]}, result,
"only the destinations of the rule shipping the allowed-user set refresh")
}
// TestAffectedPeers_SyncUserJWTGroups_OnlyAffectedPeersUpdated verifies that a JWT
// auto-group change updates only the user's peers and the peers linked to the changed
// group through policies, instead of fanning out to the whole account.
func TestAffectedPeers_SyncUserJWTGroups_OnlyAffectedPeersUpdated(t *testing.T) {
manager, updateManager, account, _, peer2, peer3 := setupNetworkMapTest(t)
ctx := context.Background()
accountID := account.Id
key, err := wgtypes.GeneratePrivateKey()
require.NoError(t, err)
userPeer, _, _, _, err := manager.AddPeer(ctx, accountID, "", userID, &nbpeer.Peer{
Key: key.PublicKey().String(),
Meta: nbpeer.PeerSystemMeta{Hostname: "user-peer"},
}, false)
require.NoError(t, err)
policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID)
require.NoError(t, err)
for _, p := range policies {
require.NoError(t, manager.Store.DeletePolicy(ctx, accountID, p.ID))
}
account, err = manager.Store.GetAccount(ctx, accountID)
require.NoError(t, err)
account.Settings.JWTGroupsEnabled = true
account.Settings.JWTGroupsClaimName = "groups"
account.Settings.GroupsPropagationEnabled = true
require.NoError(t, manager.Store.SaveAccount(ctx, account))
require.NoError(t, manager.CreateGroup(ctx, accountID, userID, &types.Group{ID: "jwt-grp", Name: "jwt-linked", Issued: types.GroupIssuedJWT, Peers: []string{}}))
require.NoError(t, manager.CreateGroup(ctx, accountID, userID, &types.Group{ID: "jwt-dest", Name: "jwt-dest", Peers: []string{peer2.ID}}))
_, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{
Enabled: true,
Rules: []*types.PolicyRule{
{
Enabled: true,
Sources: []string{"jwt-grp"},
Destinations: []string{"jwt-dest"},
Bidirectional: true,
Action: types.PolicyTrafficActionAccept,
},
},
}, true)
require.NoError(t, err)
updUser := updateManager.CreateChannel(ctx, userPeer.ID)
upd2 := updateManager.CreateChannel(ctx, peer2.ID)
upd3 := updateManager.CreateChannel(ctx, peer3.ID)
t.Cleanup(func() {
updateManager.CloseChannel(ctx, userPeer.ID)
updateManager.CloseChannel(ctx, peer2.ID)
updateManager.CloseChannel(ctx, peer3.ID)
})
userAuth := auth.UserAuth{
AccountId: accountID,
UserId: userID,
Groups: []string{"jwt-linked"},
}
t.Run("adding JWT group updates only linked peers", func(t *testing.T) {
drainPeerUpdates(updUser)
drainPeerUpdates(upd2)
drainPeerUpdates(upd3)
require.NoError(t, manager.SyncUserJWTGroups(ctx, userAuth))
peerShouldReceiveUpdate(t, updUser)
peerShouldReceiveUpdate(t, upd2)
peerShouldNotReceiveUpdate(t, upd3)
user, err := manager.Store.GetUserByUserID(ctx, store.LockingStrengthNone, userID)
require.NoError(t, err)
assert.Contains(t, user.AutoGroups, "jwt-grp")
})
t.Run("removing JWT group updates only linked peers", func(t *testing.T) {
drainPeerUpdates(updUser)
drainPeerUpdates(upd2)
drainPeerUpdates(upd3)
userAuth.Groups = nil
require.NoError(t, manager.SyncUserJWTGroups(ctx, userAuth))
peerShouldReceiveUpdate(t, updUser)
peerShouldReceiveUpdate(t, upd2)
peerShouldNotReceiveUpdate(t, upd3)
user, err := manager.Store.GetUserByUserID(ctx, store.LockingStrengthNone, userID)
require.NoError(t, err)
assert.NotContains(t, user.AutoGroups, "jwt-grp")
})
}
@@ -0,0 +1,170 @@
package server
import (
"context"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
"golang.zx2c4.com/wireguard/wgctrl/wgtypes"
"github.com/netbirdio/netbird/management/server/activity"
nbpeer "github.com/netbirdio/netbird/management/server/peer"
"github.com/netbirdio/netbird/management/server/store"
"github.com/netbirdio/netbird/management/server/types"
)
// A user update refreshes only the peers its auto-group change reaches, and a user
// update that changes no group membership refreshes nobody.
func TestAffectedPeers_SaveUser_OnlyAffectedPeersUpdated(t *testing.T) {
manager, updateManager, account, _, peer2, peer3 := setupNetworkMapTest(t)
ctx := context.Background()
accountID := account.Id
const targetUserID = "target-user"
require.NoError(t, manager.Store.SaveUser(ctx, &types.User{
Id: targetUserID, AccountID: accountID, Role: types.UserRoleUser,
}))
key, err := wgtypes.GeneratePrivateKey()
require.NoError(t, err)
targetPeer, _, _, _, err := manager.AddPeer(ctx, accountID, "", targetUserID, &nbpeer.Peer{
Key: key.PublicKey().String(),
Meta: nbpeer.PeerSystemMeta{Hostname: "target-peer"},
}, false)
require.NoError(t, err)
policies, err := manager.Store.GetAccountPolicies(ctx, store.LockingStrengthNone, accountID)
require.NoError(t, err)
for _, p := range policies {
require.NoError(t, manager.Store.DeletePolicy(ctx, accountID, p.ID))
}
account, err = manager.Store.GetAccount(ctx, accountID)
require.NoError(t, err)
account.Settings.GroupsPropagationEnabled = true
require.NoError(t, manager.Store.SaveAccount(ctx, account))
require.NoError(t, manager.CreateGroup(ctx, accountID, userID, &types.Group{ID: "ug-linked", Name: "ug-linked"}))
require.NoError(t, manager.CreateGroup(ctx, accountID, userID, &types.Group{ID: "ug-dest", Name: "ug-dest", Peers: []string{peer2.ID}}))
_, err = manager.SavePolicy(ctx, accountID, userID, &types.Policy{
Enabled: true,
Rules: []*types.PolicyRule{
{
Enabled: true,
Sources: []string{"ug-linked"},
Destinations: []string{"ug-dest"},
Bidirectional: true,
Action: types.PolicyTrafficActionAccept,
},
},
}, true)
require.NoError(t, err)
updTarget := updateManager.CreateChannel(ctx, targetPeer.ID)
upd2 := updateManager.CreateChannel(ctx, peer2.ID)
upd3 := updateManager.CreateChannel(ctx, peer3.ID)
t.Cleanup(func() {
updateManager.CloseChannel(ctx, targetPeer.ID)
updateManager.CloseChannel(ctx, peer2.ID)
updateManager.CloseChannel(ctx, peer3.ID)
})
t.Run("auto group change updates only linked peers", func(t *testing.T) {
drainPeerUpdates(updTarget)
drainPeerUpdates(upd2)
drainPeerUpdates(upd3)
_, err := manager.SaveUser(ctx, accountID, activity.SystemInitiator, &types.User{
Id: targetUserID, AccountID: accountID, Role: types.UserRoleUser,
AutoGroups: []string{"ug-linked"},
})
require.NoError(t, err)
peerShouldReceiveUpdate(t, updTarget)
peerShouldReceiveUpdate(t, upd2)
peerShouldNotReceiveUpdate(t, upd3)
})
t.Run("update without group changes refreshes nobody", func(t *testing.T) {
drainPeerUpdates(updTarget)
drainPeerUpdates(upd2)
drainPeerUpdates(upd3)
_, err := manager.SaveUser(ctx, accountID, activity.SystemInitiator, &types.User{
Id: targetUserID, AccountID: accountID, Role: types.UserRoleUser,
AutoGroups: []string{"ug-linked"}, Name: "renamed",
})
require.NoError(t, err)
peerShouldNotReceiveUpdate(t, updTarget)
peerShouldNotReceiveUpdate(t, upd2)
peerShouldNotReceiveUpdate(t, upd3)
user, err := manager.Store.GetUserByUserID(ctx, store.LockingStrengthNone, targetUserID)
require.NoError(t, err)
assert.Equal(t, "renamed", user.Name)
})
t.Run("auto group change reassigning IPv6 refreshes the changed peers and their observers", func(t *testing.T) {
account, err := manager.Store.GetAccount(ctx, accountID)
require.NoError(t, err)
account.Settings.IPv6EnabledGroups = []string{"ug-v6"}
require.NoError(t, manager.Store.SaveAccount(ctx, account))
require.NoError(t, manager.CreateGroup(ctx, accountID, userID, &types.Group{ID: "ug-v6", Name: "ug-v6"}))
drainPeerUpdates(updTarget)
drainPeerUpdates(upd2)
drainPeerUpdates(upd3)
_, err = manager.SaveUser(ctx, accountID, activity.SystemInitiator, &types.User{
Id: targetUserID, AccountID: accountID, Role: types.UserRoleUser,
AutoGroups: []string{"ug-linked", "ug-v6"}, Name: "renamed",
})
require.NoError(t, err)
// The reassigned peer refreshes with everyone it can reach: peer2 via the
// policy, but not peer3, which shares no group or policy with it.
peerShouldReceiveUpdate(t, updTarget)
peerShouldReceiveUpdate(t, upd2)
peerShouldNotReceiveUpdate(t, upd3)
})
t.Run("unblocking a user refreshes only the SSH rule destinations", func(t *testing.T) {
// An SSH rule that authorizes no group of its own ships the account's
// allowed-user set to its destinations, so those are the peers an unblock
// reaches — not the whole account.
_, err := manager.SavePolicy(ctx, accountID, userID, &types.Policy{
Enabled: true,
Rules: []*types.PolicyRule{{
Enabled: true,
Sources: []string{"ug-linked"},
Destinations: []string{"ug-dest"},
Protocol: types.PolicyRuleProtocolNetbirdSSH,
Action: types.PolicyTrafficActionAccept,
}},
}, true)
require.NoError(t, err)
blocked, err := manager.Store.GetUserByUserID(ctx, store.LockingStrengthNone, targetUserID)
require.NoError(t, err)
blocked.Blocked = true
require.NoError(t, manager.Store.SaveUser(ctx, blocked))
drainPeerUpdates(updTarget)
drainPeerUpdates(upd2)
drainPeerUpdates(upd3)
// Same auto-groups as the previous subtest left them, so no group change and
// no IPv6 reconciliation interferes: the unblock alone drives the refresh.
_, err = manager.SaveUser(ctx, accountID, activity.SystemInitiator, &types.User{
Id: targetUserID, AccountID: accountID, Role: types.UserRoleUser,
AutoGroups: []string{"ug-linked", "ug-v6"}, Name: "renamed",
})
require.NoError(t, err)
peerShouldReceiveUpdate(t, upd2)
peerShouldNotReceiveUpdate(t, upd3)
})
}
+72 -1
View File
@@ -18,6 +18,7 @@ import (
"context"
log "github.com/sirupsen/logrus"
"golang.org/x/exp/maps"
nbdns "github.com/netbirdio/netbird/dns"
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
@@ -83,7 +84,7 @@ func (snap *Snapshot) loadCollections(ctx context.Context, s store.Store, accoun
hasGroupOrPeerChange := len(c.ChangedGroupIDs) > 0 || len(c.ChangedPeerIDs) > 0 || len(c.LinkGroups) > 0 || len(c.Resources) > 0
hasNetworkObject := len(c.Routers) > 0 || len(c.Resources) > 0 || len(c.Networks) > 0
// the resource<->router bridge can fire for any of these
needsRoutersResources := hasGroupOrPeerChange || len(c.PostureCheckIDs) > 0 || len(c.Policies) > 0 || hasNetworkObject
needsRoutersResources := hasGroupOrPeerChange || len(c.PostureCheckIDs) > 0 || len(c.Policies) > 0 || hasNetworkObject || len(c.UserGroupIDs) > 0 || c.AllowedUsersChanged
if needsRoutersResources {
if err := snap.loadPolicyRoutersResources(ctx, s, accountID); err != nil {
@@ -219,6 +220,18 @@ type Change struct {
// (correct when the peer's own attributes changed, e.g. IP/status).
OutputPeerIDs []string
// UserGroupIDs are groups whose USER membership changed (a user's auto-groups),
// as opposed to their peer membership. Peers ship the group -> user mapping only
// for the groups an SSH rule authorizes, so these refresh the destinations of the
// SSH rules authorizing them — independently of any peer moving between groups.
UserGroupIDs []string
// AllowedUsersChanged marks a change to the set of users allowed to open SSH
// sessions — a user was created, blocked or unblocked. That set is account-wide,
// and peers receive it through the SSH rules that name no group or user of their
// own, so those rules' destinations refresh.
AllowedUsersChanged bool
// LinkGroups are groups used ONLY to match policies/routes/routers and walk to the
// OPPOSITE side — they are never expanded to their own members. Use this when a
// peer's group membership changed: pass the peer in ChangedPeerIDs and its
@@ -240,6 +253,8 @@ func (c Change) isEmpty() bool {
len(c.Resources) == 0 &&
len(c.Networks) == 0 &&
len(c.PostureCheckIDs) == 0 &&
len(c.UserGroupIDs) == 0 &&
!c.AllowedUsersChanged &&
len(c.DistributionGroupIDs) == 0 &&
len(c.RemovedPeersByGroup) == 0 &&
len(c.LinkGroups) == 0 &&
@@ -359,6 +374,9 @@ func (r *resolver) walk() {
r.collectFromProxyServices()
}
r.collectFromSSHAuthorizedGroups()
r.collectFromAllowedUsers()
r.collectFromChangedRoutes(r.change.Routes)
r.collectFromChangedRouters(r.change.Routers)
r.collectFromChangedResources(r.change.Resources)
@@ -811,6 +829,59 @@ func (r *resolver) collectFromNameServers() {
}
}
// collectFromSSHAuthorizedGroups folds the destinations of the enabled SSH rules that
// authorize a group whose user membership changed. Those destination peers carry the
// group -> user mapping for the groups they authorize, so they refresh even when no
// peer moved between groups.
func (r *resolver) collectFromSSHAuthorizedGroups() {
if len(r.change.UserGroupIDs) == 0 {
return
}
changed := toSet(r.change.UserGroupIDs)
for _, policy := range r.policies() {
for _, rule := range policy.Rules {
if !rule.Enabled || rule.Protocol != types.PolicyRuleProtocolNetbirdSSH {
continue
}
if !anyInSet(maps.Keys(rule.AuthorizedGroups), changed) {
continue
}
log.WithContext(r.ctx).Tracef("collectFromSSHAuthorizedGroups: rule %s authorizes a changed user group -> folding its destinations", rule.ID)
r.foldPolicySideForRule(policy, rule, sideDestination)
}
}
}
// collectFromAllowedUsers folds the destinations of the rules that make a peer carry
// the account's allowed-user set, for a change to who is in that set.
func (r *resolver) collectFromAllowedUsers() {
if !r.change.AllowedUsersChanged {
return
}
for _, policy := range r.policies() {
for _, rule := range policy.Rules {
if !rule.Enabled || !ruleShipsAllowedUsers(rule) {
continue
}
log.WithContext(r.ctx).Tracef("collectFromAllowedUsers: rule %s ships the allowed-user set -> folding its destinations", rule.ID)
r.foldPolicySideForRule(policy, rule, sideDestination)
}
}
}
// ruleShipsAllowedUsers reports whether a rule makes its destination peers carry the
// account's allowed-user set. It mirrors the network map's SSH requirements except for
// the destination peer's own SSH flag, which the snapshot does not hold — so it folds a
// superset and never misses a peer.
func ruleShipsAllowedUsers(rule *types.PolicyRule) bool {
if rule.Protocol == types.PolicyRuleProtocolNetbirdSSH {
return len(rule.AuthorizedGroups) == 0 && rule.AuthorizedUser == ""
}
return types.PolicyRuleImpliesLegacySSH(rule)
}
func (r *resolver) collectFromDNSSettings() {
if len(r.linkGroups) == 0 || r.snap.dnsSettings == nil {
return
@@ -85,6 +85,8 @@ func TestChangeIsEmpty(t *testing.T) {
assert.False(t, Change{Resources: []*resourceTypes.NetworkResource{{ID: "r"}}}.isEmpty())
assert.False(t, Change{Networks: []*networkTypes.Network{{ID: "n"}}}.isEmpty())
assert.False(t, Change{PostureCheckIDs: []string{"pc"}}.isEmpty())
assert.False(t, Change{UserGroupIDs: []string{"g"}}.isEmpty())
assert.False(t, Change{AllowedUsersChanged: true}.isEmpty())
}
func TestPolicyReferencesPostureChecks(t *testing.T) {
@@ -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{
+16 -1
View File
@@ -2,6 +2,7 @@ package proxy
import (
"context"
"errors"
"net"
"net/http"
"net/netip"
@@ -108,7 +109,7 @@ func (h *AuthCallbackHandler) handleCallback(w http.ResponseWriter, r *http.Requ
redirectURL.Scheme = "https"
query := redirectURL.Query()
query.Set("error", "access_denied")
query.Set("error_description", "Service configuration error")
query.Set("error_description", sessionTokenErrorDescription(err))
redirectURL.RawQuery = query.Encode()
http.Redirect(w, r, redirectURL.String(), http.StatusFound)
return
@@ -124,6 +125,20 @@ func (h *AuthCallbackHandler) handleCallback(w http.ResponseWriter, r *http.Requ
http.Redirect(w, r, redirectURL.String(), http.StatusFound)
}
// sessionTokenErrorDescription maps a session token failure to the text the
// proxy renders on its access denied page. Account status denials get a message
// the user can act on, while everything else stays generic so a lookup or
// signing failure does not describe management internals to the browser.
func sessionTokenErrorDescription(err error) string {
if errors.Is(err, nbgrpc.ErrUserPendingApproval) {
return "Your account is pending approval by an administrator"
}
if errors.Is(err, nbgrpc.ErrUserBlocked) {
return "Your account is blocked"
}
return "Service configuration error"
}
func extractUserIDFromToken(ctx context.Context, provider *oidc.Provider, config nbgrpc.ProxyOIDCConfig, token *oauth2.Token) string {
rawIDToken, ok := token.Extra("id_token").(string)
if !ok {
@@ -360,6 +360,51 @@ func createTestAccountsAndUsers(t *testing.T, ctx context.Context, testStore sto
Issued: "api",
}
require.NoError(t, testStore.SaveUser(ctx, allowedUser))
// A second tenant, whose users must never be issued a token signed with
// the first tenant's service session key.
otherAccount := &types.Account{
Id: "otherAccountId",
Domain: "other.com",
DomainCategory: "private",
IsDomainPrimaryAccount: true,
CreatedAt: time.Now(),
}
require.NoError(t, testStore.SaveAccount(ctx, otherAccount))
otherAccountUser := &types.User{
Id: "otherAccountUserId",
AccountID: "otherAccountId",
Role: types.UserRoleUser,
CreatedAt: time.Now(),
Issued: "api",
}
require.NoError(t, testStore.SaveUser(ctx, otherAccountUser))
// A user awaiting approval is stored as blocked and pending approval, and
// carries the same group membership as the approved one.
pendingUser := &types.User{
Id: "pendingUserId",
AccountID: "testAccountId",
Role: types.UserRoleUser,
AutoGroups: []string{"allowedGroupId"},
Blocked: true,
PendingApproval: true,
CreatedAt: time.Now(),
Issued: "api",
}
require.NoError(t, testStore.SaveUser(ctx, pendingUser))
blockedUser := &types.User{
Id: "blockedUserId",
AccountID: "testAccountId",
Role: types.UserRoleUser,
AutoGroups: []string{"allowedGroupId"},
Blocked: true,
CreatedAt: time.Now(),
Issued: "api",
}
require.NoError(t, testStore.SaveUser(ctx, blockedUser))
}
// testServiceManager is a minimal implementation for testing.
@@ -490,6 +535,64 @@ func TestAuthCallback_UserAllowedToLogin(t *testing.T) {
require.Empty(t, parsedLocation.Query().Get("error"), "Should not have error parameter")
}
// 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.
func TestAuthCallback_UserDeniedByAccountStatus(t *testing.T) {
tests := []struct {
name string
subject string
expectErrorDesc string
}{
{
name: "pending approval",
subject: "pendingUserId",
expectErrorDesc: "Your account is pending approval by an administrator",
},
{
name: "blocked",
subject: "blockedUserId",
expectErrorDesc: "Your account is blocked",
},
{
name: "unknown to management",
subject: "userMissingFromStoreId",
expectErrorDesc: "Service configuration error",
},
{
// The account topology stays out of the browser-visible message.
name: "belongs to another account",
subject: "otherAccountUserId",
expectErrorDesc: "Service configuration error",
},
}
for _, tt := range tests {
t.Run(tt.name, func(t *testing.T) {
setup := setupAuthCallbackTest(t)
defer setup.cleanup()
setup.oidcServer.tokenSubject = tt.subject
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)
parsedLocation, err := url.Parse(rec.Header().Get("Location"))
require.NoError(t, err)
require.Empty(t, parsedLocation.Query().Get("session_token"), "Denied user must not receive a session token")
require.Equal(t, "access_denied", parsedLocation.Query().Get("error"))
require.Equal(t, tt.expectErrorDesc, parsedLocation.Query().Get("error_description"))
})
}
}
func TestAuthCallback_ProxyNotFound(t *testing.T) {
setup := setupAuthCallbackTest(t)
defer setup.cleanup()
@@ -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")
}
+24
View File
@@ -6340,6 +6340,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.
@@ -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)
+8 -1
View File
@@ -328,6 +328,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 +361,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 +612,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)
},
}
}
+64 -6
View File
@@ -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()
+3
View File
@@ -89,6 +89,8 @@ type Account struct {
Onboarding AccountOnboarding `gorm:"foreignKey:AccountID;references:id;constraint:OnDelete:CASCADE"`
ReverseProxyFreeDomainNonce string
PostureValidation map[string]map[string]bool `gorm:"-"`
}
// this class is used by gorm only
@@ -789,6 +791,7 @@ func (a *Account) Copy() *Account {
Services: services,
Onboarding: a.Onboarding,
Domains: domains,
PostureValidation: a.PostureValidation,
}
}
@@ -106,3 +106,15 @@ func (a *Account) GetPeerNetworkMapComponents(
nmd := a.toNetworkMapData(accountZones, validatedPeersMap, resourcePolicies, routers, groupIDToUserIDs)
return nmd.GetPeerNetworkMapComponents(peerID, TwinCustomZone(peersCustomZone))
}
// PrecomputePostureValidation evaluates every posture check referenced by an enabled
// policy once and stores the results on the account, so the per-peer components
// calculations that follow look them up instead of re-evaluating checks for every
// peer pair. The evaluation itself runs on the twin store; every twin built from
// this account afterwards inherits the results. It must be called before the
// account is shared across goroutines.
func (a *Account) PrecomputePostureValidation(ctx context.Context) {
nmd := a.toNetworkMapData(nil, nil, nil, nil, nil)
nmd.PrecomputePostureValidation()
a.PostureValidation = nmd.PostureValidation
}
@@ -36,6 +36,7 @@ func (a *Account) toNetworkMapData(
Routers: make(map[string]map[string]*nmdata.NetworkRouter, len(routers)),
ValidatedPeers: validatedPeersMap,
GroupIDToUserIDs: groupIDToUserIDs,
PostureValidation: a.PostureValidation,
AllowedUserIDs: a.getAllowedUserIDs(),
NetworkXIDToPublicID: make(map[string]string, len(a.Networks)),
PostureCheckXIDToPublicID: make(map[string]string, len(a.PostureChecks)),
@@ -0,0 +1,72 @@
package types_test
import (
"context"
"testing"
"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
nbdns "github.com/netbirdio/netbird/dns"
"github.com/netbirdio/netbird/management/server/posture"
)
func TestPrecomputePostureValidation_MatchesDirectEvaluation(t *testing.T) {
account, validatedPeers := scalableTestAccount(60, 5)
account.PostureChecks = append(account.PostureChecks, &posture.Checks{
ID: "posture-check-strict", Name: "Strict version",
Checks: posture.ChecksDefinition{
NBVersionCheck: &posture.NBVersionCheck{MinVersion: "0.50.0"},
},
})
account.Policies[0].SourcePostureChecks = []string{"posture-check-ver", "posture-check-unknown"}
account.Policies[1].SourcePostureChecks = []string{"posture-check-strict"}
account.Policies[2].SourcePostureChecks = []string{"posture-check-ver"}
account.Policies[2].Enabled = false
ctx := context.Background()
resourcePolicies := account.GetResourcePoliciesMap()
routers := account.GetResourceRoutersMap()
type result struct {
peers map[string]struct{}
postureFailedPeers map[string]map[string]struct{}
}
snapshot := func() map[string]result {
results := make(map[string]result, len(account.Peers))
for peerID := range account.Peers {
components := account.GetPeerNetworkMapComponents(ctx, peerID, nbdns.CustomZone{}, nil, validatedPeers, resourcePolicies, routers, nil)
require.NotNil(t, components)
peerSet := make(map[string]struct{}, len(components.Peers))
for id := range components.Peers {
peerSet[id] = struct{}{}
}
results[peerID] = result{peers: peerSet, postureFailedPeers: components.PostureFailedPeers}
}
return results
}
direct := snapshot()
account.PrecomputePostureValidation(ctx)
memoized := snapshot()
require.Equal(t, len(direct), len(memoized))
for peerID, want := range direct {
got := memoized[peerID]
assert.Equal(t, want.peers, got.peers, "visible peers changed for %s", peerID)
assert.Equal(t, want.postureFailedPeers, got.postureFailedPeers, "posture failed peers changed for %s", peerID)
}
}
func TestPrecomputePostureValidation_NoPostureChecks(t *testing.T) {
account, validatedPeers := scalableTestAccount(10, 2)
account.PostureChecks = nil
ctx := context.Background()
account.PrecomputePostureValidation(ctx)
components := account.GetPeerNetworkMapComponents(ctx, "peer-0", nbdns.CustomZone{}, nil, validatedPeers, account.GetResourcePoliciesMap(), account.GetResourceRoutersMap(), nil)
require.NotNil(t, components)
assert.NotEmpty(t, components.Peers)
}
@@ -86,6 +86,43 @@ func BenchmarkNetworkMapGeneration_AllPeers(b *testing.B) {
b.ReportAllocs()
b.ResetTimer()
for range b.N {
account.PrecomputePostureValidation(ctx)
for _, peerID := range peerIDs {
_ = account.GetPeerNetworkMapFromComponents(ctx, peerID, nbdns.CustomZone{}, nil, validatedPeers, resourcePolicies, routers, nil, groupIDToUserIDs)
}
}
})
}
}
// BenchmarkNetworkMapGeneration_AllPeersPostureChecks benchmarks the UpdateAccountPeers
// hot path with a posture check attached to the account-wide policy, so posture
// validation runs for every source peer of every target peer's map.
func BenchmarkNetworkMapGeneration_AllPeersPostureChecks(b *testing.B) {
skipCIBenchmark(b)
scales := []benchmarkScale{
{"500peers_20groups", 500, 20},
{"1000peers_50groups", 1000, 50},
}
for _, scale := range scales {
account, validatedPeers := scalableTestAccount(scale.peers, scale.groups)
account.Policies[0].SourcePostureChecks = []string{"posture-check-ver"}
ctx := context.Background()
peerIDs := make([]string, 0, len(account.Peers))
for peerID := range account.Peers {
peerIDs = append(peerIDs, peerID)
}
b.Run("components/"+scale.name, func(b *testing.B) {
resourcePolicies := account.GetResourcePoliciesMap()
routers := account.GetResourceRoutersMap()
groupIDToUserIDs := account.GetActiveGroupUsers()
b.ReportAllocs()
b.ResetTimer()
for range b.N {
account.PrecomputePostureValidation(ctx)
for _, peerID := range peerIDs {
_ = account.GetPeerNetworkMapFromComponents(ctx, peerID, nbdns.CustomZone{}, nil, validatedPeers, resourcePolicies, routers, nil, groupIDToUserIDs)
}
+63 -17
View File
@@ -593,7 +593,8 @@ func (am *DefaultAccountManager) SaveOrAddUsers(ctx context.Context, accountID,
return nil, err
}
var updateAccountPeers bool
var snaps []*affectedpeers.Snapshot
var changes []affectedpeers.Change
var peersToExpire []*nbpeer.Peer
var addUserEvents []func()
var usersToSave = make([]*types.User, 0, len(updates))
@@ -629,20 +630,25 @@ func (am *DefaultAccountManager) SaveOrAddUsers(ctx context.Context, accountID,
}
err = am.Store.ExecuteInTransaction(ctx, func(transaction store.Store) error {
_, updatedUser, userPeersToExpire, userEvents, err := am.processUserUpdate(
change, updatedUser, userPeersToExpire, userEvents, err := am.processUserUpdate(
ctx, transaction, groupsMap, accountID, initiatorUserID, initiatorUser, update, addIfNotExists, settings,
)
if err != nil {
return fmt.Errorf("failed to process update for user %s: %w", update.Id, err)
}
updateAccountPeers = true
err = transaction.SaveUser(ctx, updatedUser)
if err != nil {
return fmt.Errorf("failed to save updated user %s: %w", update.Id, err)
}
snap, err := affectedpeers.Load(ctx, transaction, accountID, change)
if err != nil {
return err
}
snaps = append(snaps, snap)
changes = append(changes, change)
usersToSave = append(usersToSave, updatedUser)
addUserEvents = append(addUserEvents, userEvents...)
peersToExpire = append(peersToExpire, userPeersToExpire...)
@@ -683,11 +689,11 @@ func (am *DefaultAccountManager) SaveOrAddUsers(ctx context.Context, accountID,
log.WithContext(ctx).Errorf("failed update expired peers: %s", err)
return nil, err
}
} else if updateAccountPeers {
} else if len(usersToSave) > 0 {
if err = am.Store.IncrementNetworkSerial(ctx, accountID); err != nil {
return nil, fmt.Errorf("failed to increment network serial: %w", err)
}
am.UpdateAccountPeers(ctx, accountID, types.UpdateReason{Resource: types.UpdateResourceUser, Operation: types.UpdateOperationUpdate})
go am.dispatchAffected(ctx, accountID, snaps, changes)
}
return updatedUsersInfo, globalErr
@@ -759,19 +765,21 @@ func (am *DefaultAccountManager) prepareUserUpdateEvents(ctx context.Context, ac
}
func (am *DefaultAccountManager) processUserUpdate(ctx context.Context, transaction store.Store, groupsMap map[string]*types.Group,
accountID, initiatorUserId string, initiatorUser, update *types.User, addIfNotExists bool, settings *types.Settings) (bool, *types.User, []*nbpeer.Peer, []func(), error) {
accountID, initiatorUserId string, initiatorUser, update *types.User, addIfNotExists bool, settings *types.Settings) (affectedpeers.Change, *types.User, []*nbpeer.Peer, []func(), error) {
var change affectedpeers.Change
if update == nil {
return false, nil, nil, nil, status.Errorf(status.InvalidArgument, "provided user update is nil")
return change, nil, nil, nil, status.Errorf(status.InvalidArgument, "provided user update is nil")
}
oldUser, isNewUser, err := getUserOrCreateIfNotExists(ctx, transaction, accountID, update, addIfNotExists)
if err != nil {
return false, nil, nil, nil, err
return change, nil, nil, nil, err
}
if err := validateUserUpdate(groupsMap, initiatorUser, oldUser, update); err != nil {
return false, nil, nil, nil, err
return change, nil, nil, nil, err
}
// only auto groups, revoked status, and integration reference can be updated for now
@@ -792,13 +800,13 @@ func (am *DefaultAccountManager) processUserUpdate(ctx context.Context, transact
var transferredOwnerRole bool
result, err := handleOwnerRoleTransfer(ctx, transaction, initiatorUser, update)
if err != nil {
return false, nil, nil, nil, err
return change, nil, nil, nil, err
}
transferredOwnerRole = result
userPeers, err := transaction.GetUserPeers(ctx, store.LockingStrengthNone, updatedUser.AccountID, update.Id)
if err != nil {
return false, nil, nil, nil, err
return change, nil, nil, nil, err
}
var peersToExpire []*nbpeer.Peer
@@ -807,6 +815,32 @@ func (am *DefaultAccountManager) processUserUpdate(ctx context.Context, transact
peersToExpire = userPeers
}
// A user reaches a peer's network map only through the SSH rules: as part of a
// group -> user mapping, and as part of the account's allowed-user set. Creating,
// blocking or unblocking a user adds it to or removes it from both, so every group
// it maps into changes — including the All group that holds every active user.
// Otherwise only the auto-groups it joined or left do.
if isNewUser || oldUser.IsBlocked() != updatedUser.IsBlocked() {
change.AllowedUsersChanged = true
change.UserGroupIDs = slices.Concat(oldUser.AutoGroups, updatedUser.AutoGroups, allGroupIDs(groupsMap))
} else {
change.UserGroupIDs = slices.Concat(
util.Difference(oldUser.AutoGroups, updatedUser.AutoGroups),
util.Difference(updatedUser.AutoGroups, oldUser.AutoGroups),
)
}
// The user's peers are the changed entity in every scenario the update can
// produce — group membership, IPv6 assignment, SSH mappings — so they refresh
// together with every peer they can connect to, like on a regular peer update.
// An update that changes neither the auto-groups nor the active-user set has no
// peer-visible effect and refreshes nobody.
if len(change.UserGroupIDs) > 0 || change.AllowedUsersChanged {
for _, peer := range userPeers {
change.ChangedPeerIDs = append(change.ChangedPeerIDs, peer.ID)
}
}
var removedGroups, addedGroups []string
if update.AutoGroups != nil && settings.GroupsPropagationEnabled {
removedGroups = util.Difference(oldUser.AutoGroups, update.AutoGroups)
@@ -814,26 +848,38 @@ func (am *DefaultAccountManager) processUserUpdate(ctx context.Context, transact
for _, peer := range userPeers {
for _, groupID := range removedGroups {
if err := transaction.RemovePeerFromGroup(ctx, peer.ID, groupID); err != nil {
return false, nil, nil, nil, fmt.Errorf("failed to remove peer %s from group %s: %w", peer.ID, groupID, err)
return change, nil, nil, nil, fmt.Errorf("failed to remove peer %s from group %s: %w", peer.ID, groupID, err)
}
}
for _, groupID := range addedGroups {
if err := transaction.AddPeerToGroup(ctx, accountID, peer.ID, groupID); err != nil {
return false, nil, nil, nil, fmt.Errorf("failed to add peer %s to group %s: %w", peer.ID, groupID, err)
return change, nil, nil, nil, fmt.Errorf("failed to add peer %s to group %s: %w", peer.ID, groupID, err)
}
}
}
allGroupChanges := slices.Concat(removedGroups, addedGroups)
change.LinkGroups = allGroupChanges
if err := am.reconcileIPv6ForGroupChanges(ctx, transaction, accountID, allGroupChanges); err != nil {
return false, nil, nil, nil, fmt.Errorf("reconcile IPv6 for group changes: %w", err)
return change, nil, nil, nil, fmt.Errorf("reconcile IPv6 for group changes: %w", err)
}
}
updateAccountPeers := len(userPeers) > 0
userEventsToAdd := am.prepareUserUpdateEvents(ctx, updatedUser.AccountID, initiatorUserId, oldUser, updatedUser, transferredOwnerRole, isNewUser, removedGroups, addedGroups, transaction)
return updateAccountPeers, updatedUser, peersToExpire, userEventsToAdd, nil
return change, updatedUser, peersToExpire, userEventsToAdd, nil
}
// allGroupIDs returns the ID of the account's All group, which every active user maps
// into, as a slice so callers can concatenate it.
func allGroupIDs(groupsMap map[string]*types.Group) []string {
for _, group := range groupsMap {
if group.IsGroupAll() {
return []string{group.ID}
}
}
return nil
}
// getUserOrCreateIfNotExists retrieves the existing user or creates a new one if it doesn't exist.