mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-02 11:39:06 +02:00
Merge main into embedded-vnc
# Conflicts: # client/internal/engine.go
This commit is contained in:
@@ -17,6 +17,7 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
"github.com/netbirdio/netbird/management/server/account"
|
||||
nbcontext "github.com/netbirdio/netbird/management/server/context"
|
||||
"github.com/netbirdio/netbird/management/server/permissions"
|
||||
"github.com/netbirdio/netbird/management/server/store"
|
||||
@@ -61,10 +62,23 @@ func newAgentNetworkHandlerFixture(t *testing.T) *agentNetworkHandlerFixture {
|
||||
Return(true, context.Background(), nil).
|
||||
AnyTimes()
|
||||
|
||||
manager := agentnetwork.NewManager(st, perms, nil, nil)
|
||||
// Swallow activity events so the mutation paths (create/update/delete)
|
||||
// are exercisable through the HTTP layer.
|
||||
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()
|
||||
|
||||
manager := agentnetwork.NewManager(st, perms, accounts, nil)
|
||||
h := &handler{manager: manager}
|
||||
|
||||
router := mux.NewRouter()
|
||||
router.HandleFunc("/agent-network/providers", h.createProvider).Methods("POST")
|
||||
router.HandleFunc("/agent-network/providers/{providerId}", h.getProvider).Methods("GET")
|
||||
router.HandleFunc("/agent-network/providers/{providerId}", h.updateProvider).Methods("PUT")
|
||||
h.addPolicyEndpoints(router)
|
||||
h.addConsumptionEndpoints(router)
|
||||
h.addBudgetRuleEndpoints(router)
|
||||
|
||||
@@ -1,7 +1,9 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"math"
|
||||
nethttp "net/http"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
@@ -51,3 +53,50 @@ func TestValidate_ModelRates(t *testing.T) {
|
||||
assert.Error(t, validate(base(m), true), "case %q must be rejected", name)
|
||||
}
|
||||
}
|
||||
|
||||
// TestProviderHandler_UpdateReplacesFullState pins the update contract shared
|
||||
// with the other PUT endpoints: the request replaces the provider's mutable
|
||||
// state, so optional fields absent from the JSON land as their zero values.
|
||||
// The two exceptions are server-side: the api_key (a secret — omitted means
|
||||
// "not rotated") and the session keypair, both preserved by the manager. The
|
||||
// identity headers stay on the wire as explicit empty strings so a cleared
|
||||
// value round-trips.
|
||||
func TestProviderHandler_UpdateReplacesFullState(t *testing.T) {
|
||||
f := newAgentNetworkHandlerFixture(t)
|
||||
|
||||
create := `{
|
||||
"provider_id": "openai_api",
|
||||
"name": "openai",
|
||||
"upstream_url": "https://api.openai.com",
|
||||
"api_key": "sk-test",
|
||||
"enabled": true,
|
||||
"metadata_disabled": true,
|
||||
"skip_tls_verification": true,
|
||||
"extra_values": {"x-portkey-config": "pc-prod-3f2a"},
|
||||
"identity_header_user_id": "x-bf-dim-netbird_user_id",
|
||||
"models": [{"id": "gpt-4o", "input_per_1k": 0.0025, "output_per_1k": 0.01}]
|
||||
}`
|
||||
rec := f.do(t, nethttp.MethodPost, "/agent-network/providers", create)
|
||||
require.Equal(t, nethttp.StatusOK, rec.Code, "create must succeed: %s", rec.Body.String())
|
||||
|
||||
var created api.AgentNetworkProvider
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &created))
|
||||
|
||||
// Minimal update: only the required fields, no api_key. Everything
|
||||
// optional must land as its zero value.
|
||||
update := `{"provider_id": "openai_api", "name": "openai-renamed", "upstream_url": "https://api.openai.com", "enabled": true}`
|
||||
rec = f.do(t, nethttp.MethodPut, "/agent-network/providers/"+created.Id, update)
|
||||
require.Equal(t, nethttp.StatusOK, rec.Code, "update without api_key must succeed (key is preserved): %s", rec.Body.String())
|
||||
|
||||
var updated api.AgentNetworkProvider
|
||||
require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &updated))
|
||||
assert.Equal(t, "openai-renamed", updated.Name, "sent field must apply")
|
||||
assert.True(t, updated.Enabled, "sent field must apply")
|
||||
assert.False(t, updated.MetadataDisabled, "omitted metadata_disabled must land as false — PUT replaces the full state")
|
||||
assert.False(t, updated.SkipTlsVerification, "omitted skip_tls_verification must land as false")
|
||||
assert.Nil(t, updated.ExtraValues, "omitted extra_values must be cleared")
|
||||
assert.Equal(t, "", updated.IdentityHeaderUserId, "omitted identity header must be cleared yet stay on the wire")
|
||||
assert.Empty(t, updated.Models, "omitted models must be cleared")
|
||||
assert.Contains(t, rec.Body.String(), `"identity_header_user_id":""`,
|
||||
"cleared identity header must round-trip as an explicit empty string")
|
||||
}
|
||||
|
||||
@@ -2,7 +2,6 @@ package handlers
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
|
||||
"github.com/gorilla/mux"
|
||||
@@ -11,19 +10,20 @@ import (
|
||||
nbcontext "github.com/netbirdio/netbird/management/server/context"
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
"github.com/netbirdio/netbird/shared/management/http/util"
|
||||
"github.com/netbirdio/netbird/shared/management/status"
|
||||
)
|
||||
|
||||
// addSettingsEndpoints registers the Agent Network settings routes. The
|
||||
// settings row is bootstrapped server-side on first provider create; GET reads
|
||||
// it and PUT updates the mutable collection toggles (cluster/subdomain stay
|
||||
// immutable).
|
||||
// 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).
|
||||
func (h *handler) addSettingsEndpoints(router *mux.Router) {
|
||||
router.HandleFunc("/agent-network/settings", h.getSettings).Methods("GET", "OPTIONS")
|
||||
router.HandleFunc("/agent-network/settings", h.updateSettings).Methods("PUT", "OPTIONS")
|
||||
}
|
||||
|
||||
// updateSettings applies the collection toggles to the account's settings row.
|
||||
// updateSettings replaces the mutable settings fields on the account's row.
|
||||
// A request carrying a cluster bootstraps the row when the account doesn't
|
||||
// have one yet.
|
||||
func (h *handler) updateSettings(w http.ResponseWriter, r *http.Request) {
|
||||
userAuth, err := nbcontext.GetUserAuthFromContext(r.Context())
|
||||
if err != nil {
|
||||
@@ -48,11 +48,9 @@ func (h *handler) updateSettings(w http.ResponseWriter, r *http.Request) {
|
||||
util.WriteJSONObject(r.Context(), w, updated.ToAPIResponse())
|
||||
}
|
||||
|
||||
// getSettings returns the account's agent-network settings. The settings
|
||||
// row is bootstrapped on first provider create, so freshly-onboarded
|
||||
// accounts have nothing to read. Rather than 404-ing in that case (which
|
||||
// the dashboard would have to special-case), return a JSON null with 200
|
||||
// so consumers can branch on the body alone.
|
||||
// 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.
|
||||
func (h *handler) getSettings(w http.ResponseWriter, r *http.Request) {
|
||||
userAuth, err := nbcontext.GetUserAuthFromContext(r.Context())
|
||||
if err != nil {
|
||||
@@ -62,11 +60,6 @@ func (h *handler) getSettings(w http.ResponseWriter, r *http.Request) {
|
||||
|
||||
settings, err := h.manager.GetSettings(r.Context(), userAuth.AccountId, userAuth.UserId)
|
||||
if err != nil {
|
||||
var sErr *status.Error
|
||||
if errors.As(err, &sErr) && sErr.Type() == status.NotFound {
|
||||
util.WriteJSONObject(r.Context(), w, nil)
|
||||
return
|
||||
}
|
||||
util.WriteError(r.Context(), err, w)
|
||||
return
|
||||
}
|
||||
|
||||
@@ -0,0 +1,137 @@
|
||||
package handlers
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"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"
|
||||
// signal) and no timestamps — never a 404 and never the legacy null body.
|
||||
func TestSettingsHandler_GetUnbootstrappedReturnsDefaults(t *testing.T) {
|
||||
f := newAgentNetworkHandlerFixture(t)
|
||||
|
||||
rec := f.do(t, http.MethodGet, "/agent-network/settings", "")
|
||||
require.Equal(t, http.StatusOK, rec.Code,
|
||||
"unbootstrapped account must read as 200 with defaults: got %d body=%s", rec.Code, rec.Body.String())
|
||||
require.NotEqual(t, "null", trimSpace(rec.Body.String()),
|
||||
"the legacy 200+null shape must not come back")
|
||||
|
||||
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.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")
|
||||
require.NotNil(t, got.AccessLogRetentionDays)
|
||||
assert.Equal(t, 30, *got.AccessLogRetentionDays, "defaults must show the bootstrap retention")
|
||||
assert.Nil(t, got.CreatedAt, "no timestamps before a row exists")
|
||||
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) {
|
||||
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())
|
||||
|
||||
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.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")
|
||||
|
||||
// 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")
|
||||
}
|
||||
|
||||
// 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) {
|
||||
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())
|
||||
}
|
||||
|
||||
// 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.
|
||||
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())
|
||||
|
||||
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}`)
|
||||
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.True(t, got.EnableLogCollection, "sent toggle must apply")
|
||||
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")
|
||||
}
|
||||
|
||||
// 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) {
|
||||
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.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())
|
||||
|
||||
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())
|
||||
|
||||
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")
|
||||
}
|
||||
@@ -157,14 +157,14 @@ func NewManager(
|
||||
}
|
||||
|
||||
func (m *managerImpl) GetAllProviders(ctx context.Context, accountID, userID string) ([]*types.Provider, error) {
|
||||
if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil {
|
||||
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkProviders, operations.Read); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return m.store.GetAccountAgentNetworkProviders(ctx, store.LockingStrengthNone, accountID)
|
||||
}
|
||||
|
||||
func (m *managerImpl) GetProvider(ctx context.Context, accountID, userID, providerID string) (*types.Provider, error) {
|
||||
if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil {
|
||||
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkProviders, operations.Read); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return m.store.GetAgentNetworkProviderByID(ctx, store.LockingStrengthNone, accountID, providerID)
|
||||
@@ -175,9 +175,14 @@ func (m *managerImpl) GetProvider(ctx context.Context, accountID, userID, provid
|
||||
// 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) {
|
||||
if err := m.requirePermission(ctx, provider.AccountID, userID, operations.Create); err != nil {
|
||||
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
|
||||
@@ -202,7 +207,7 @@ func (m *managerImpl) CreateProvider(ctx context.Context, userID string, provide
|
||||
}
|
||||
|
||||
if strings.TrimSpace(bootstrapCluster) != "" {
|
||||
if _, err := m.bootstrapSettingsIfNeeded(ctx, provider.AccountID, bootstrapCluster); err != nil {
|
||||
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
|
||||
@@ -218,7 +223,7 @@ func (m *managerImpl) CreateProvider(ctx context.Context, userID string, provide
|
||||
}
|
||||
|
||||
func (m *managerImpl) UpdateProvider(ctx context.Context, userID string, provider *types.Provider) (*types.Provider, error) {
|
||||
if err := m.requirePermission(ctx, provider.AccountID, userID, operations.Update); err != nil {
|
||||
if err := m.requirePermission(ctx, provider.AccountID, userID, modules.AgentNetworkProviders, operations.Update); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -257,7 +262,7 @@ func (m *managerImpl) UpdateProvider(ctx context.Context, userID string, provide
|
||||
}
|
||||
|
||||
func (m *managerImpl) DeleteProvider(ctx context.Context, accountID, userID, providerID string) error {
|
||||
if err := m.requirePermission(ctx, accountID, userID, operations.Delete); err != nil {
|
||||
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkProviders, operations.Delete); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -306,21 +311,21 @@ func pluralize(n int, singular, plural string) string {
|
||||
}
|
||||
|
||||
func (m *managerImpl) GetAllPolicies(ctx context.Context, accountID, userID string) ([]*types.Policy, error) {
|
||||
if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil {
|
||||
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkPolicies, operations.Read); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return m.store.GetAccountAgentNetworkPolicies(ctx, store.LockingStrengthNone, accountID)
|
||||
}
|
||||
|
||||
func (m *managerImpl) GetPolicy(ctx context.Context, accountID, userID, policyID string) (*types.Policy, error) {
|
||||
if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil {
|
||||
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkPolicies, operations.Read); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return m.store.GetAgentNetworkPolicyByID(ctx, store.LockingStrengthNone, accountID, policyID)
|
||||
}
|
||||
|
||||
func (m *managerImpl) CreatePolicy(ctx context.Context, userID string, policy *types.Policy) (*types.Policy, error) {
|
||||
if err := m.requirePermission(ctx, policy.AccountID, userID, operations.Create); err != nil {
|
||||
if err := m.requirePermission(ctx, policy.AccountID, userID, modules.AgentNetworkPolicies, operations.Create); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -346,7 +351,7 @@ func (m *managerImpl) CreatePolicy(ctx context.Context, userID string, policy *t
|
||||
}
|
||||
|
||||
func (m *managerImpl) UpdatePolicy(ctx context.Context, userID string, policy *types.Policy) (*types.Policy, error) {
|
||||
if err := m.requirePermission(ctx, policy.AccountID, userID, operations.Update); err != nil {
|
||||
if err := m.requirePermission(ctx, policy.AccountID, userID, modules.AgentNetworkPolicies, operations.Update); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -373,7 +378,7 @@ func (m *managerImpl) UpdatePolicy(ctx context.Context, userID string, policy *t
|
||||
}
|
||||
|
||||
func (m *managerImpl) DeletePolicy(ctx context.Context, accountID, userID, policyID string) error {
|
||||
if err := m.requirePermission(ctx, accountID, userID, operations.Delete); err != nil {
|
||||
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkPolicies, operations.Delete); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -393,21 +398,21 @@ func (m *managerImpl) DeletePolicy(ctx context.Context, accountID, userID, polic
|
||||
}
|
||||
|
||||
func (m *managerImpl) GetAllGuardrails(ctx context.Context, accountID, userID string) ([]*types.Guardrail, error) {
|
||||
if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil {
|
||||
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkGuardrails, operations.Read); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return m.store.GetAccountAgentNetworkGuardrails(ctx, store.LockingStrengthNone, accountID)
|
||||
}
|
||||
|
||||
func (m *managerImpl) GetGuardrail(ctx context.Context, accountID, userID, guardrailID string) (*types.Guardrail, error) {
|
||||
if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil {
|
||||
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkGuardrails, operations.Read); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return m.store.GetAgentNetworkGuardrailByID(ctx, store.LockingStrengthNone, accountID, guardrailID)
|
||||
}
|
||||
|
||||
func (m *managerImpl) CreateGuardrail(ctx context.Context, userID string, guardrail *types.Guardrail) (*types.Guardrail, error) {
|
||||
if err := m.requirePermission(ctx, guardrail.AccountID, userID, operations.Create); err != nil {
|
||||
if err := m.requirePermission(ctx, guardrail.AccountID, userID, modules.AgentNetworkGuardrails, operations.Create); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -429,7 +434,7 @@ func (m *managerImpl) CreateGuardrail(ctx context.Context, userID string, guardr
|
||||
}
|
||||
|
||||
func (m *managerImpl) UpdateGuardrail(ctx context.Context, userID string, guardrail *types.Guardrail) (*types.Guardrail, error) {
|
||||
if err := m.requirePermission(ctx, guardrail.AccountID, userID, operations.Update); err != nil {
|
||||
if err := m.requirePermission(ctx, guardrail.AccountID, userID, modules.AgentNetworkGuardrails, operations.Update); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -452,7 +457,7 @@ func (m *managerImpl) UpdateGuardrail(ctx context.Context, userID string, guardr
|
||||
}
|
||||
|
||||
func (m *managerImpl) DeleteGuardrail(ctx context.Context, accountID, userID, guardrailID string) error {
|
||||
if err := m.requirePermission(ctx, accountID, userID, operations.Delete); err != nil {
|
||||
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkGuardrails, operations.Delete); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -473,7 +478,7 @@ func (m *managerImpl) DeleteGuardrail(ctx context.Context, accountID, userID, gu
|
||||
|
||||
// GetAllBudgetRules returns every account-level budget rule for the account.
|
||||
func (m *managerImpl) GetAllBudgetRules(ctx context.Context, accountID, userID string) ([]*types.AccountBudgetRule, error) {
|
||||
if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil {
|
||||
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkBudgets, operations.Read); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return m.store.GetAccountAgentNetworkBudgetRules(ctx, store.LockingStrengthNone, accountID)
|
||||
@@ -481,7 +486,7 @@ func (m *managerImpl) GetAllBudgetRules(ctx context.Context, accountID, userID s
|
||||
|
||||
// GetBudgetRule returns a single account-level budget rule.
|
||||
func (m *managerImpl) GetBudgetRule(ctx context.Context, accountID, userID, ruleID string) (*types.AccountBudgetRule, error) {
|
||||
if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil {
|
||||
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkBudgets, operations.Read); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return m.store.GetAgentNetworkBudgetRuleByID(ctx, store.LockingStrengthNone, accountID, ruleID)
|
||||
@@ -491,7 +496,7 @@ func (m *managerImpl) GetBudgetRule(ctx context.Context, accountID, userID, rule
|
||||
// enforced at request time (CheckLLMPolicyLimits), not baked into the synth
|
||||
// proxy config, so no reconcile is needed.
|
||||
func (m *managerImpl) CreateBudgetRule(ctx context.Context, userID string, rule *types.AccountBudgetRule) (*types.AccountBudgetRule, error) {
|
||||
if err := m.requirePermission(ctx, rule.AccountID, userID, operations.Create); err != nil {
|
||||
if err := m.requirePermission(ctx, rule.AccountID, userID, modules.AgentNetworkBudgets, operations.Create); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -513,7 +518,7 @@ func (m *managerImpl) CreateBudgetRule(ctx context.Context, userID string, rule
|
||||
|
||||
// UpdateBudgetRule updates an existing account-level budget rule.
|
||||
func (m *managerImpl) UpdateBudgetRule(ctx context.Context, userID string, rule *types.AccountBudgetRule) (*types.AccountBudgetRule, error) {
|
||||
if err := m.requirePermission(ctx, rule.AccountID, userID, operations.Update); err != nil {
|
||||
if err := m.requirePermission(ctx, rule.AccountID, userID, modules.AgentNetworkBudgets, operations.Update); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
@@ -536,7 +541,7 @@ func (m *managerImpl) UpdateBudgetRule(ctx context.Context, userID string, rule
|
||||
|
||||
// DeleteBudgetRule removes an account-level budget rule.
|
||||
func (m *managerImpl) DeleteBudgetRule(ctx context.Context, accountID, userID, ruleID string) error {
|
||||
if err := m.requirePermission(ctx, accountID, userID, operations.Delete); err != nil {
|
||||
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkBudgets, operations.Delete); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
@@ -554,40 +559,83 @@ func (m *managerImpl) DeleteBudgetRule(ctx context.Context, accountID, userID, r
|
||||
return nil
|
||||
}
|
||||
|
||||
// UpdateSettings applies the mutable account-level settings — the collection
|
||||
// toggles — onto the existing row. Cluster and Subdomain are immutable and are
|
||||
// preserved from the persisted row regardless of the input. 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.
|
||||
// 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.
|
||||
func (m *managerImpl) UpdateSettings(ctx context.Context, userID string, settings *types.Settings) (*types.Settings, error) {
|
||||
if err := m.requirePermission(ctx, settings.AccountID, userID, operations.Update); err != nil {
|
||||
if err := m.requirePermission(ctx, settings.AccountID, userID, modules.AgentNetworkSettings, operations.Update); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
existing, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthUpdate, settings.AccountID)
|
||||
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.
|
||||
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
|
||||
}
|
||||
default:
|
||||
return fmt.Errorf("get agent network settings: %w", err)
|
||||
}
|
||||
|
||||
existing.EnableLogCollection = settings.EnableLogCollection
|
||||
existing.EnablePromptCollection = settings.EnablePromptCollection
|
||||
existing.RedactPii = settings.RedactPii
|
||||
existing.AccessLogRetentionDays = settings.AccessLogRetentionDays
|
||||
existing.UpdatedAt = time.Now().UTC()
|
||||
|
||||
if err := tx.SaveAgentNetworkSettings(ctx, existing); err != nil {
|
||||
return fmt.Errorf("save agent network settings: %w", err)
|
||||
}
|
||||
updated = existing
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("get agent network settings: %w", err)
|
||||
}
|
||||
|
||||
existing.EnableLogCollection = settings.EnableLogCollection
|
||||
existing.EnablePromptCollection = settings.EnablePromptCollection
|
||||
existing.RedactPii = settings.RedactPii
|
||||
existing.AccessLogRetentionDays = settings.AccessLogRetentionDays
|
||||
existing.UpdatedAt = time.Now().UTC()
|
||||
|
||||
if err := m.store.SaveAgentNetworkSettings(ctx, existing); err != nil {
|
||||
return nil, fmt.Errorf("save agent network settings: %w", err)
|
||||
return nil, err
|
||||
}
|
||||
|
||||
m.accountManager.StoreEvent(ctx, userID, settings.AccountID, settings.AccountID, activity.AgentNetworkSettingsUpdated, map[string]any{
|
||||
"log_collection": existing.EnableLogCollection,
|
||||
"prompt_collection": existing.EnablePromptCollection,
|
||||
"redact_pii": existing.RedactPii,
|
||||
"log_collection": updated.EnableLogCollection,
|
||||
"prompt_collection": updated.EnablePromptCollection,
|
||||
"redact_pii": updated.RedactPii,
|
||||
})
|
||||
m.reconcile(ctx, settings.AccountID)
|
||||
|
||||
return existing, nil
|
||||
return updated, nil
|
||||
}
|
||||
|
||||
// isNotFound reports whether err is a status.NotFound error.
|
||||
func isNotFound(err error) bool {
|
||||
var sErr *status.Error
|
||||
return errors.As(err, &sErr) && sErr.Type() == status.NotFound
|
||||
}
|
||||
|
||||
// validateProviderRefs ensures every destination provider id refers to a
|
||||
@@ -611,14 +659,38 @@ func (m *managerImpl) validateProviderRefs(ctx context.Context, accountID string
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetSettings returns the agent-network settings row for the account.
|
||||
// Returns the underlying status.NotFound when no row has been
|
||||
// bootstrapped yet (i.e. the account has no providers).
|
||||
// GetSettings returns the agent-network settings row for the account. When no
|
||||
// row has been bootstrapped yet, the defaults are returned (without
|
||||
// persisting) with cluster and subdomain empty — settings always read as an
|
||||
// object, like the account and DNS settings endpoints.
|
||||
func (m *managerImpl) GetSettings(ctx context.Context, accountID, userID string) (*types.Settings, error) {
|
||||
if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil {
|
||||
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkSettings, operations.Read); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID)
|
||||
settings, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID)
|
||||
switch {
|
||||
case err == nil:
|
||||
return settings, nil
|
||||
case isNotFound(err):
|
||||
return types.DefaultSettings(accountID), nil
|
||||
default:
|
||||
return nil, err
|
||||
}
|
||||
}
|
||||
|
||||
// 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)
|
||||
}
|
||||
|
||||
// bootstrapSettingsIfNeeded creates the per-account agent-network
|
||||
@@ -626,8 +698,9 @@ func (m *managerImpl) GetSettings(ctx context.Context, accountID, userID string)
|
||||
// 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.
|
||||
func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, accountID, providerCluster string) (*types.Settings, error) {
|
||||
// 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")
|
||||
}
|
||||
@@ -635,16 +708,15 @@ func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, accountID,
|
||||
return nil, fmt.Errorf("bootstrap settings: provider cluster is required")
|
||||
}
|
||||
|
||||
existing, err := m.store.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID)
|
||||
existing, err := st.GetAgentNetworkSettings(ctx, store.LockingStrengthNone, accountID)
|
||||
if err == nil {
|
||||
return existing, nil
|
||||
}
|
||||
var sErr *status.Error
|
||||
if !errors.As(err, &sErr) || sErr.Type() != status.NotFound {
|
||||
if !isNotFound(err) {
|
||||
return nil, fmt.Errorf("get agent network settings: %w", err)
|
||||
}
|
||||
|
||||
siblings, err := m.store.GetAgentNetworkSettingsByCluster(ctx, store.LockingStrengthNone, providerCluster)
|
||||
siblings, err := st.GetAgentNetworkSettingsByCluster(ctx, store.LockingStrengthNone, providerCluster)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("list agent network settings on cluster: %w", err)
|
||||
}
|
||||
@@ -663,18 +735,12 @@ func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, accountID,
|
||||
m.labelRngMu.Unlock()
|
||||
|
||||
now := time.Now().UTC()
|
||||
settings := &types.Settings{
|
||||
AccountID: accountID,
|
||||
Cluster: providerCluster,
|
||||
Subdomain: subdomain,
|
||||
// Logs on by default; usage is collected regardless. Retention bounds
|
||||
// how long full log rows are kept.
|
||||
EnableLogCollection: true,
|
||||
AccessLogRetentionDays: types.DefaultAccessLogRetentionDays,
|
||||
CreatedAt: now,
|
||||
UpdatedAt: now,
|
||||
}
|
||||
if err := m.store.SaveAgentNetworkSettings(ctx, settings); err != nil {
|
||||
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)
|
||||
}
|
||||
return settings, nil
|
||||
@@ -685,7 +751,7 @@ func (m *managerImpl) bootstrapSettingsIfNeeded(ctx context.Context, accountID,
|
||||
// counter view; permission gate is the same Read role that gates
|
||||
// every other agent-network surface.
|
||||
func (m *managerImpl) ListConsumption(ctx context.Context, accountID, userID string) ([]*types.Consumption, error) {
|
||||
if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil {
|
||||
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkUsage, operations.Read); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
return m.store.ListAgentNetworkConsumption(ctx, store.LockingStrengthNone, accountID)
|
||||
@@ -694,7 +760,7 @@ func (m *managerImpl) ListConsumption(ctx context.Context, accountID, userID str
|
||||
// ListAccessLogs returns a paginated, server-side-filtered page of
|
||||
// agent-network access logs plus the total count matching the filter.
|
||||
func (m *managerImpl) ListAccessLogs(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter) ([]*types.AgentNetworkAccessLog, int64, error) {
|
||||
if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil {
|
||||
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkLogs, operations.Read); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
return m.store.GetAgentNetworkAccessLogs(ctx, store.LockingStrengthNone, accountID, filter)
|
||||
@@ -704,7 +770,7 @@ func (m *managerImpl) ListAccessLogs(ctx context.Context, accountID, userID stri
|
||||
// agent-network access logs grouped by session, plus the total number of
|
||||
// sessions matching the filter.
|
||||
func (m *managerImpl) ListAccessLogSessions(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter) ([]*types.AgentNetworkAccessLogSession, int64, error) {
|
||||
if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil {
|
||||
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkLogs, operations.Read); err != nil {
|
||||
return nil, 0, err
|
||||
}
|
||||
return m.store.GetAgentNetworkAccessLogSessions(ctx, store.LockingStrengthNone, accountID, filter)
|
||||
@@ -713,7 +779,7 @@ func (m *managerImpl) ListAccessLogSessions(ctx context.Context, accountID, user
|
||||
// GetUsageOverview returns the filtered usage rows aggregated into time buckets
|
||||
// at the requested granularity, oldest-first.
|
||||
func (m *managerImpl) GetUsageOverview(ctx context.Context, accountID, userID string, filter types.AgentNetworkAccessLogFilter, granularity types.UsageGranularity) ([]*types.AgentNetworkUsageBucket, error) {
|
||||
if err := m.requirePermission(ctx, accountID, userID, operations.Read); err != nil {
|
||||
if err := m.requirePermission(ctx, accountID, userID, modules.AgentNetworkUsage, operations.Read); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
rows, err := m.store.GetAgentNetworkUsageRows(ctx, store.LockingStrengthNone, accountID, filter)
|
||||
@@ -787,8 +853,8 @@ func (m *managerImpl) RecordConsumption(ctx context.Context, accountID string, k
|
||||
return m.store.IncrementAgentNetworkConsumption(ctx, accountID, kind, dimID, windowSeconds, windowStart, tokensIn, tokensOut, costUSD)
|
||||
}
|
||||
|
||||
func (m *managerImpl) requirePermission(ctx context.Context, accountID, userID string, op operations.Operation) error {
|
||||
ok, _, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, modules.AgentNetwork, op)
|
||||
func (m *managerImpl) requirePermission(ctx context.Context, accountID, userID string, module modules.Module, op operations.Operation) error {
|
||||
ok, _, err := m.permissionsManager.ValidateUserPermissions(ctx, accountID, userID, module, op)
|
||||
if err != nil {
|
||||
return status.NewPermissionValidationError(err)
|
||||
}
|
||||
@@ -877,8 +943,8 @@ func (*mockManager) UpdateBudgetRule(_ context.Context, _ string, r *types.Accou
|
||||
|
||||
func (*mockManager) DeleteBudgetRule(_ context.Context, _, _, _ string) error { return nil }
|
||||
|
||||
func (*mockManager) GetSettings(_ context.Context, _, _ string) (*types.Settings, error) {
|
||||
return nil, status.Errorf(status.NotFound, "agent network settings not found")
|
||||
func (*mockManager) GetSettings(_ context.Context, accountID, _ string) (*types.Settings, error) {
|
||||
return types.DefaultSettings(accountID), nil
|
||||
}
|
||||
|
||||
func (*mockManager) UpdateSettings(_ context.Context, _ string, s *types.Settings) (*types.Settings, error) {
|
||||
|
||||
@@ -0,0 +1,134 @@
|
||||
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")
|
||||
})
|
||||
}
|
||||
@@ -164,9 +164,7 @@ func (p *Provider) FromAPIRequest(req *api.AgentNetworkProviderRequest) {
|
||||
p.MetadataDisabled = *req.MetadataDisabled
|
||||
}
|
||||
// Identity-header overrides for catalogs flagged Customizable.
|
||||
// nil pointer = "field omitted on the wire" → leave the stored
|
||||
// value untouched (per the openapi description). Empty string is
|
||||
// an explicit clear that disables stamping for this dimension.
|
||||
// Empty or omitted disables stamping for this dimension.
|
||||
if req.IdentityHeaderUserId != nil {
|
||||
p.IdentityHeaderUserID = strings.TrimSpace(*req.IdentityHeaderUserId)
|
||||
}
|
||||
@@ -192,16 +190,20 @@ func (p *Provider) ToAPIResponse() *api.AgentNetworkProvider {
|
||||
created := p.CreatedAt
|
||||
updated := p.UpdatedAt
|
||||
resp := &api.AgentNetworkProvider{
|
||||
Id: p.ID,
|
||||
ProviderId: p.ProviderID,
|
||||
Name: p.Name,
|
||||
UpstreamUrl: p.UpstreamURL,
|
||||
Models: models,
|
||||
Enabled: p.Enabled,
|
||||
SkipTlsVerification: p.SkipTLSVerification,
|
||||
MetadataDisabled: p.MetadataDisabled,
|
||||
CreatedAt: &created,
|
||||
UpdatedAt: &updated,
|
||||
Id: p.ID,
|
||||
ProviderId: p.ProviderID,
|
||||
Name: p.Name,
|
||||
UpstreamUrl: p.UpstreamURL,
|
||||
Models: models,
|
||||
// Always present on the wire so an explicitly cleared header
|
||||
// round-trips as "" instead of vanishing from the response.
|
||||
IdentityHeaderUserId: p.IdentityHeaderUserID,
|
||||
IdentityHeaderGroups: p.IdentityHeaderGroups,
|
||||
Enabled: p.Enabled,
|
||||
SkipTlsVerification: p.SkipTLSVerification,
|
||||
MetadataDisabled: p.MetadataDisabled,
|
||||
CreatedAt: &created,
|
||||
UpdatedAt: &updated,
|
||||
}
|
||||
if len(p.ExtraValues) > 0 {
|
||||
out := make(map[string]string, len(p.ExtraValues))
|
||||
@@ -210,14 +212,6 @@ func (p *Provider) ToAPIResponse() *api.AgentNetworkProvider {
|
||||
}
|
||||
resp.ExtraValues = &out
|
||||
}
|
||||
if p.IdentityHeaderUserID != "" {
|
||||
v := p.IdentityHeaderUserID
|
||||
resp.IdentityHeaderUserId = &v
|
||||
}
|
||||
if p.IdentityHeaderGroups != "" {
|
||||
v := p.IdentityHeaderGroups
|
||||
resp.IdentityHeaderGroups = &v
|
||||
}
|
||||
return resp
|
||||
}
|
||||
|
||||
|
||||
@@ -77,3 +77,41 @@ func TestProvider_MetadataDisabled_RoundTrip(t *testing.T) {
|
||||
assert.False(t, p.MetadataDisabled, "explicit false must clear metadata_disabled")
|
||||
assert.False(t, p.ToAPIResponse().MetadataDisabled, "response must reflect the cleared value")
|
||||
}
|
||||
|
||||
// TestProvider_IdentityHeaders_AlwaysOnWire pins that the identity header
|
||||
// fields are always present in the API response — an explicitly cleared
|
||||
// ("") header must round-trip as "" rather than vanish, so API consumers
|
||||
// (e.g. the Terraform provider) never observe a value other than the one
|
||||
// they wrote.
|
||||
func TestProvider_IdentityHeaders_AlwaysOnWire(t *testing.T) {
|
||||
set := "x-bf-dim-netbird_user_id"
|
||||
empty := ""
|
||||
|
||||
base := func() *api.AgentNetworkProviderRequest {
|
||||
return &api.AgentNetworkProviderRequest{
|
||||
ProviderId: "custom",
|
||||
Name: "bifrost",
|
||||
UpstreamUrl: "https://bifrost.internal",
|
||||
}
|
||||
}
|
||||
|
||||
p := NewProvider("acc-1")
|
||||
resp := p.ToAPIResponse()
|
||||
assert.Equal(t, "", resp.IdentityHeaderUserId, "unset header must surface as empty string, not be omitted")
|
||||
assert.Equal(t, "", resp.IdentityHeaderGroups, "unset header must surface as empty string, not be omitted")
|
||||
|
||||
req := base()
|
||||
req.IdentityHeaderUserId = &set
|
||||
p.FromAPIRequest(req)
|
||||
assert.Equal(t, set, p.ToAPIResponse().IdentityHeaderUserId, "configured header must round-trip")
|
||||
|
||||
// Omitting the field preserves it.
|
||||
p.FromAPIRequest(base())
|
||||
assert.Equal(t, set, p.ToAPIResponse().IdentityHeaderUserId, "omitted header must preserve the stored value")
|
||||
|
||||
// An explicit "" clears it AND stays visible on the wire.
|
||||
req = base()
|
||||
req.IdentityHeaderUserId = &empty
|
||||
p.FromAPIRequest(req)
|
||||
assert.Equal(t, "", p.ToAPIResponse().IdentityHeaderUserId, "cleared header must round-trip as empty string")
|
||||
}
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package types
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/netbirdio/netbird/shared/management/http/api"
|
||||
@@ -42,18 +43,34 @@ type Settings struct {
|
||||
// schema cohesive.
|
||||
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.
|
||||
func DefaultSettings(accountID string) *Settings {
|
||||
return &Settings{
|
||||
AccountID: accountID,
|
||||
EnableLogCollection: true,
|
||||
AccessLogRetentionDays: DefaultAccessLogRetentionDays,
|
||||
}
|
||||
}
|
||||
|
||||
// Endpoint returns the bare hostname agents reach this account at:
|
||||
// `<subdomain>.<cluster>`.
|
||||
// `<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
|
||||
}
|
||||
|
||||
// ToAPIResponse renders the settings as the API representation.
|
||||
// ToAPIResponse renders the settings as the API representation. The
|
||||
// timestamps are omitted while zero — a default (not yet bootstrapped) view
|
||||
// has no persisted row to date.
|
||||
func (s *Settings) ToAPIResponse() *api.AgentNetworkSettings {
|
||||
created := s.CreatedAt
|
||||
updated := s.UpdatedAt
|
||||
retention := s.AccessLogRetentionDays
|
||||
return &api.AgentNetworkSettings{
|
||||
resp := &api.AgentNetworkSettings{
|
||||
Cluster: s.Cluster,
|
||||
Subdomain: s.Subdomain,
|
||||
Endpoint: s.Endpoint(),
|
||||
@@ -61,14 +78,27 @@ func (s *Settings) ToAPIResponse() *api.AgentNetworkSettings {
|
||||
EnablePromptCollection: s.EnablePromptCollection,
|
||||
RedactPii: s.RedactPii,
|
||||
AccessLogRetentionDays: &retention,
|
||||
CreatedAt: &created,
|
||||
UpdatedAt: &updated,
|
||||
}
|
||||
if !s.CreatedAt.IsZero() {
|
||||
created := s.CreatedAt
|
||||
resp.CreatedAt = &created
|
||||
}
|
||||
if !s.UpdatedAt.IsZero() {
|
||||
updated := s.UpdatedAt
|
||||
resp.UpdatedAt = &updated
|
||||
}
|
||||
return resp
|
||||
}
|
||||
|
||||
// FromAPIRequest applies the mutable settings fields from the request. Cluster
|
||||
// and Subdomain are immutable and intentionally not touched here.
|
||||
// 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.
|
||||
func (s *Settings) FromAPIRequest(req *api.AgentNetworkSettingsRequest) {
|
||||
if req.Cluster != nil {
|
||||
s.Cluster = strings.TrimSpace(*req.Cluster)
|
||||
}
|
||||
s.EnableLogCollection = req.EnableLogCollection
|
||||
s.EnablePromptCollection = req.EnablePromptCollection
|
||||
s.RedactPii = req.RedactPii
|
||||
|
||||
@@ -24,13 +24,13 @@ import (
|
||||
|
||||
"github.com/netbirdio/netbird/encryption"
|
||||
"github.com/netbirdio/netbird/formatter/hook"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs"
|
||||
accesslogsmanager "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/accesslogs/manager"
|
||||
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
nbgrpc "github.com/netbirdio/netbird/management/internals/shared/grpc"
|
||||
"github.com/netbirdio/netbird/management/server/activity"
|
||||
activitystore "github.com/netbirdio/netbird/management/server/activity/store"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/agentnetwork"
|
||||
nbcache "github.com/netbirdio/netbird/management/server/cache"
|
||||
nbContext "github.com/netbirdio/netbird/management/server/context"
|
||||
nbhttp "github.com/netbirdio/netbird/management/server/http"
|
||||
@@ -184,6 +184,10 @@ func (s *BaseServer) GRPCServer() *grpc.Server {
|
||||
grpc.ChainStreamInterceptor(realip.StreamServerInterceptorOpts(realipOpts...), streamInterceptor, proxyStream),
|
||||
}
|
||||
|
||||
// Append interceptors contributed by registered gRPC extensions. These
|
||||
// run after the built-in chain (ChainUnaryInterceptor is additive).
|
||||
gRPCOpts = appendExtensionInterceptors(gRPCOpts, s.grpcExtensions)
|
||||
|
||||
if s.Config.HttpConfig.LetsEncryptDomain != "" {
|
||||
certManager, err := encryption.CreateCertManager(s.Config.Datadir, s.Config.HttpConfig.LetsEncryptDomain)
|
||||
if err != nil {
|
||||
@@ -215,6 +219,9 @@ func (s *BaseServer) GRPCServer() *grpc.Server {
|
||||
mgmtProto.RegisterProxyServiceServer(gRPCAPIHandler, s.ReverseProxyGRPCServer())
|
||||
log.Info("ProxyService registered on gRPC server")
|
||||
|
||||
// Register services contributed by external modules via the extension seam.
|
||||
registerExtensions(gRPCAPIHandler, s.grpcExtensions)
|
||||
|
||||
return gRPCAPIHandler
|
||||
})
|
||||
}
|
||||
|
||||
@@ -0,0 +1,74 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"google.golang.org/grpc"
|
||||
)
|
||||
|
||||
// GRPCExtension bundles an external module's contribution to the management
|
||||
// gRPC server: the registration of one or more services onto the shared
|
||||
// grpc.Server, any server-wide interceptors those services require, and an
|
||||
// optional shutdown hook. It is a generic extension point with no knowledge of
|
||||
// any specific service.
|
||||
type GRPCExtension struct {
|
||||
// Register is invoked with the shared grpc.Server (as a ServiceRegistrar)
|
||||
// after the built-in services are registered. It may register any number of
|
||||
// services. May be nil.
|
||||
Register func(grpc.ServiceRegistrar)
|
||||
// UnaryInterceptors are appended to the server's unary interceptor chain,
|
||||
// running after the built-in interceptors. May be empty.
|
||||
UnaryInterceptors []grpc.UnaryServerInterceptor
|
||||
// StreamInterceptors are appended to the server's stream interceptor chain,
|
||||
// running after the built-in interceptors. May be empty.
|
||||
StreamInterceptors []grpc.StreamServerInterceptor
|
||||
// Shutdown, if non-nil, is called once during Stop() with the context
|
||||
// governing server shutdown, which carries a deadline. The hook MUST
|
||||
// return promptly and MUST abandon its work once that context is
|
||||
// cancelled or expires: it runs before the rest of Stop()'s cleanup
|
||||
// (store, event store, embedded IdP) and before Stop() itself checks the
|
||||
// context's deadline, so a hook that ignores the context will delay all
|
||||
// of that cleanup and prevent Stop() from returning on time. May be nil.
|
||||
Shutdown func(ctx context.Context)
|
||||
}
|
||||
|
||||
// RegisterGRPCExtension registers a gRPC extension. Call before the gRPC server
|
||||
// is first built (i.e. before Start); registrations after that have no effect.
|
||||
func (s *BaseServer) RegisterGRPCExtension(ext GRPCExtension) {
|
||||
s.grpcExtensions = append(s.grpcExtensions, ext)
|
||||
}
|
||||
|
||||
// appendExtensionInterceptors appends each extension's interceptors to the gRPC
|
||||
// server options as additional chained interceptors. grpc.ChainUnaryInterceptor
|
||||
// and grpc.ChainStreamInterceptor are additive, so the returned options run the
|
||||
// extension interceptors after any interceptors already present in opts.
|
||||
func appendExtensionInterceptors(opts []grpc.ServerOption, exts []GRPCExtension) []grpc.ServerOption {
|
||||
for _, ext := range exts {
|
||||
if len(ext.UnaryInterceptors) > 0 {
|
||||
opts = append(opts, grpc.ChainUnaryInterceptor(ext.UnaryInterceptors...))
|
||||
}
|
||||
if len(ext.StreamInterceptors) > 0 {
|
||||
opts = append(opts, grpc.ChainStreamInterceptor(ext.StreamInterceptors...))
|
||||
}
|
||||
}
|
||||
return opts
|
||||
}
|
||||
|
||||
// registerExtensions registers each extension's services onto reg.
|
||||
func registerExtensions(reg grpc.ServiceRegistrar, exts []GRPCExtension) {
|
||||
for _, ext := range exts {
|
||||
if ext.Register != nil {
|
||||
ext.Register(reg)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// runExtensionShutdownHooks calls each extension's shutdown hook, if set,
|
||||
// passing ctx through so hooks can honor its deadline/cancellation.
|
||||
func runExtensionShutdownHooks(ctx context.Context, exts []GRPCExtension) {
|
||||
for _, ext := range exts {
|
||||
if ext.Shutdown != nil {
|
||||
ext.Shutdown(ctx)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,160 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net"
|
||||
"sync/atomic"
|
||||
"testing"
|
||||
|
||||
"google.golang.org/grpc"
|
||||
"google.golang.org/grpc/credentials/insecure"
|
||||
"google.golang.org/grpc/health"
|
||||
healthgrpc "google.golang.org/grpc/health/grpc_health_v1"
|
||||
"google.golang.org/grpc/test/bufconn"
|
||||
)
|
||||
|
||||
// Test that an extension's interceptors and service registration are actually
|
||||
// wired onto a real in-process gRPC server via the helpers, and that shutdown
|
||||
// hooks run. This validates the load-bearing assumption that
|
||||
// grpc.ChainUnaryInterceptor is additive (extension interceptors run in
|
||||
// addition to any base chain).
|
||||
func TestGRPCExtensionAppliedToServer(t *testing.T) {
|
||||
var unaryCalls atomic.Int32
|
||||
var streamShutdownCalled atomic.Bool
|
||||
|
||||
ext := GRPCExtension{
|
||||
Register: func(reg grpc.ServiceRegistrar) {
|
||||
healthgrpc.RegisterHealthServer(reg, health.NewServer())
|
||||
},
|
||||
UnaryInterceptors: []grpc.UnaryServerInterceptor{
|
||||
func(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) {
|
||||
unaryCalls.Add(1)
|
||||
return handler(ctx, req)
|
||||
},
|
||||
},
|
||||
Shutdown: func(ctx context.Context) { streamShutdownCalled.Store(true) },
|
||||
}
|
||||
exts := []GRPCExtension{ext}
|
||||
|
||||
// Base options mimic GRPCServer(): a pre-existing chain the extension appends to.
|
||||
var baseUnaryCalls atomic.Int32
|
||||
opts := []grpc.ServerOption{
|
||||
grpc.ChainUnaryInterceptor(func(ctx context.Context, req any, info *grpc.UnaryServerInfo, handler grpc.UnaryHandler) (any, error) {
|
||||
baseUnaryCalls.Add(1)
|
||||
return handler(ctx, req)
|
||||
}),
|
||||
}
|
||||
opts = appendExtensionInterceptors(opts, exts)
|
||||
|
||||
srv := grpc.NewServer(opts...)
|
||||
registerExtensions(srv, exts)
|
||||
|
||||
lis := bufconn.Listen(1024 * 1024)
|
||||
go func() { _ = srv.Serve(lis) }()
|
||||
t.Cleanup(srv.Stop)
|
||||
|
||||
conn, err := grpc.NewClient("passthrough:///bufnet",
|
||||
grpc.WithContextDialer(func(ctx context.Context, _ string) (net.Conn, error) { return lis.DialContext(ctx) }),
|
||||
grpc.WithTransportCredentials(insecure.NewCredentials()))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = conn.Close() })
|
||||
|
||||
_, err = healthgrpc.NewHealthClient(conn).Check(context.Background(), &healthgrpc.HealthCheckRequest{})
|
||||
if err != nil {
|
||||
t.Fatalf("health check via extension-registered service failed: %v", err)
|
||||
}
|
||||
if baseUnaryCalls.Load() != 1 {
|
||||
t.Errorf("base interceptor calls = %d, want 1 (base chain must be preserved)", baseUnaryCalls.Load())
|
||||
}
|
||||
if unaryCalls.Load() != 1 {
|
||||
t.Errorf("extension interceptor calls = %d, want 1", unaryCalls.Load())
|
||||
}
|
||||
|
||||
runExtensionShutdownHooks(context.Background(), exts)
|
||||
if !streamShutdownCalled.Load() {
|
||||
t.Error("extension shutdown hook was not called")
|
||||
}
|
||||
}
|
||||
|
||||
// TestGRPCExtensionShutdownHookReceivesCallerContext asserts that each hook receives
|
||||
// a non-nil context and that it is the very same context the caller passed
|
||||
// in, so hooks can rely on values/deadlines placed on it by Stop().
|
||||
func TestGRPCExtensionShutdownHookReceivesCallerContext(t *testing.T) {
|
||||
type sentinelKey struct{}
|
||||
want := "shutdown-ctx-sentinel"
|
||||
ctx := context.WithValue(context.Background(), sentinelKey{}, want)
|
||||
|
||||
var called bool
|
||||
ext := GRPCExtension{
|
||||
Shutdown: func(hookCtx context.Context) {
|
||||
called = true
|
||||
if hookCtx == nil {
|
||||
t.Fatal("hook received a nil context")
|
||||
}
|
||||
got, _ := hookCtx.Value(sentinelKey{}).(string)
|
||||
if got != want {
|
||||
t.Errorf("hook context sentinel = %q, want %q (not the caller's context)", got, want)
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
runExtensionShutdownHooks(ctx, []GRPCExtension{ext})
|
||||
if !called {
|
||||
t.Fatal("shutdown hook was not called")
|
||||
}
|
||||
}
|
||||
|
||||
// TestGRPCExtensionShutdownHookObservesCancellation documents, by test, that
|
||||
// hooks can honor cancellation/deadlines: a hook given an already-cancelled
|
||||
// context must see ctx.Err() != nil and a closed Done() channel.
|
||||
func TestGRPCExtensionShutdownHookObservesCancellation(t *testing.T) {
|
||||
ctx, cancel := context.WithCancel(context.Background())
|
||||
cancel()
|
||||
|
||||
var called bool
|
||||
ext := GRPCExtension{
|
||||
Shutdown: func(hookCtx context.Context) {
|
||||
called = true
|
||||
if hookCtx.Err() == nil {
|
||||
t.Error("hook context Err() = nil, want non-nil for a cancelled context")
|
||||
}
|
||||
select {
|
||||
case <-hookCtx.Done():
|
||||
default:
|
||||
t.Error("hook context Done() channel is not closed for a cancelled context")
|
||||
}
|
||||
},
|
||||
}
|
||||
|
||||
runExtensionShutdownHooks(ctx, []GRPCExtension{ext})
|
||||
if !called {
|
||||
t.Fatal("shutdown hook was not called")
|
||||
}
|
||||
}
|
||||
|
||||
// TestGRPCExtensionShutdownHookNilSkipped asserts that an extension
|
||||
// with a nil Shutdown hook is skipped without panicking, and that hooks for
|
||||
// other extensions still run.
|
||||
func TestGRPCExtensionShutdownHookNilSkipped(t *testing.T) {
|
||||
var called atomic.Bool
|
||||
exts := []GRPCExtension{
|
||||
{Shutdown: nil},
|
||||
{Shutdown: func(context.Context) { called.Store(true) }},
|
||||
}
|
||||
|
||||
runExtensionShutdownHooks(context.Background(), exts)
|
||||
if !called.Load() {
|
||||
t.Error("shutdown hook for non-nil extension was not called")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterGRPCExtensionAccumulates(t *testing.T) {
|
||||
s := &BaseServer{}
|
||||
s.RegisterGRPCExtension(GRPCExtension{})
|
||||
s.RegisterGRPCExtension(GRPCExtension{})
|
||||
if len(s.grpcExtensions) != 2 {
|
||||
t.Fatalf("grpcExtensions len = %d, want 2", len(s.grpcExtensions))
|
||||
}
|
||||
}
|
||||
@@ -68,6 +68,11 @@ type BaseServer struct {
|
||||
|
||||
proxyAuthClose func()
|
||||
|
||||
// grpcExtensions holds additional gRPC services, interceptors, and shutdown
|
||||
// hooks registered by external modules via RegisterGRPCExtension. Populated
|
||||
// during boot (single-threaded), consumed by GRPCServer() and Stop().
|
||||
grpcExtensions []GRPCExtension
|
||||
|
||||
listener net.Listener
|
||||
certManager *autocert.Manager
|
||||
update *version.Update
|
||||
@@ -257,6 +262,7 @@ func (s *BaseServer) Stop() error {
|
||||
s.proxyAuthClose()
|
||||
s.proxyAuthClose = nil
|
||||
}
|
||||
runExtensionShutdownHooks(ctx, s.grpcExtensions)
|
||||
_ = s.Store().Close(ctx)
|
||||
_ = s.EventStore().Close(ctx)
|
||||
if s.update != nil {
|
||||
|
||||
@@ -61,6 +61,8 @@ func EncodeNetworkMapEnvelope(in ComponentsEnvelopeInput) *proto.NetworkMapEnvel
|
||||
return &proto.NetworkMapEnvelope{
|
||||
Payload: &proto.NetworkMapEnvelope_Full{
|
||||
Full: &proto.NetworkMapComponentsFull{
|
||||
Serial: networkSerial(c.Network),
|
||||
Network: toAccountNetwork(c.Network),
|
||||
PeerConfig: in.PeerConfig,
|
||||
// components.Peers always contains the target peer
|
||||
Peers: []*proto.PeerCompact{toPeerCompact(c.Peers[c.PeerID])},
|
||||
|
||||
@@ -758,6 +758,9 @@ func TestEncodeNetworkMapEnvelope_NilComponentsGracefulDegrade(t *testing.T) {
|
||||
assert.Equal(t, "netbird.cloud", full.DnsDomain)
|
||||
assert.Len(t, full.Peers, 1)
|
||||
assert.Empty(t, full.Policies)
|
||||
require.NotNil(t, full.Network, "client runs Calculate() over the envelope and dereferences Network unconditionally; a nil here would crash the receiver")
|
||||
assert.Equal(t, "net-empty", full.Network.Identifier)
|
||||
assert.Equal(t, uint64(9), full.Serial)
|
||||
}
|
||||
|
||||
func TestEncodeNetworkMapEnvelope_AccountSettingsAlwaysEmitted(t *testing.T) {
|
||||
@@ -776,6 +779,12 @@ func TestEncodeNetworkMapEnvelope_AccountSettingsAlwaysEmitted(t *testing.T) {
|
||||
func emptyNetworkMapComponents() *types.NetworkMapComponents {
|
||||
return types.EmptyNetworkMapComponents(
|
||||
&types.NetworkMapComponents{
|
||||
PeerID: "peer-id", Peers: map[string]*types.ComponentPeer{"peer-id": {}}},
|
||||
PeerID: "peer-id", Peers: map[string]*types.ComponentPeer{"peer-id": {}},
|
||||
Network: &types.Network{
|
||||
Identifier: "net-empty",
|
||||
Net: net.IPNet{IP: net.IP{100, 64, 0, 0}, Mask: net.CIDRMask(10, 32)},
|
||||
Serial: 9,
|
||||
},
|
||||
},
|
||||
)
|
||||
}
|
||||
|
||||
@@ -102,11 +102,20 @@ func TestAgentNetwork_UpdateSettings_PreservesImmutableAndTogglesCollection(t *t
|
||||
require.NotEmpty(t, before.Subdomain, "subdomain pinned at bootstrap")
|
||||
assert.False(t, before.EnablePromptCollection, "prompt collection defaults off")
|
||||
|
||||
// Attempt to flip toggles AND smuggle a different cluster/subdomain — the
|
||||
// immutable fields must be ignored.
|
||||
// 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")
|
||||
|
||||
// Flipping the toggles works with the pinned cluster echoed back (and
|
||||
// with it omitted); the subdomain is never taken from the request.
|
||||
updated, err := mgr.UpdateSettings(ctx, adminUserID, &agenttypes.Settings{
|
||||
AccountID: accountID,
|
||||
Cluster: "attacker.cluster",
|
||||
Cluster: clusterAddr,
|
||||
Subdomain: "evil",
|
||||
EnableLogCollection: true,
|
||||
EnablePromptCollection: true,
|
||||
|
||||
@@ -6,6 +6,8 @@ import (
|
||||
"fmt"
|
||||
"slices"
|
||||
|
||||
agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
"github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
"github.com/rs/xid"
|
||||
log "github.com/sirupsen/logrus"
|
||||
|
||||
@@ -744,6 +746,14 @@ func validateDeleteGroup(ctx context.Context, transaction store.Store, group *ty
|
||||
return &GroupLinkError{"network router", linkedRouter.ID}
|
||||
}
|
||||
|
||||
if isLinked, linkedService := isGroupLinkedToReverseProxyService(ctx, transaction, group.AccountID, group.ID); isLinked {
|
||||
return &GroupLinkError{"reverse proxy service", linkedService.Domain}
|
||||
}
|
||||
|
||||
if isLinked, linkedPolicy := isGroupLinkedToAgentNetworkPolicy(ctx, transaction, group.AccountID, group.ID); isLinked {
|
||||
return &GroupLinkError{"agent network policy", linkedPolicy.Name}
|
||||
}
|
||||
|
||||
return checkGroupLinkedToSettings(ctx, transaction, group)
|
||||
}
|
||||
|
||||
@@ -875,6 +885,46 @@ func isGroupLinkedToNetworkRouter(ctx context.Context, transaction store.Store,
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// isGroupLinkedToReverseProxyService checks if a group is used as an access group
|
||||
// of a private reverse proxy service or as a bearer-auth distribution group.
|
||||
func isGroupLinkedToReverseProxyService(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *service.Service) {
|
||||
services, err := transaction.GetAccountServices(ctx, store.LockingStrengthNone, accountID)
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Errorf("error retrieving reverse proxy services while checking group linkage: %v", err)
|
||||
return false, nil
|
||||
}
|
||||
|
||||
for _, svc := range services {
|
||||
if svc.Private && slices.Contains(svc.AccessGroups, groupID) {
|
||||
return true, svc
|
||||
}
|
||||
if svc.Auth.BearerAuth != nil && svc.Auth.BearerAuth.Enabled && slices.Contains(svc.Auth.BearerAuth.DistributionGroups, groupID) {
|
||||
return true, svc
|
||||
}
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// isGroupLinkedToAgentNetworkPolicy checks if a group is used as a source group by any
|
||||
// agent network policy in the account.
|
||||
func isGroupLinkedToAgentNetworkPolicy(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *agentNetworkTypes.Policy) {
|
||||
policies, err := transaction.GetAccountAgentNetworkPolicies(ctx, store.LockingStrengthNone, accountID)
|
||||
if err != nil {
|
||||
log.WithContext(ctx).Errorf("error retrieving agent network policies while checking group linkage: %v", err)
|
||||
return false, nil
|
||||
}
|
||||
|
||||
for _, policy := range policies {
|
||||
if policy == nil {
|
||||
continue
|
||||
}
|
||||
if slices.Contains(policy.SourceGroups, groupID) {
|
||||
return true, policy
|
||||
}
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
// areGroupChangesAffectPeers checks if any changes to the specified groups will affect peers.
|
||||
// It fetches each collection once and checks all groupIDs against them in memory.
|
||||
func areGroupChangesAffectPeers(ctx context.Context, transaction store.Store, accountID string, groupIDs []string) (bool, error) {
|
||||
|
||||
@@ -18,6 +18,8 @@ import (
|
||||
"golang.org/x/exp/maps"
|
||||
|
||||
nbdns "github.com/netbirdio/netbird/dns"
|
||||
agentNetworkTypes "github.com/netbirdio/netbird/management/internals/modules/agentnetwork/types"
|
||||
rpservice "github.com/netbirdio/netbird/management/internals/modules/reverseproxy/service"
|
||||
"github.com/netbirdio/netbird/management/server/groups"
|
||||
"github.com/netbirdio/netbird/management/server/networks"
|
||||
"github.com/netbirdio/netbird/management/server/networks/resources"
|
||||
@@ -125,6 +127,21 @@ func TestDefaultAccountManager_DeleteGroup(t *testing.T) {
|
||||
"grp-for-integration",
|
||||
"only service users with admin power can delete integration group",
|
||||
},
|
||||
{
|
||||
"agent network policy",
|
||||
"grp-for-agent-network-policy",
|
||||
"agent network policy",
|
||||
},
|
||||
{
|
||||
"reverse proxy private service access group",
|
||||
"grp-for-rp-private",
|
||||
"reverse proxy service",
|
||||
},
|
||||
{
|
||||
"reverse proxy bearer distribution group",
|
||||
"grp-for-rp-bearer",
|
||||
"reverse proxy service",
|
||||
},
|
||||
}
|
||||
|
||||
for _, testCase := range testCases {
|
||||
@@ -218,6 +235,17 @@ func TestDefaultAccountManager_DeleteGroups(t *testing.T) {
|
||||
groupIDs: []string{"grp-for-integration"},
|
||||
expectedReasons: []string{"only service users with admin power can delete integration group"},
|
||||
},
|
||||
{
|
||||
name: "agent network policy",
|
||||
groupIDs: []string{"grp-for-agent-network-policy"},
|
||||
expectedReasons: []string{"agent network policy"},
|
||||
},
|
||||
{
|
||||
name: "reverse proxy services",
|
||||
groupIDs: []string{"grp-for-rp-private", "grp-for-rp-bearer"},
|
||||
expectedReasons: []string{"reverse proxy service", "reverse proxy service"},
|
||||
expectedNotDeleted: []string{"grp-for-rp-private", "grp-for-rp-bearer"},
|
||||
},
|
||||
{
|
||||
name: "successfully delete multiple groups",
|
||||
groupIDs: []string{"group-1", "group-2"},
|
||||
@@ -285,6 +313,65 @@ func TestDefaultAccountManager_DeleteGroups(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultAccountManager_DeleteGroupUnlinkedFromReverseProxyService(t *testing.T) {
|
||||
am, _, err := createManager(t)
|
||||
require.NoError(t, err, "Failed to create account manager")
|
||||
|
||||
_, account, err := initTestGroupAccount(am)
|
||||
require.NoError(t, err, "Failed to init testing account")
|
||||
|
||||
deletableGroups := []*types.Group{
|
||||
{
|
||||
ID: "grp-rp-bearer-disabled",
|
||||
AccountID: account.Id,
|
||||
Name: "Group only in a disabled bearer auth",
|
||||
Issued: types.GroupIssuedAPI,
|
||||
Peers: make([]string, 0),
|
||||
},
|
||||
{
|
||||
ID: "grp-rp-nonprivate-access",
|
||||
AccountID: account.Id,
|
||||
Name: "Group only in a non-private service's access groups",
|
||||
Issued: types.GroupIssuedAPI,
|
||||
Peers: make([]string, 0),
|
||||
},
|
||||
}
|
||||
for _, group := range deletableGroups {
|
||||
require.NoError(t, am.CreateGroup(context.Background(), account.Id, groupAdminUserID, group))
|
||||
}
|
||||
|
||||
// Disabled bearer auth and stale access groups on a non-private service
|
||||
// are inert configuration and must not block group deletion.
|
||||
services := []*rpservice.Service{
|
||||
{
|
||||
ID: "rp-svc-bearer-disabled",
|
||||
AccountID: account.Id,
|
||||
Domain: "bearer-disabled.services.example.com",
|
||||
Auth: rpservice.AuthConfig{
|
||||
BearerAuth: &rpservice.BearerAuthConfig{
|
||||
Enabled: false,
|
||||
DistributionGroups: []string{"grp-rp-bearer-disabled"},
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
ID: "rp-svc-nonprivate-access",
|
||||
AccountID: account.Id,
|
||||
Domain: "nonprivate.services.example.com",
|
||||
Private: false,
|
||||
AccessGroups: []string{"grp-rp-nonprivate-access"},
|
||||
},
|
||||
}
|
||||
for _, svc := range services {
|
||||
require.NoError(t, am.Store.CreateService(context.Background(), svc))
|
||||
}
|
||||
|
||||
for _, group := range deletableGroups {
|
||||
err = am.DeleteGroup(context.Background(), account.Id, groupAdminUserID, group.ID)
|
||||
assert.NoError(t, err, "group %s is not referenced by an active reverse proxy gate and should be deletable", group.ID)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDefaultAccountManager_DeleteGroupLinkedToFlowGroup(t *testing.T) {
|
||||
am, _, err := createManager(t)
|
||||
require.NoError(t, err)
|
||||
@@ -406,6 +493,30 @@ func initTestGroupAccount(am *DefaultAccountManager) (*DefaultAccountManager, *t
|
||||
Peers: make([]string, 0),
|
||||
}
|
||||
|
||||
groupForAgentNetworkPolicy := &types.Group{
|
||||
ID: "grp-for-agent-network-policy",
|
||||
AccountID: "account-id",
|
||||
Name: "Group for agent network policies",
|
||||
Issued: types.GroupIssuedAPI,
|
||||
Peers: make([]string, 0),
|
||||
}
|
||||
|
||||
groupForRPPrivate := &types.Group{
|
||||
ID: "grp-for-rp-private",
|
||||
AccountID: "account-id",
|
||||
Name: "Group for private reverse proxy service",
|
||||
Issued: types.GroupIssuedAPI,
|
||||
Peers: make([]string, 0),
|
||||
}
|
||||
|
||||
groupForRPBearer := &types.Group{
|
||||
ID: "grp-for-rp-bearer",
|
||||
AccountID: "account-id",
|
||||
Name: "Group for bearer reverse proxy service",
|
||||
Issued: types.GroupIssuedAPI,
|
||||
Peers: make([]string, 0),
|
||||
}
|
||||
|
||||
routeResource := &route.Route{
|
||||
ID: "example route",
|
||||
Groups: []string{groupForRoute.ID},
|
||||
@@ -461,6 +572,66 @@ func initTestGroupAccount(am *DefaultAccountManager) (*DefaultAccountManager, *t
|
||||
_ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForSetupKeys)
|
||||
_ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForUsers)
|
||||
_ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForIntegration)
|
||||
_ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForAgentNetworkPolicy)
|
||||
_ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForRPPrivate)
|
||||
_ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForRPBearer)
|
||||
|
||||
agentNetworkPolicy := &agentNetworkTypes.Policy{
|
||||
ID: "example agent network policy",
|
||||
AccountID: accountID,
|
||||
Name: "Example agent network policy",
|
||||
Enabled: true,
|
||||
SourceGroups: []string{groupForAgentNetworkPolicy.ID},
|
||||
}
|
||||
if err := am.Store.SaveAgentNetworkPolicy(context.Background(), agentNetworkPolicy); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
|
||||
// The decoy services are created first so the linkage check has to scan
|
||||
// past services that do not reference the groups under test.
|
||||
rpServices := []*rpservice.Service{
|
||||
{
|
||||
ID: "rp-svc-private-decoy",
|
||||
AccountID: accountID,
|
||||
Domain: "private-decoy.services.example.com",
|
||||
Private: true,
|
||||
AccessGroups: []string{"unrelated-group"},
|
||||
},
|
||||
{
|
||||
ID: "rp-svc-bearer-decoy",
|
||||
AccountID: accountID,
|
||||
Domain: "bearer-decoy.services.example.com",
|
||||
Auth: rpservice.AuthConfig{
|
||||
BearerAuth: &rpservice.BearerAuthConfig{
|
||||
Enabled: true,
|
||||
DistributionGroups: []string{"unrelated-group"},
|
||||
},
|
||||
},
|
||||
},
|
||||
{
|
||||
ID: "rp-svc-private",
|
||||
AccountID: accountID,
|
||||
Domain: "private.services.example.com",
|
||||
Private: true,
|
||||
AccessGroups: []string{groupForRPPrivate.ID},
|
||||
},
|
||||
{
|
||||
ID: "rp-svc-bearer",
|
||||
AccountID: accountID,
|
||||
Domain: "bearer.services.example.com",
|
||||
Auth: rpservice.AuthConfig{
|
||||
BearerAuth: &rpservice.BearerAuthConfig{
|
||||
Enabled: true,
|
||||
DistributionGroups: []string{groupForRPBearer.ID},
|
||||
},
|
||||
},
|
||||
},
|
||||
}
|
||||
for _, svc := range rpServices {
|
||||
if err := am.Store.CreateService(context.Background(), svc); err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
}
|
||||
|
||||
acc, err := am.Store.GetAccount(context.Background(), account.Id)
|
||||
if err != nil {
|
||||
|
||||
@@ -82,6 +82,9 @@ func (m *managerImpl) ValidateUserPermissions(
|
||||
return m.ValidateRoleModuleAccess(ctx, accountID, role, module, operation), ctxEnriched, nil
|
||||
}
|
||||
|
||||
// ValidateRoleModuleAccess resolves an operation against the role's explicit
|
||||
// grant for the module, then the grant for its parent module when the module
|
||||
// is a dotted submodule, and finally the role's AutoAllowNew default.
|
||||
func (m *managerImpl) ValidateRoleModuleAccess(
|
||||
ctx context.Context,
|
||||
accountID string,
|
||||
@@ -89,7 +92,7 @@ func (m *managerImpl) ValidateRoleModuleAccess(
|
||||
module modules.Module,
|
||||
operation operations.Operation,
|
||||
) bool {
|
||||
if permissions, ok := role.Permissions[module]; ok {
|
||||
if permissions, ok := lookupModulePermissions(role, module); ok {
|
||||
if allowed, exists := permissions[operation]; exists {
|
||||
return allowed
|
||||
}
|
||||
@@ -100,6 +103,21 @@ func (m *managerImpl) ValidateRoleModuleAccess(
|
||||
return role.AutoAllowNew[operation]
|
||||
}
|
||||
|
||||
// lookupModulePermissions returns the role's explicit permission set for the
|
||||
// module, falling back to the parent module's set for dotted submodules. The
|
||||
// second return reports whether any explicit set was found.
|
||||
func lookupModulePermissions(role roles.RolePermissions, module modules.Module) (map[operations.Operation]bool, bool) {
|
||||
if permissions, ok := role.Permissions[module]; ok {
|
||||
return permissions, true
|
||||
}
|
||||
if parent, hasParent := module.Parent(); hasParent {
|
||||
if permissions, ok := role.Permissions[parent]; ok {
|
||||
return permissions, true
|
||||
}
|
||||
}
|
||||
return nil, false
|
||||
}
|
||||
|
||||
func (m *managerImpl) ValidateAccountAccess(ctx context.Context, accountID string, user *types.User, allowOwnerAndAdmin bool) (context.Context, error) {
|
||||
if user.AccountID != accountID {
|
||||
return ctx, status.NewUserNotPartOfAccountError()
|
||||
@@ -119,7 +137,7 @@ func (m *managerImpl) GetPermissionsByRole(ctx context.Context, role types.UserR
|
||||
permissions := roles.Permissions{}
|
||||
|
||||
for k := range modules.All {
|
||||
if rolePermissions, ok := roleMap.Permissions[k]; ok {
|
||||
if rolePermissions, ok := lookupModulePermissions(roleMap, k); ok {
|
||||
permissions[k] = rolePermissions
|
||||
continue
|
||||
}
|
||||
|
||||
@@ -0,0 +1,139 @@
|
||||
package permissions
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/stretchr/testify/assert"
|
||||
"github.com/stretchr/testify/require"
|
||||
|
||||
"github.com/netbirdio/netbird/management/server/permissions/modules"
|
||||
"github.com/netbirdio/netbird/management/server/permissions/operations"
|
||||
"github.com/netbirdio/netbird/management/server/permissions/roles"
|
||||
"github.com/netbirdio/netbird/management/server/types"
|
||||
)
|
||||
|
||||
func TestValidateRoleModuleAccessSubmoduleCascade(t *testing.T) {
|
||||
manager := NewManager(nil)
|
||||
ctx := context.Background()
|
||||
|
||||
fullAccess := map[operations.Operation]bool{
|
||||
operations.Read: true,
|
||||
operations.Create: true,
|
||||
operations.Update: true,
|
||||
operations.Delete: true,
|
||||
}
|
||||
readOnly := map[operations.Operation]bool{
|
||||
operations.Read: true,
|
||||
operations.Create: false,
|
||||
operations.Update: false,
|
||||
operations.Delete: false,
|
||||
}
|
||||
denyAll := map[operations.Operation]bool{
|
||||
operations.Read: false,
|
||||
operations.Create: false,
|
||||
operations.Update: false,
|
||||
operations.Delete: false,
|
||||
}
|
||||
|
||||
t.Run("parent grant covers submodules", func(t *testing.T) {
|
||||
role := roles.RolePermissions{
|
||||
AutoAllowNew: denyAll,
|
||||
Permissions: roles.Permissions{modules.AgentNetwork: fullAccess},
|
||||
}
|
||||
assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkProviders, operations.Create),
|
||||
"parent full grant should allow create on a submodule")
|
||||
assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkLogs, operations.Read),
|
||||
"parent full grant should allow read on a submodule")
|
||||
})
|
||||
|
||||
t.Run("submodule grant does not leak to parent or siblings", func(t *testing.T) {
|
||||
role := roles.RolePermissions{
|
||||
AutoAllowNew: denyAll,
|
||||
Permissions: roles.Permissions{modules.AgentNetworkUsage: readOnly},
|
||||
}
|
||||
assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkUsage, operations.Read),
|
||||
"explicit submodule read should be allowed")
|
||||
assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkUsage, operations.Create),
|
||||
"read-only submodule grant should not allow create")
|
||||
assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetwork, operations.Read),
|
||||
"submodule grant should not grant the parent module")
|
||||
assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkProviders, operations.Read),
|
||||
"submodule grant should not grant a sibling submodule")
|
||||
})
|
||||
|
||||
t.Run("explicit submodule entry wins over parent grant", func(t *testing.T) {
|
||||
role := roles.RolePermissions{
|
||||
AutoAllowNew: denyAll,
|
||||
Permissions: roles.Permissions{
|
||||
modules.AgentNetwork: fullAccess,
|
||||
modules.AgentNetworkLogs: denyAll,
|
||||
},
|
||||
}
|
||||
assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkLogs, operations.Read),
|
||||
"explicit submodule deny should override the parent grant")
|
||||
assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkUsage, operations.Read),
|
||||
"sibling submodules should still resolve through the parent grant")
|
||||
})
|
||||
|
||||
t.Run("auto allow applies when neither submodule nor parent is granted", func(t *testing.T) {
|
||||
role := roles.RolePermissions{
|
||||
AutoAllowNew: readOnly,
|
||||
}
|
||||
assert.True(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkProviders, operations.Read),
|
||||
"auto-allow read should apply to submodules")
|
||||
assert.False(t, manager.ValidateRoleModuleAccess(ctx, "account", role, modules.AgentNetworkProviders, operations.Delete),
|
||||
"auto-allow should not grant unlisted operations")
|
||||
})
|
||||
}
|
||||
|
||||
// TestExistingRolesKeepAgentNetworkBehaviorOnSubmodules pins the behavior the
|
||||
// submodule split must not change: every built-in role resolves the new
|
||||
// submodules exactly as it resolved the agent_network module before.
|
||||
func TestExistingRolesKeepAgentNetworkBehaviorOnSubmodules(t *testing.T) {
|
||||
manager := NewManager(nil)
|
||||
ctx := context.Background()
|
||||
|
||||
submodules := []modules.Module{
|
||||
modules.AgentNetworkProviders,
|
||||
modules.AgentNetworkPolicies,
|
||||
modules.AgentNetworkGuardrails,
|
||||
modules.AgentNetworkBudgets,
|
||||
modules.AgentNetworkUsage,
|
||||
modules.AgentNetworkLogs,
|
||||
modules.AgentNetworkSettings,
|
||||
}
|
||||
allOperations := []operations.Operation{operations.Read, operations.Create, operations.Update, operations.Delete}
|
||||
|
||||
for _, role := range []types.UserRole{types.UserRoleOwner, types.UserRoleAdmin, types.UserRoleAuditor, types.UserRoleNetworkAdmin, types.UserRoleUser} {
|
||||
rolePermissions, ok := roles.RolesMap[role]
|
||||
require.True(t, ok, "role %s must exist in RolesMap", role)
|
||||
|
||||
for _, sub := range submodules {
|
||||
for _, op := range allOperations {
|
||||
expected := manager.ValidateRoleModuleAccess(ctx, "account", rolePermissions, modules.AgentNetwork, op)
|
||||
actual := manager.ValidateRoleModuleAccess(ctx, "account", rolePermissions, sub, op)
|
||||
assert.Equal(t, expected, actual, "role %s: %s on %s should match the agent_network module", role, op, sub)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestGetPermissionsByRoleIncludesSubmodules(t *testing.T) {
|
||||
manager := NewManager(nil)
|
||||
ctx := context.Background()
|
||||
|
||||
permissions, err := manager.GetPermissionsByRole(ctx, types.UserRoleAuditor)
|
||||
require.NoError(t, err, "auditor role must resolve")
|
||||
|
||||
usage, ok := permissions[modules.AgentNetworkUsage]
|
||||
require.True(t, ok, "permissions map should contain the usage submodule")
|
||||
assert.True(t, usage[operations.Read], "auditor should read the usage submodule")
|
||||
assert.False(t, usage[operations.Update], "auditor should not update the usage submodule")
|
||||
|
||||
adminPermissions, err := manager.GetPermissionsByRole(ctx, types.UserRoleAdmin)
|
||||
require.NoError(t, err, "admin role must resolve")
|
||||
providers, ok := adminPermissions[modules.AgentNetworkProviders]
|
||||
require.True(t, ok, "permissions map should contain the providers submodule")
|
||||
assert.True(t, providers[operations.Delete], "admin should delete on the providers submodule")
|
||||
}
|
||||
@@ -1,5 +1,7 @@
|
||||
package modules
|
||||
|
||||
import "strings"
|
||||
|
||||
type Module string
|
||||
|
||||
const (
|
||||
@@ -20,6 +22,17 @@ const (
|
||||
IdentityProviders Module = "identity_providers"
|
||||
Services Module = "services"
|
||||
AgentNetwork Module = "agent_network"
|
||||
|
||||
// Agent Network submodules. A role may grant one of these directly
|
||||
// or grant the AgentNetwork parent, which covers all of them (see
|
||||
// permissions.Manager cascade resolution).
|
||||
AgentNetworkProviders Module = "agent_network.providers"
|
||||
AgentNetworkPolicies Module = "agent_network.policies"
|
||||
AgentNetworkGuardrails Module = "agent_network.guardrails"
|
||||
AgentNetworkBudgets Module = "agent_network.budgets"
|
||||
AgentNetworkUsage Module = "agent_network.usage"
|
||||
AgentNetworkLogs Module = "agent_network.logs"
|
||||
AgentNetworkSettings Module = "agent_network.settings"
|
||||
)
|
||||
|
||||
var All = map[Module]struct{}{
|
||||
@@ -40,4 +53,21 @@ var All = map[Module]struct{}{
|
||||
IdentityProviders: {},
|
||||
Services: {},
|
||||
AgentNetwork: {},
|
||||
|
||||
AgentNetworkProviders: {},
|
||||
AgentNetworkPolicies: {},
|
||||
AgentNetworkGuardrails: {},
|
||||
AgentNetworkBudgets: {},
|
||||
AgentNetworkUsage: {},
|
||||
AgentNetworkLogs: {},
|
||||
AgentNetworkSettings: {},
|
||||
}
|
||||
|
||||
// Parent returns the module owning a dotted submodule name and true, or the
|
||||
// module itself and false when it has no parent.
|
||||
func (m Module) Parent() (Module, bool) {
|
||||
if i := strings.IndexByte(string(m), '.'); i > 0 {
|
||||
return Module(string(m)[:i]), true
|
||||
}
|
||||
return m, false
|
||||
}
|
||||
|
||||
@@ -1685,14 +1685,34 @@ func (a *Account) injectPrivateServicePolicies(svc *service.Service, proxyPeers
|
||||
if len(proxyPeers) == 0 {
|
||||
return
|
||||
}
|
||||
// A service's AccessGroups can name groups that no longer exist — persisted
|
||||
// services and the agent-network synthesiser both carry the ids verbatim from
|
||||
// their own state. An unresolvable source authorises nothing, so drop it here
|
||||
// rather than let the network-map assembly resolve it to a nil group.
|
||||
sources := a.existingGroupIDs(svc.AccessGroups)
|
||||
if len(sources) == 0 {
|
||||
return
|
||||
}
|
||||
for _, proxyPeer := range proxyPeers {
|
||||
a.Policies = append(a.Policies, a.createPrivateServicePolicy(svc, proxyPeer))
|
||||
a.Policies = append(a.Policies, a.createPrivateServicePolicy(svc, proxyPeer, sources))
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Account) createPrivateServicePolicy(svc *service.Service, proxyPeer *nbpeer.Peer) *Policy {
|
||||
// existingGroupIDs returns the subset of groupIDs that resolve to a group in the account,
|
||||
// preserving the input order.
|
||||
func (a *Account) existingGroupIDs(groupIDs []string) []string {
|
||||
out := make([]string, 0, len(groupIDs))
|
||||
for _, groupID := range groupIDs {
|
||||
if _, ok := a.Groups[groupID]; ok {
|
||||
out = append(out, groupID)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func (a *Account) createPrivateServicePolicy(svc *service.Service, proxyPeer *nbpeer.Peer, accessGroups []string) *Policy {
|
||||
policyID := fmt.Sprintf("private-access-%s-%s", svc.ID, proxyPeer.ID)
|
||||
sources := append([]string(nil), svc.AccessGroups...)
|
||||
sources := append([]string(nil), accessGroups...)
|
||||
return &Policy{
|
||||
ID: policyID,
|
||||
Name: fmt.Sprintf("Private Access to %s", svc.Name),
|
||||
|
||||
@@ -68,7 +68,7 @@ type ProxyAccessTokenGenerated struct {
|
||||
// CreateNewProxyAccessToken generates a new proxy access token.
|
||||
// Returns the token with hashed value stored and plain token for one-time display.
|
||||
func CreateNewProxyAccessToken(name string, expiresIn time.Duration, accountID *string, createdBy string) (*ProxyAccessTokenGenerated, error) {
|
||||
hashedToken, plainToken, err := generateProxyToken()
|
||||
hashedToken, plainToken, err := GenerateProxyToken()
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
@@ -94,7 +94,10 @@ func CreateNewProxyAccessToken(name string, expiresIn time.Duration, accountID *
|
||||
}, nil
|
||||
}
|
||||
|
||||
func generateProxyToken() (HashedProxyToken, PlainProxyToken, error) {
|
||||
// GenerateProxyToken generates a new random proxy token, returning its SHA-256
|
||||
// hash (for storage) and the one-time plaintext. Exported so external modules
|
||||
// can mint tokens in the canonical proxy-token format.
|
||||
func GenerateProxyToken() (HashedProxyToken, PlainProxyToken, error) {
|
||||
secret, err := b.Random(ProxyTokenSecretLength)
|
||||
if err != nil {
|
||||
return "", "", err
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
package types
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
@@ -123,6 +124,22 @@ func TestCreateNewProxyAccessToken(t *testing.T) {
|
||||
})
|
||||
}
|
||||
|
||||
func TestGenerateProxyToken(t *testing.T) {
|
||||
hashed, plain, err := GenerateProxyToken()
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if err := plain.Validate(); err != nil {
|
||||
t.Errorf("generated token failed Validate(): %v", err)
|
||||
}
|
||||
if plain.Hash() != hashed {
|
||||
t.Error("returned hashed token does not match Hash(plain)")
|
||||
}
|
||||
if !strings.HasPrefix(string(plain), ProxyTokenPrefix) {
|
||||
t.Errorf("token %q missing prefix %q", plain, ProxyTokenPrefix)
|
||||
}
|
||||
}
|
||||
|
||||
func TestProxyAccessToken_IsExpired(t *testing.T) {
|
||||
past := time.Now().Add(-1 * time.Hour)
|
||||
future := time.Now().Add(1 * time.Hour)
|
||||
|
||||
Reference in New Issue
Block a user