diff --git a/management/internals/modules/agentnetwork/manager.go b/management/internals/modules/agentnetwork/manager.go index 4f31e0791..94f5de35f 100644 --- a/management/internals/modules/agentnetwork/manager.go +++ b/management/internals/modules/agentnetwork/manager.go @@ -1036,6 +1036,9 @@ func (m *managerImpl) bootstrapSelfAddressed(ctx context.Context, settings *type if err != nil { return status.Errorf(status.InvalidArgument, "invalid endpoint: %s", err) } + if err := m.requireHostNotForeign(ctx, settings.AccountID, hostname); err != nil { + return err + } settings.Domain = hostname settings.ProxyAddress = hostname @@ -1162,6 +1165,9 @@ func (m *managerImpl) bootstrapLabeled(ctx context.Context, settings *types.Sett if err != nil { return status.Errorf(status.InvalidArgument, "invalid proxy_address: %s", err) } + if err := m.requireHostNotForeign(ctx, settings.AccountID, parent); err != nil { + return err + } if err := m.validateGatewayCluster(ctx, settings.AccountID, parent); err != nil { return err @@ -1212,6 +1218,24 @@ func (m *managerImpl) bootstrapLabeled(ctx context.Context, settings *types.Sett return fmt.Errorf("allocate agent network endpoint for account %s: %d attempts exhausted", settings.AccountID, maxDomainAllocationAttempts) } +// requireHostNotForeign refuses to pin the account's gateway onto a host that +// another account's proxy declares. The pin's proxy_address is what selects +// the proxy that serves the endpoint, and an account-scoped proxy only ever +// receives its own account's mappings, so such a pin could never be served — +// and the endpoint it assigns is immutable. Shared proxies are not foreign, and +// a host no proxy has declared stays pinnable: claiming the address before the +// proxy's first connection is the documented order. +func (m *managerImpl) requireHostNotForeign(ctx context.Context, accountID, host string) error { + foreign, err := m.store.HasForeignAccountProxyAtHost(ctx, host, accountID) + if err != nil { + return fmt.Errorf("check proxy host ownership: %w", err) + } + if foreign { + return status.Errorf(status.InvalidArgument, "proxy cluster %s is not available to this account", host) + } + return nil +} + // isUniqueConstraintError reports whether err is a database unique-constraint // violation, matched on the driver message because CreateAgentNetworkSettings // deliberately returns the driver error unwrapped. diff --git a/management/server/group.go b/management/server/group.go index ca20a6b08..88295e2f6 100644 --- a/management/server/group.go +++ b/management/server/group.go @@ -774,6 +774,14 @@ func validateDeleteGroup(ctx context.Context, transaction store.Store, group *ty return &GroupLinkError{"agent network policy", linkedPolicy.Name} } + isLinked, linkedRule, err := isGroupLinkedToAgentNetworkBudgetRule(ctx, transaction, group.AccountID, group.ID) + if err != nil { + return status.Errorf(status.Internal, "failed to check agent network budget rules") + } + if isLinked { + return &GroupLinkError{"agent network budget rule", linkedRule.Name} + } + return checkGroupLinkedToSettings(ctx, transaction, group) } @@ -945,6 +953,26 @@ func isGroupLinkedToAgentNetworkPolicy(ctx context.Context, transaction store.St return false, nil } +// isGroupLinkedToAgentNetworkBudgetRule checks if a group is a target of any +// account-level agent network budget rule. +func isGroupLinkedToAgentNetworkBudgetRule(ctx context.Context, transaction store.Store, accountID string, groupID string) (bool, *agentNetworkTypes.AccountBudgetRule, error) { + rules, err := transaction.GetAccountAgentNetworkBudgetRules(ctx, store.LockingStrengthNone, accountID) + if err != nil { + log.WithContext(ctx).Errorf("error retrieving agent network budget rules while checking group linkage: %v", err) + return false, nil, err + } + + for _, rule := range rules { + if rule == nil { + continue + } + if slices.Contains(rule.TargetGroups, groupID) { + return true, rule, nil + } + } + return false, nil, 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) { diff --git a/management/server/group_test.go b/management/server/group_test.go index da056c8a9..fa351a43e 100644 --- a/management/server/group_test.go +++ b/management/server/group_test.go @@ -132,6 +132,11 @@ func TestDefaultAccountManager_DeleteGroup(t *testing.T) { "grp-for-agent-network-policy", "agent network policy", }, + { + "agent network budget rule", + "grp-for-agent-network-budget-rule", + "agent network budget rule", + }, { "reverse proxy private service access group", "grp-for-rp-private", @@ -152,6 +157,16 @@ func TestDefaultAccountManager_DeleteGroup(t *testing.T) { return } + group, getErr := am.GetGroup(context.Background(), account.Id, testCase.groupID, groupAdminUserID) + if getErr != nil { + t.Errorf("group %s should still exist after failed deletion: %s", testCase.groupID, getErr) + return + } + if group == nil { + t.Errorf("group %s was deleted despite the failed deletion", testCase.groupID) + return + } + var sErr *status.Error if errors.As(err, &sErr) { if sErr.Message != testCase.expectedReason { @@ -240,6 +255,12 @@ func TestDefaultAccountManager_DeleteGroups(t *testing.T) { groupIDs: []string{"grp-for-agent-network-policy"}, expectedReasons: []string{"agent network policy"}, }, + { + name: "agent network budget rule", + groupIDs: []string{"grp-for-agent-network-budget-rule"}, + expectedReasons: []string{"agent network budget rule"}, + expectedNotDeleted: []string{"grp-for-agent-network-budget-rule"}, + }, { name: "reverse proxy services", groupIDs: []string{"grp-for-rp-private", "grp-for-rp-bearer"}, @@ -501,6 +522,14 @@ func initTestGroupAccount(am *DefaultAccountManager) (*DefaultAccountManager, *t Peers: make([]string, 0), } + groupForAgentNetworkBudgetRule := &types.Group{ + ID: "grp-for-agent-network-budget-rule", + AccountID: "account-id", + Name: "Group for agent network budget rules", + Issued: types.GroupIssuedAPI, + Peers: make([]string, 0), + } + groupForRPPrivate := &types.Group{ ID: "grp-for-rp-private", AccountID: "account-id", @@ -573,6 +602,7 @@ func initTestGroupAccount(am *DefaultAccountManager) (*DefaultAccountManager, *t _ = 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, groupForAgentNetworkBudgetRule) _ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForRPPrivate) _ = am.CreateGroup(context.Background(), accountID, groupAdminUserID, groupForRPBearer) @@ -587,6 +617,20 @@ func initTestGroupAccount(am *DefaultAccountManager) (*DefaultAccountManager, *t return nil, nil, err } + budgetRuleDecoy := agentNetworkTypes.NewAccountBudgetRule(accountID) + budgetRuleDecoy.Name = "Unrelated agent network budget rule" + budgetRuleDecoy.TargetGroups = []string{"unrelated-group"} + if err := am.Store.SaveAgentNetworkBudgetRule(context.Background(), budgetRuleDecoy); err != nil { + return nil, nil, err + } + + budgetRule := agentNetworkTypes.NewAccountBudgetRule(accountID) + budgetRule.Name = "Example agent network budget rule" + budgetRule.TargetGroups = []string{groupForAgentNetworkBudgetRule.ID} + if err := am.Store.SaveAgentNetworkBudgetRule(context.Background(), budgetRule); 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{ diff --git a/management/server/store/sql_store.go b/management/server/store/sql_store.go index 08ec45395..ec4c2bb55 100644 --- a/management/server/store/sql_store.go +++ b/management/server/store/sql_store.go @@ -6446,6 +6446,25 @@ func (s *SqlStore) IsClusterAddressConflicting(ctx context.Context, clusterAddre return count > 0, nil } +// HasForeignAccountProxyAtHost reports whether a proxy owned by a different +// account declares this host. Shared proxies (account_id IS NULL) are not +// foreign: a shared cluster is what most accounts pin their agent network +// gateway to. The match folds case because proxies declare their address as +// the operator spelled it while the caller's host is normalised; that costs a +// scan of the proxies table, taken once per account when its gateway is +// bootstrapped, not on the per-connect path IsClusterAddressConflicting serves. +func (s *SqlStore) HasForeignAccountProxyAtHost(ctx context.Context, host, accountID string) (bool, error) { + var count int64 + result := s.db. + Model(&proxy.Proxy{}). + Where("LOWER(cluster_address) = LOWER(?) AND account_id IS NOT NULL AND account_id != ?", host, accountID). + Count(&count) + if result.Error != nil { + return false, status.Errorf(status.Internal, "check proxy host ownership: %v", result.Error) + } + return count > 0, nil +} + func (s *SqlStore) DeleteAccountCluster(ctx context.Context, clusterAddress, accountID string) error { result := s.db. Where("cluster_address = ? AND account_id = ?", clusterAddress, accountID). diff --git a/management/server/store/store.go b/management/server/store/store.go index 6886536b9..55a8c319c 100644 --- a/management/server/store/store.go +++ b/management/server/store/store.go @@ -340,6 +340,7 @@ type Store interface { CountProxiesByAccountID(ctx context.Context, accountID string) (int64, error) IsClusterAddressConflicting(ctx context.Context, clusterAddress, accountID string) (bool, error) HasActiveProxyAtClusterAddress(ctx context.Context, clusterAddress string) (bool, error) + HasForeignAccountProxyAtHost(ctx context.Context, host, accountID string) (bool, error) DeleteAccountCluster(ctx context.Context, clusterAddress, accountID string) error GetCustomDomainsCounts(ctx context.Context) (total int64, validated int64, err error) diff --git a/management/server/store/store_mock.go b/management/server/store/store_mock.go index 04f79d30a..4b8284212 100644 --- a/management/server/store/store_mock.go +++ b/management/server/store/store_mock.go @@ -3065,6 +3065,21 @@ func (mr *MockStoreMockRecorder) HasActiveProxyAtClusterAddress(ctx, clusterAddr return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "HasActiveProxyAtClusterAddress", reflect.TypeOf((*MockStore)(nil).HasActiveProxyAtClusterAddress), ctx, clusterAddress) } +// HasForeignAccountProxyAtHost mocks base method. +func (m *MockStore) HasForeignAccountProxyAtHost(ctx context.Context, host, accountID string) (bool, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "HasForeignAccountProxyAtHost", ctx, host, accountID) + ret0, _ := ret[0].(bool) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// HasForeignAccountProxyAtHost indicates an expected call of HasForeignAccountProxyAtHost. +func (mr *MockStoreMockRecorder) HasForeignAccountProxyAtHost(ctx, host, accountID any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "HasForeignAccountProxyAtHost", reflect.TypeOf((*MockStore)(nil).HasForeignAccountProxyAtHost), ctx, host, accountID) +} + // IncrementAgentNetworkConsumption mocks base method. func (m *MockStore) IncrementAgentNetworkConsumption(ctx context.Context, accountID string, kind types.ConsumptionDimension, dimID string, windowSeconds int64, windowStart time.Time, tokensIn, tokensOut int64, costUSD float64) error { m.ctrl.T.Helper()