Fix headless authentication matching logic for watcher (#28843)

* Fix headless authentication matching logic for watcher and add test.

* Move hasWatchPermissionForKind to a separate function.

* Clean up hasWatchPermissionForKind.

* Cleanup test code with suggestions from review.
This commit is contained in:
Brian Joerger
2023-07-11 20:05:17 +00:00
committed by GitHub
parent d4b3afe9a1
commit 1b73bafca8
5 changed files with 237 additions and 74 deletions
+4
View File
@@ -143,6 +143,10 @@ func (kind WatchKind) Matches(e Event) (bool, error) {
var filter CertAuthorityFilter
filter.FromMap(kind.Filter)
return filter.Match(res), nil
case *HeadlessAuthentication:
var filter HeadlessAuthenticationFilter
filter.FromMap(kind.Filter)
return filter.Match(res), nil
default:
// we don't know about this filter, let the event through
}
+1 -1
View File
@@ -158,7 +158,7 @@ func (f *HeadlessAuthenticationFilter) FromMap(m map[string]string) error {
}
// Match checks if a given headless authentication matches this filter.
func (f *HeadlessAuthenticationFilter) Match(req HeadlessAuthentication) bool {
func (f *HeadlessAuthenticationFilter) Match(req *HeadlessAuthentication) bool {
if f.Name != "" && req.GetName() != f.Name {
return false
}
+61 -69
View File
@@ -1414,76 +1414,12 @@ func (a *ServerWithRoles) NewWatcher(ctx context.Context, watch types.Watch) (ty
validKinds := make([]types.WatchKind, 0, len(watch.Kinds))
for _, kind := range watch.Kinds {
// Check the permissions for data of each kind. For watching, most
// kinds of data just need a Read permission, but some have more
// complicated logic.
switch kind.Kind {
case types.KindCertAuthority:
verb := types.VerbReadNoSecrets
if kind.LoadSecrets {
verb = types.VerbRead
}
if err := a.action(apidefaults.Namespace, types.KindCertAuthority, verb); err != nil {
if watch.AllowPartialSuccess {
continue
}
return nil, trace.Wrap(err)
}
case types.KindAccessRequest:
var filter types.AccessRequestFilter
if err := filter.FromMap(kind.Filter); err != nil {
if watch.AllowPartialSuccess {
continue
}
return nil, trace.Wrap(err)
}
if filter.User == "" || a.currentUserAction(filter.User) != nil {
if err := a.action(apidefaults.Namespace, types.KindAccessRequest, types.VerbRead); err != nil {
if watch.AllowPartialSuccess {
continue
}
return nil, trace.Wrap(err)
}
}
case types.KindWebSession:
var filter types.WebSessionFilter
if err := filter.FromMap(kind.Filter); err != nil {
if watch.AllowPartialSuccess {
continue
}
return nil, trace.Wrap(err)
}
// Allow reading Snowflake sessions to DB service.
if !(kind.SubKind == types.KindSnowflakeSession && a.hasBuiltinRole(types.RoleDatabase)) {
if filter.User == "" || a.currentUserAction(filter.User) != nil {
if err := a.action(apidefaults.Namespace, types.KindWebSession, types.VerbRead); err != nil {
if watch.AllowPartialSuccess {
continue
}
return nil, trace.Wrap(err)
}
}
}
case types.KindHeadlessAuthentication:
var filter types.HeadlessAuthenticationFilter
if err := filter.FromMap(kind.Filter); err != nil {
if watch.AllowPartialSuccess {
continue
}
return nil, trace.Wrap(err)
}
// Only users can watch their own headless authentications.
if !hasLocalUserRole(a.context) || filter.Username != a.context.User.GetName() {
return nil, trace.AccessDenied("user %q cannot watch headless authentications of %q", a.context.User.GetName(), filter.Username)
}
default:
if err := a.action(apidefaults.Namespace, kind.Kind, types.VerbRead); err != nil {
if watch.AllowPartialSuccess {
continue
}
return nil, trace.Wrap(err)
err := a.hasWatchPermissionForKind(kind)
if err != nil {
if watch.AllowPartialSuccess {
continue
}
return nil, trace.Wrap(err)
}
validKinds = append(validKinds, kind)
@@ -1503,6 +1439,62 @@ func (a *ServerWithRoles) NewWatcher(ctx context.Context, watch types.Watch) (ty
return a.authServer.NewWatcher(ctx, watch)
}
// hasWatchPermissionForKind checks the permissions for data of each kind.
// For watching, most kinds of data just need a Read permission, but some
// have more complicated logic.
func (a *ServerWithRoles) hasWatchPermissionForKind(kind types.WatchKind) error {
verb := types.VerbRead
switch kind.Kind {
case types.KindCertAuthority:
if !kind.LoadSecrets {
verb = types.VerbReadNoSecrets
}
case types.KindAccessRequest:
var filter types.AccessRequestFilter
if err := filter.FromMap(kind.Filter); err != nil {
return trace.Wrap(err)
}
// Users can watch their own access requests.
if filter.User != "" && a.currentUserAction(filter.User) == nil {
return nil
}
case types.KindWebSession:
var filter types.WebSessionFilter
if err := filter.FromMap(kind.Filter); err != nil {
return trace.Wrap(err)
}
// Allow reading Snowflake sessions to DB service.
if kind.SubKind == types.KindSnowflakeSession && a.hasBuiltinRole(types.RoleDatabase) {
return nil
}
// Users can watch their own web sessions.
if filter.User != "" && a.currentUserAction(filter.User) == nil {
return nil
}
case types.KindHeadlessAuthentication:
var filter types.HeadlessAuthenticationFilter
if err := filter.FromMap(kind.Filter); err != nil {
return trace.Wrap(err)
}
// Users can only watch their own headless authentications, meaning we don't fallback to
// the generalized verb-kind-action check below.
if !hasLocalUserRole(a.context) {
return trace.AccessDenied("non-local user roles cannot watch headless authentications")
} else if filter.Username == "" {
return trace.AccessDenied("user cannot watch headless authentications without a filter for their username")
} else if filter.Username != a.context.User.GetName() {
return trace.AccessDenied("user %q cannot watch headless authentications of %q", a.context.User.GetName(), filter.Username)
}
return nil
}
return trace.Wrap(a.action(apidefaults.Namespace, kind.Kind, verb))
}
// DeleteAllNodes deletes all nodes in a given namespace
func (a *ServerWithRoles) DeleteAllNodes(ctx context.Context, namespace string) error {
if err := a.action(namespace, types.KindNode, types.VerbDelete); err != nil {
+146
View File
@@ -5706,3 +5706,149 @@ func mustResourceID(clusterName, kind, name string) types.ResourceID {
Name: name,
}
}
func TestWatchHeadlessAuthentications_usersCanOnlyWatchThemselves(t *testing.T) {
t.Parallel()
ctx := context.Background()
srv := newTestTLSServer(t)
alice, bob, admin := createSessionTestUsers(t, srv.Auth())
// For each user, prepare 4 different headless authentications with the varying states.
// These will be created during each test, and the watcher will return a subset of the
// collected events based on the test's filter.
var headlessAuthns []*types.HeadlessAuthentication
for _, username := range []string{alice, bob} {
for _, state := range []types.HeadlessAuthenticationState{
types.HeadlessAuthenticationState_HEADLESS_AUTHENTICATION_STATE_UNSPECIFIED,
types.HeadlessAuthenticationState_HEADLESS_AUTHENTICATION_STATE_PENDING,
types.HeadlessAuthenticationState_HEADLESS_AUTHENTICATION_STATE_DENIED,
types.HeadlessAuthenticationState_HEADLESS_AUTHENTICATION_STATE_APPROVED,
} {
ha, err := types.NewHeadlessAuthentication(username, uuid.NewString(), srv.Clock().Now().Add(time.Minute))
require.NoError(t, err)
ha.State = state
headlessAuthns = append(headlessAuthns, ha)
}
}
aliceAuthns := headlessAuthns[:4]
bobAuthns := headlessAuthns[4:]
testCases := []struct {
name string
identity TestIdentity
filter types.HeadlessAuthenticationFilter
expectWatchError string
expectResources []*types.HeadlessAuthentication
}{
{
name: "NOK non local users cannot watch headless authentications",
identity: TestAdmin(),
expectWatchError: "non-local user roles cannot watch headless authentications",
},
{
name: "NOK must filter for username",
identity: TestUser(admin),
filter: types.HeadlessAuthenticationFilter{},
expectWatchError: "user cannot watch headless authentications without a filter for their username",
}, {
name: "NOK alice cannot filter for username=bob",
identity: TestUser(alice),
filter: types.HeadlessAuthenticationFilter{
Username: bob,
},
expectWatchError: "user \"alice\" cannot watch headless authentications of \"bob\"",
}, {
name: "OK alice can filter for username=alice",
identity: TestUser(alice),
filter: types.HeadlessAuthenticationFilter{
Username: alice,
},
expectResources: aliceAuthns,
}, {
name: "OK bob can filter for username=bob",
identity: TestUser(bob),
filter: types.HeadlessAuthenticationFilter{
Username: bob,
},
expectResources: bobAuthns,
}, {
name: "OK alice can filter for pending requests",
identity: TestUser(alice),
filter: types.HeadlessAuthenticationFilter{
Username: alice,
State: types.HeadlessAuthenticationState_HEADLESS_AUTHENTICATION_STATE_PENDING,
},
expectResources: []*types.HeadlessAuthentication{aliceAuthns[types.HeadlessAuthenticationState_HEADLESS_AUTHENTICATION_STATE_PENDING]},
}, {
name: "OK alice can filter for a specific request",
identity: TestUser(alice),
filter: types.HeadlessAuthenticationFilter{
Username: alice,
Name: headlessAuthns[2].GetName(),
},
expectResources: aliceAuthns[2:3],
},
}
for _, tc := range testCases {
t.Run(tc.name, func(t *testing.T) {
client, err := srv.NewClient(tc.identity)
require.NoError(t, err)
watchCtx, cancel := context.WithTimeout(ctx, 5*time.Second)
defer cancel()
watcher, err := client.NewWatcher(watchCtx, types.Watch{
Kinds: []types.WatchKind{
{
Kind: types.KindHeadlessAuthentication,
Filter: tc.filter.IntoMap(),
},
},
})
require.NoError(t, err)
select {
case event := <-watcher.Events():
require.Equal(t, types.OpInit, event.Type, "Expected watcher init event but got %v", event)
case <-time.After(time.Second):
t.Fatal("Failed to receive watcher init event before timeout")
case <-watcher.Done():
if tc.expectWatchError != "" {
require.True(t, trace.IsAccessDenied(watcher.Error()), "Expected access denied error but got %v", err)
require.ErrorContains(t, watcher.Error(), tc.expectWatchError)
return
}
t.Fatalf("Watcher unexpectedly closed with error: %v", watcher.Error())
}
for _, ha := range headlessAuthns {
err = srv.Auth().UpsertHeadlessAuthentication(ctx, ha)
require.NoError(t, err)
}
var expectEvents []types.Event
for _, expectResource := range tc.expectResources {
expectEvents = append(expectEvents, types.Event{
Type: types.OpPut,
Resource: expectResource,
})
}
var events []types.Event
loop:
for {
select {
case event := <-watcher.Events():
events = append(events, event)
case <-time.After(100 * time.Millisecond):
break loop
case <-watcher.Done():
t.Fatalf("Watcher unexpectedly closed with error: %v", watcher.Error())
}
}
require.Equal(t, expectEvents, events)
})
}
}
+25 -4
View File
@@ -163,7 +163,14 @@ func (e *EventsService) NewWatcher(ctx context.Context, watch types.Watch) (type
case types.KindIntegration:
parser = newIntegrationParser()
case types.KindHeadlessAuthentication:
parser = newHeadlessAuthenticationParser()
p, err := newHeadlessAuthenticationParser(kind.Filter)
if err != nil {
if watch.AllowPartialSuccess {
continue
}
return nil, trace.Wrap(err)
}
parser = p
default:
if watch.AllowPartialSuccess {
continue
@@ -1553,14 +1560,21 @@ func (p *integrationParser) parse(event backend.Event) (types.Resource, error) {
}
}
func newHeadlessAuthenticationParser() *headlessAuthenticationParser {
func newHeadlessAuthenticationParser(m map[string]string) (*headlessAuthenticationParser, error) {
var filter types.HeadlessAuthenticationFilter
if err := filter.FromMap(m); err != nil {
return nil, trace.Wrap(err)
}
return &headlessAuthenticationParser{
baseParser: newBaseParser(backend.Key(headlessAuthenticationPrefix)),
}
filter: filter,
}, nil
}
type headlessAuthenticationParser struct {
baseParser
filter types.HeadlessAuthenticationFilter
}
func (p *headlessAuthenticationParser) parse(event backend.Event) (types.Resource, error) {
@@ -1568,7 +1582,14 @@ func (p *headlessAuthenticationParser) parse(event backend.Event) (types.Resourc
case types.OpDelete:
return resourceHeader(event, types.KindIntegration, types.V1, 0)
case types.OpPut:
return unmarshalHeadlessAuthentication(event.Item.Value)
ha, err := unmarshalHeadlessAuthentication(event.Item.Value)
if err != nil {
return nil, trace.Wrap(err)
}
if !p.filter.Match(ha) {
return nil, nil
}
return ha, nil
default:
return nil, trace.BadParameter("event %v is not supported", event.Type)
}