Access Lists: Add status fixing functions (#62090)

This commit is contained in:
Pawel Kopiczko
2025-12-10 17:18:58 +00:00
committed by GitHub
parent 3b21ad909e
commit 339277f49f
5 changed files with 499 additions and 2 deletions
@@ -0,0 +1,69 @@
// Copyright 2025 Gravitational, Inc.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package backendtest
import (
"context"
"errors"
"iter"
"math/rand/v2"
"github.com/gravitational/teleport/lib/backend"
)
var ErrRandomBackend = errors.New("RandomlyErroringBackend error")
// RandomlyErroringBackend wraps Backend reading methods making them fail 50% of the time with
// [ErrRandomBackend].
type RandomlyErroringBackend struct {
backend.Backend
}
// NewRandomlyErroringBackend creates new instance of RandomlyErroringBackend for a given backend.
func NewRandomlyErroringBackend(b backend.Backend) *RandomlyErroringBackend {
return &RandomlyErroringBackend{
Backend: b,
}
}
func (b *RandomlyErroringBackend) Get(ctx context.Context, key backend.Key) (*backend.Item, error) {
if rand.IntN(2) == 0 {
return nil, ErrRandomBackend
}
return b.Backend.Get(ctx, key)
}
func (b *RandomlyErroringBackend) Items(ctx context.Context, params backend.ItemsParams) iter.Seq2[backend.Item, error] {
return func(yield func(backend.Item, error) bool) {
for item, err := range b.Backend.Items(ctx, params) {
var ok bool
if rand.IntN(2) == 0 {
ok = yield(backend.Item{}, ErrRandomBackend)
} else {
ok = yield(item, err)
}
if !ok {
return
}
}
}
}
func (b *RandomlyErroringBackend) GetRange(ctx context.Context, startKey, endKey backend.Key, limit int) (*backend.GetResult, error) {
if rand.IntN(2) == 0 {
return nil, ErrRandomBackend
}
return b.Backend.GetRange(ctx, startKey, endKey, limit)
}
@@ -0,0 +1,88 @@
// Copyright 2025 Gravitational, Inc.
//
// Licensed under the Apache License, Version 2.0 (the "License");
// you may not use this file except in compliance with the License.
// You may obtain a copy of the License at
//
// http://www.apache.org/licenses/LICENSE-2.0
//
// Unless required by applicable law or agreed to in writing, software
// distributed under the License is distributed on an "AS IS" BASIS,
// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
// See the License for the specific language governing permissions and
// limitations under the License.
package backendtest
import (
"testing"
"github.com/stretchr/testify/require"
"github.com/gravitational/teleport/lib/backend"
"github.com/gravitational/teleport/lib/backend/memory"
)
func TestRandomlyErroringBackend(t *testing.T) {
t.Parallel()
ctx := t.Context()
mem, err := memory.New(memory.Config{Context: ctx})
require.NoError(t, err)
var key, value = backend.NewKey("test_key_1"), []byte("test_value_1")
randomlyErroringBackend := NewRandomlyErroringBackend(mem)
_, err = randomlyErroringBackend.Put(ctx, backend.Item{Key: key, Value: value})
require.NoError(t, err)
const iterations, minBoundary, maxBoundary = 200, 50, 150
// Get
randomErrCnt := 0
for range iterations {
item, err := randomlyErroringBackend.Get(ctx, key)
if err != nil {
require.ErrorIs(t, err, ErrRandomBackend)
randomErrCnt++
} else {
require.Equal(t, key, item.Key)
require.Equal(t, value, item.Value)
}
}
require.GreaterOrEqual(t, randomErrCnt, minBoundary)
require.LessOrEqual(t, randomErrCnt, maxBoundary)
// GetRange
randomErrCnt = 0
for range iterations {
res, err := randomlyErroringBackend.GetRange(ctx, key, key, 100)
if err != nil {
require.ErrorIs(t, err, ErrRandomBackend)
randomErrCnt++
} else {
require.Len(t, res.Items, 1)
item := res.Items[0]
require.Equal(t, key, item.Key)
require.Equal(t, value, item.Value)
}
}
require.GreaterOrEqual(t, randomErrCnt, minBoundary)
require.LessOrEqual(t, randomErrCnt, maxBoundary)
// Items
randomErrCnt = 0
for range iterations {
for item, err := range randomlyErroringBackend.Items(ctx, backend.ItemsParams{StartKey: key, EndKey: key, Limit: 100}) {
if err != nil {
require.ErrorIs(t, err, ErrRandomBackend)
randomErrCnt++
} else {
require.Equal(t, key, item.Key)
require.Equal(t, value, item.Value)
}
}
}
require.GreaterOrEqual(t, randomErrCnt, minBoundary)
require.LessOrEqual(t, randomErrCnt, maxBoundary)
}
+7
View File
@@ -95,6 +95,13 @@ type AccessListsInternal interface {
// overwriting the list's members if successful.
UpdateAccessListAndOverwriteMembers(context.Context, *accesslist.AccessList, []*accesslist.AccessListMember) (*accesslist.AccessList, []*accesslist.AccessListMember, error)
// CleanupAccessListStatus removes invalid Status.OwnerOf and Status.MemberOf references.
CleanupAccessListStatus(ctx context.Context, accessListName string) (*accesslist.AccessList, error)
// EnsureNestedAccessListStatuses goes over all nested owners and nested members of the named
// access list and ensures nested lists' statuses owner_of/member_of contain the access list name.
EnsureNestedAccessListStatuses(ctx context.Context, accessListName string) error
// InsertAccessListCollection inserts a complete collection of access lists and their members from a single
// upstream source (e.g. EntraID) using a batch operation for improved performance.
//
+97
View File
@@ -1181,6 +1181,85 @@ func (a *AccessListService) updatedMembersNestedRelationships(ctx context.Contex
return nil
}
// CleanupAccessListStatus removes invalid Status.OwnerOf and Status.MemberOf references.
func (a *AccessListService) CleanupAccessListStatus(ctx context.Context, accessListName string) (*accesslist.AccessList, error) {
return a.runWithGlobalLockAccessList(ctx, accessListName, func() (*accesslist.AccessList, error) {
accessList, err := a.service.GetResource(ctx, accessListName)
if err != nil {
return nil, trace.Wrap(err)
}
var ownerRefreshErr error
accessList.Status.OwnerOf = slices.DeleteFunc(accessList.Status.OwnerOf, func(ownerOf string) bool {
ownedList, err := a.service.GetResource(ctx, ownerOf)
if err != nil {
if trace.IsNotFound(err) {
return true
}
ownerRefreshErr = err
return false
}
isActualOwner := slices.ContainsFunc(ownedList.Spec.Owners, func(ownedListOwner accesslist.Owner) bool {
return ownedListOwner.MembershipKind == accesslist.MembershipKindList && ownedListOwner.Name == accessList.GetName()
})
return !isActualOwner
})
if ownerRefreshErr != nil {
return nil, trace.Wrap(ownerRefreshErr)
}
var memberRefreshErr error
accessList.Status.MemberOf = slices.DeleteFunc(accessList.Status.MemberOf, func(memberOf string) bool {
if _, err := a.memberService.WithPrefix(memberOf).GetResource(ctx, accessList.GetName()); err != nil {
if trace.IsNotFound(err) {
return true
}
memberRefreshErr = err
}
return false
})
if memberRefreshErr != nil {
return nil, trace.Wrap(memberRefreshErr)
}
accessList, err = a.service.UpdateResource(ctx, accessList)
return accessList, trace.Wrap(err)
})
}
// EnsureNestedAccessListStatuses goes over all nested owners and nested members of the named
// access list and ensures nested lists' statuses owner_of/member_of contain the access list name.
func (a *AccessListService) EnsureNestedAccessListStatuses(ctx context.Context, accessListName string) error {
return a.runWithGlobalLock(ctx, accessListName, func() error {
accessList, err := a.service.GetResource(ctx, accessListName)
if err != nil {
return trace.Wrap(err)
}
for _, owner := range accessList.Spec.Owners {
if owner.MembershipKind == accesslist.MembershipKindList {
if err := a.updateAccessListOwnerOf(ctx, accessListName, owner.Name, true); err != nil {
return trace.Wrap(err)
}
}
}
members, err := a.memberService.WithPrefix(accessListName).GetResources(ctx)
if err != nil {
return trace.Wrap(err)
}
for _, member := range members {
if member.Spec.MembershipKind == accesslist.MembershipKindList {
if err := a.updateAccessListMemberOf(ctx, accessListName, member.GetName(), true); err != nil {
return trace.Wrap(err)
}
}
}
return nil
})
}
// InsertAccessListCollection inserts a complete collection of access lists and their members from a single
// upstream source (e.g. EntraID) using a batch operation for improved performance.
//
@@ -1249,3 +1328,21 @@ func (a *AccessListService) collectionToBackendItemsIter(collection *accesslists
}
}
}
func (a *AccessListService) runWithGlobalLock(ctx context.Context, accessListName string, fn func() error) error {
return a.service.RunWhileLocked(ctx, []string{accessListResourceLockName}, 2*accessListLockTTL, func(ctx context.Context, _ backend.Backend) error {
return a.service.RunWhileLocked(ctx, lockName(accessListName), 2*accessListLockTTL, func(ctx context.Context, _ backend.Backend) error {
return trace.Wrap(fn())
})
})
}
func (a *AccessListService) runWithGlobalLockAccessList(ctx context.Context, accessListName string, fn func() (*accesslist.AccessList, error)) (*accesslist.AccessList, error) {
var res *accesslist.AccessList
err := a.runWithGlobalLock(ctx, accessListName, func() error {
var err error
res, err = fn()
return trace.Wrap(err)
})
return res, err
}
+238 -2
View File
@@ -41,6 +41,7 @@ import (
"github.com/gravitational/teleport/entitlements"
"github.com/gravitational/teleport/lib/accesslists"
"github.com/gravitational/teleport/lib/backend"
"github.com/gravitational/teleport/lib/backend/backendtest"
"github.com/gravitational/teleport/lib/backend/memory"
"github.com/gravitational/teleport/lib/itertools/stream"
"github.com/gravitational/teleport/lib/modules"
@@ -1792,6 +1793,16 @@ func createAccessList(t *testing.T, service *AccessListService, name string, clo
return upserted
}
func getAccessList(t *testing.T, service *AccessListService, name string) *accesslist.AccessList {
t.Helper()
ctx := context.Background()
al, err := service.GetAccessList(ctx, name)
require.NoError(t, err)
return al
}
type accessListMemberOptions struct {
membershipKind string
expires time.Time
@@ -1834,6 +1845,17 @@ func newAccessListMember(t *testing.T, accessList, name string, opts ...accessLi
return member
}
func createAccessListMember(t *testing.T, service *AccessListService, accessList, name string, opts ...accessListMemberOpt) *accesslist.AccessListMember {
t.Helper()
ctx := t.Context()
m := newAccessListMember(t, accessList, name, opts...)
m, err := service.UpsertAccessListMember(ctx, m)
require.NoError(t, err)
return m
}
func newAccessListReview(t *testing.T, accessList, name string) *accesslist.Review {
t.Helper()
@@ -2228,6 +2250,220 @@ func TestAccessListService_Status_MemberOf(t *testing.T) {
})
}
func TestAccessListService_CleanupAccessListStatus(t *testing.T) {
ctx := t.Context()
clock := clockwork.NewFakeClock()
mem, err := memory.New(memory.Config{
Context: ctx,
Clock: clock,
})
require.NoError(t, err)
service := newAccessListService(t, mem, clock, true /* igsEnabled */)
const a1, a2, a3, a4, a5, a6 = "al_1", "al_2", "al_3", "al_4", "al_5", "al_6"
userOwner := accesslist.Owner{MembershipKind: accesslist.MembershipKindUser, Name: "test_user_1"}
a1Owner := accesslist.Owner{MembershipKind: accesslist.MembershipKindList, Name: a1}
// a1 is the target access list which status will be fixed
_ = createAccessList(t, service, a1, clock)
// a2 is a list owned by a1
_ = createAccessList(t, service, a2, clock,
withOwners([]accesslist.Owner{userOwner, a1Owner}))
// a3 is an owned list which will be deleted without updating a1 status
_ = createAccessList(t, service, a3, clock,
withOwners([]accesslist.Owner{userOwner, a1Owner}))
// a4 is another owned list by a1, for this one ownership will be removed without updating a1 status
a4List := createAccessList(t, service, a4, clock,
withOwners([]accesslist.Owner{userOwner, a1Owner}))
// a5 is a parent for a1 (a1 is a member of a3)
_ = createAccessList(t, service, a5, clock)
_ = createAccessListMember(t, service, a5, a1, withMembershipKind(accesslist.MembershipKindList))
// a6 is the list of which a1 is a member; the membership will be removed without updating a1 status
_ = createAccessList(t, service, a6, clock)
_ = createAccessListMember(t, service, a6, a1, withMembershipKind(accesslist.MembershipKindList))
// Let's verify the a1 status is correct
requireStatusOwnerOf(t, service, a1, []string{a2, a3, a4})
requireStatusMemberOf(t, service, a1, []string{a5, a6})
// Let's now break the a1 status, we need to use generic service directly to bypass status
// updates
service.service.DeleteResource(ctx, a3)
a4List.Spec.Owners = []accesslist.Owner{userOwner} // remove a1Owner
service.service.UpdateResource(ctx, a4List)
service.memberService.WithPrefix(a6).DeleteResource(ctx, a1)
// Let's check the status remain untouched:
// - a3 should be removed from owner_of because it doesn't exist anymore
// - a4 should be removed from owner_of because a1 is not an owner anymore
requireStatusOwnerOf(t, service, a1, []string{a2, a3, a4})
// - a6 should be removed from member_of because a6/a1 membership doesn't exist anymore
requireStatusMemberOf(t, service, a1, []string{a5, a6})
// Run cleanup
fixedA1List, err := service.CleanupAccessListStatus(ctx, a1)
require.NoError(t, err)
// Check status of the returned list
require.Equal(t, []string{a2}, fixedA1List.Status.OwnerOf)
require.Equal(t, []string{a5}, fixedA1List.Status.MemberOf)
// Also check the status in the storage
requireStatusOwnerOf(t, service, a1, []string{a2})
requireStatusMemberOf(t, service, a1, []string{a5})
// Now let's see if CleanupStatus can cope with emptying status completely
// Let's remove remove a1->a2 ownership and a5->a1 membership using the generic service and
// therefore bypassing status update
a2List := getAccessList(t, service, a2)
a2List.Spec.Owners = []accesslist.Owner{userOwner} // remove a1Owner
service.service.UpdateResource(ctx, a2List)
service.memberService.WithPrefix(a5).DeleteResource(ctx, a1)
// Verify the status is broken now as it should be empty
requireStatusOwnerOf(t, service, a1, []string{a2})
requireStatusMemberOf(t, service, a1, []string{a5})
// Run cleanup
fixedA1List, err = service.CleanupAccessListStatus(ctx, a1)
require.NoError(t, err)
// Check status of the returned list
require.Empty(t, fixedA1List.Status.OwnerOf)
require.Empty(t, fixedA1List.Status.MemberOf)
// Also check the status in the storage
requireStatusOwnerOf(t, service, a1, []string{})
requireStatusMemberOf(t, service, a1, []string{})
}
func TestAccessListService_CleanupAccessListStatus_does_not_panic(t *testing.T) {
for range 10 {
ctx := t.Context()
clock := clockwork.NewFakeClock()
mem, err := memory.New(memory.Config{
Context: ctx,
Clock: clock,
})
require.NoError(t, err)
service := newAccessListService(t, mem, clock, true /* igsEnabled */)
const a1, a2, a3, a4, a5 = "al_1", "al_2", "al_3", "al_4", "al_5"
a1Owner := accesslist.Owner{MembershipKind: accesslist.MembershipKindList, Name: a1}
// a1 is owner of a2 and a3 and member of a4 and a5
_ = createAccessList(t, service, a1, clock)
_ = createAccessList(t, service, a2, clock,
withOwners([]accesslist.Owner{a1Owner}))
_ = createAccessList(t, service, a3, clock,
withOwners([]accesslist.Owner{a1Owner}))
_ = createAccessList(t, service, a4, clock)
_ = createAccessListMember(t, service, a4, a1, withMembershipKind(accesslist.MembershipKindList))
_ = createAccessList(t, service, a5, clock)
_ = createAccessListMember(t, service, a5, a1, withMembershipKind(accesslist.MembershipKindList))
// Recreate the service, but with the randomly erroring backend. This is a wrapper so all
// the data will be retained.
service = newAccessListService(t, backendtest.NewRandomlyErroringBackend(mem), clock, true /* igsEnabled */)
require.NotPanics(t, func() {
_, _ = service.CleanupAccessListStatus(ctx, a1)
})
}
}
func TestAccessListService_EnsureNestedAccessListStatuses(t *testing.T) {
ctx := t.Context()
clock := clockwork.NewFakeClock()
mem, err := memory.New(memory.Config{
Context: ctx,
Clock: clock,
})
require.NoError(t, err)
service := newAccessListService(t, mem, clock, true /* igsEnabled */)
const a1, a2, a3, a4, a5, a6 = "al_1", "al_2", "al_3", "al_4", "al_5", "al_6"
const ghost = "ghost_list"
a2Owner := accesslist.Owner{MembershipKind: accesslist.MembershipKindList, Name: a2}
a3Owner := accesslist.Owner{MembershipKind: accesslist.MembershipKindList, Name: a3}
// Setup:
// - a1 is the list that will be missing from other list status owner_of/member_of
// - a2 is the owned list which status.owned_of will be fixed by adding a1
// - a3 is the owned list which status.owned_of will be partially fixed by adding a1
// - a4 is the member list which status.member_of will be fixed by adding a1
// - a5 is the member list which status.member_of will be partially fixed by adding a1
a2List := createAccessList(t, service, a2, clock)
a3List := createAccessList(t, service, a3, clock)
_ = createAccessList(t, service, a1, clock, withOwners([]accesslist.Owner{a2Owner, a3Owner}))
a4List := createAccessList(t, service, a4, clock)
_ = createAccessListMember(t, service, a1, a4, withMembershipKind(accesslist.MembershipKindList))
a5List := createAccessList(t, service, a5, clock)
_ = createAccessListMember(t, service, a1, a5, withMembershipKind(accesslist.MembershipKindList))
// Verify the target statuses
requireStatusOwnerOf(t, service, a2, []string{a1})
requireStatusOwnerOf(t, service, a3, []string{a1})
requireStatusOwnerOf(t, service, a4, []string{})
requireStatusOwnerOf(t, service, a5, []string{})
requireStatusMemberOf(t, service, a2, []string{})
requireStatusMemberOf(t, service, a3, []string{})
requireStatusMemberOf(t, service, a4, []string{a1})
requireStatusMemberOf(t, service, a5, []string{a1})
// Let's break the statuses
a2List.Status.OwnerOf = nil
a3List.Status.OwnerOf = []string{ghost}
a4List.Status.MemberOf = nil
a5List.Status.MemberOf = []string{ghost}
for _, al := range []*accesslist.AccessList{a2List, a3List, a4List, a5List} {
_, err := service.service.UpdateResource(ctx, al)
require.NoError(t, err, "access_list = %q", al.GetName())
}
// Verify the statuses are broken now:
requireStatusOwnerOf(t, service, a2, []string{})
requireStatusOwnerOf(t, service, a3, []string{ghost})
requireStatusOwnerOf(t, service, a4, []string{})
requireStatusOwnerOf(t, service, a5, []string{})
requireStatusMemberOf(t, service, a2, []string{})
requireStatusMemberOf(t, service, a3, []string{})
requireStatusMemberOf(t, service, a4, []string{})
requireStatusMemberOf(t, service, a5, []string{ghost})
// Ensure a1 is present where it should be (apply the fix)
err = service.EnsureNestedAccessListStatuses(ctx, a1)
require.NoError(t, err)
// Verify the statuses are fixed (a1 is added back, but ghost is not removed as we applied
// the fix for a1 only):
requireStatusOwnerOf(t, service, a2, []string{a1})
requireStatusOwnerOf(t, service, a3, []string{ghost, a1})
requireStatusOwnerOf(t, service, a4, []string{})
requireStatusOwnerOf(t, service, a5, []string{})
requireStatusMemberOf(t, service, a2, []string{})
requireStatusMemberOf(t, service, a3, []string{})
requireStatusMemberOf(t, service, a4, []string{a1})
requireStatusMemberOf(t, service, a5, []string{ghost, a1})
}
func requireStatusOwnerOf(t *testing.T, service *AccessListService, accessListName string, ownerOf []string) {
t.Helper()
ctx := context.Background()
@@ -2248,7 +2484,7 @@ func requireStatusMemberOf(t *testing.T, service *AccessListService, accessListN
require.ElementsMatch(t, memberOf, accessList.Status.MemberOf)
}
func newAccessListService(t *testing.T, mem *memory.Memory, clock clockwork.Clock, igsEnabled bool) *AccessListService {
func newAccessListService(t *testing.T, b backend.Backend, clock clockwork.Clock, igsEnabled bool) *AccessListService {
t.Helper()
modulestest.SetTestModules(t, modulestest.Modules{
@@ -2260,7 +2496,7 @@ func newAccessListService(t *testing.T, mem *memory.Memory, clock clockwork.Cloc
},
})
service, err := NewAccessListService(backend.NewSanitizer(mem), clock)
service, err := NewAccessListService(backend.NewSanitizer(b), clock)
require.NoError(t, err)
return service