From 25973b3974bc8842e7d15e9f87bcc576adeda701 Mon Sep 17 00:00:00 2001 From: Juan Fernandez Date: Wed, 24 Jun 2026 08:27:53 -0400 Subject: [PATCH] Authorize query read when creating a policy from query_id When creating a fleet or global policy from an existing query (via query_id) load the referenced query and authorize ActionRead on it before its fields are copied, in both fleet policies and global policies. --- changes/16291-policy-query-id-authz | 1 + server/service/global_policies.go | 10 +++ server/service/global_policies_test.go | 81 ++++++++++++++++++++++ server/service/team_policies.go | 10 +++ server/service/team_policies_test.go | 94 ++++++++++++++++++++++++++ 5 files changed, 196 insertions(+) create mode 100644 changes/16291-policy-query-id-authz diff --git a/changes/16291-policy-query-id-authz b/changes/16291-policy-query-id-authz new file mode 100644 index 0000000000..4243f1c83b --- /dev/null +++ b/changes/16291-policy-query-id-authz @@ -0,0 +1 @@ +* Improved query validation logic around policy creation. diff --git a/server/service/global_policies.go b/server/service/global_policies.go index 62a3c01698..a1232c8291 100644 --- a/server/service/global_policies.go +++ b/server/service/global_policies.go @@ -62,6 +62,16 @@ func (svc Service) NewGlobalPolicy(ctx context.Context, p fleet.PolicyPayload) ( }) } + if p.QueryID != nil { + query, err := svc.ds.Query(ctx, *p.QueryID) + if err != nil { + return nil, ctxerr.Wrap(ctx, err, "get query for policy") + } + if err := svc.authz.Authorize(ctx, query, fleet.ActionRead); err != nil { + return nil, err + } + } + if (len(p.LabelsIncludeAll) > 0 || len(p.LabelsExcludeAll) > 0 || len(p.LabelsIncludeAny) > 0 || len(p.LabelsExcludeAny) > 0) && !license.IsPremium(ctx) { return nil, fleet.ErrMissingLicense } diff --git a/server/service/global_policies_test.go b/server/service/global_policies_test.go index dfc798954c..7d6396599d 100644 --- a/server/service/global_policies_test.go +++ b/server/service/global_policies_test.go @@ -768,3 +768,84 @@ func TestResetPolicyEmitsActivity(t *testing.T) { require.Nil(t, act.TeamName) }) } + +func TestNewGlobalPolicyQueryIDAuth(t *testing.T) { + const ( + queryID = uint(99) + secretSQL = "SELECT secret FROM restricted;" + ) + + testCases := []struct { + name string + user *fleet.User + payload fleet.PolicyPayload + queryErr error + wantQueryLoaded bool + wantErr bool + }{ + { + name: "global admin from query_id loads and authorizes the query", + user: &fleet.User{ID: 1, GlobalRole: new(fleet.RoleAdmin)}, + payload: fleet.PolicyPayload{QueryID: new(queryID)}, + wantQueryLoaded: true, + }, + { + name: "global maintainer from query_id loads and authorizes the query", + user: &fleet.User{ID: 1, GlobalRole: new(fleet.RoleMaintainer)}, + payload: fleet.PolicyPayload{QueryID: new(queryID)}, + wantQueryLoaded: true, + }, + { + name: "global gitops from query_id loads and authorizes the query", + user: &fleet.User{ID: 1, GlobalRole: new(fleet.RoleGitOps)}, + payload: fleet.PolicyPayload{QueryID: new(queryID)}, + wantQueryLoaded: true, + }, + { + name: "no query_id does not load any query", + user: &fleet.User{ID: 1, GlobalRole: new(fleet.RoleAdmin)}, + payload: fleet.PolicyPayload{Name: "inline", Query: "SELECT 1;"}, + wantQueryLoaded: false, + }, + { + name: "missing referenced query fails", + user: &fleet.User{ID: 1, GlobalRole: new(fleet.RoleAdmin)}, + payload: fleet.PolicyPayload{QueryID: new(queryID)}, + queryErr: ¬FoundError{}, + wantQueryLoaded: true, + wantErr: true, + }, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + ds := new(mock.Store) + opts := &TestServerOpts{} + svc, baseCtx := newTestService(t, ds, nil, nil, opts) + opts.ActivityMock.NewActivityFunc = func(_ context.Context, _ *activity_api.User, _ activity_api.ActivityDetails) error { + return nil + } + + ds.QueryFunc = func(ctx context.Context, id uint) (*fleet.Query, error) { + require.Equal(t, queryID, id) + if tc.queryErr != nil { + return nil, tc.queryErr + } + return &fleet.Query{ID: id, Name: "referenced query", Query: secretSQL}, nil + } + ds.NewGlobalPolicyFunc = func(ctx context.Context, authorID *uint, args fleet.PolicyPayload) (*fleet.Policy, error) { + return &fleet.Policy{PolicyData: fleet.PolicyData{ID: 1, Name: "referenced query", Query: secretSQL}}, nil + } + + ctx := viewer.NewContext(baseCtx, viewer.Viewer{User: tc.user}) + + _, err := svc.NewGlobalPolicy(ctx, tc.payload) + if tc.wantErr { + require.Error(t, err) + } else { + require.NoError(t, err) + } + require.Equal(t, tc.wantQueryLoaded, ds.QueryFuncInvoked) + }) + } +} diff --git a/server/service/team_policies.go b/server/service/team_policies.go index 76c0817831..0e898fca0a 100644 --- a/server/service/team_policies.go +++ b/server/service/team_policies.go @@ -75,6 +75,16 @@ func (svc Service) NewTeamPolicy(ctx context.Context, teamID uint, tp fleet.NewT }) } + if p.QueryID != nil { + query, err := svc.ds.Query(ctx, *p.QueryID) + if err != nil { + return nil, ctxerr.Wrap(ctx, err, "get query for policy") + } + if err := svc.authz.Authorize(ctx, query, fleet.ActionRead); err != nil { + return nil, err + } + } + if (len(tp.LabelsIncludeAll) > 0 || len(tp.LabelsExcludeAll) > 0 || len(tp.LabelsIncludeAny) > 0 || len(tp.LabelsExcludeAny) > 0) && !license.IsPremium(ctx) { return nil, fleet.ErrMissingLicense } diff --git a/server/service/team_policies_test.go b/server/service/team_policies_test.go index 2e3c3b517a..de02fef22e 100644 --- a/server/service/team_policies_test.go +++ b/server/service/team_policies_test.go @@ -510,6 +510,100 @@ func TestPopulateSoftwareIconURLs(t *testing.T) { ) } +func TestNewTeamPolicyQueryIDAuth(t *testing.T) { + const ( + callerTeamID = uint(1) + otherTeamID = uint(2) + queryID = uint(99) + secretSQL = "SELECT secret FROM restricted;" + ) + + otherTeam := otherTeamID + callerTeam := callerTeamID + + testCases := []struct { + name string + user *fleet.User + queryTeamID *uint + shouldFail bool + }{ + { + name: "team admin references another team's query", + user: &fleet.User{ID: 1, Teams: []fleet.UserTeam{{Team: fleet.Team{ID: callerTeamID}, Role: fleet.RoleAdmin}}}, + queryTeamID: &otherTeam, + shouldFail: true, + }, + { + name: "team gitops references a global query", + user: &fleet.User{ID: 1, Teams: []fleet.UserTeam{{Team: fleet.Team{ID: callerTeamID}, Role: fleet.RoleGitOps}}}, + queryTeamID: nil, + shouldFail: true, + }, + { + name: "team admin references a global query", + user: &fleet.User{ID: 1, Teams: []fleet.UserTeam{{Team: fleet.Team{ID: callerTeamID}, Role: fleet.RoleAdmin}}}, + queryTeamID: nil, + shouldFail: false, + }, + { + name: "team admin references their own team's query", + user: &fleet.User{ID: 1, Teams: []fleet.UserTeam{{Team: fleet.Team{ID: callerTeamID}, Role: fleet.RoleAdmin}}}, + queryTeamID: &callerTeam, + shouldFail: false, + }, + { + name: "global admin references another team's query", + user: &fleet.User{ID: 1, GlobalRole: new(fleet.RoleAdmin)}, + queryTeamID: &otherTeam, + shouldFail: false, + }, + } + + for _, tc := range testCases { + t.Run(tc.name, func(t *testing.T) { + ds := new(mock.Store) + opts := &TestServerOpts{} + svc, baseCtx := newTestService(t, ds, nil, nil, opts) + opts.ActivityMock.NewActivityFunc = func(_ context.Context, _ *activity_api.User, _ activity_api.ActivityDetails) error { + return nil + } + + ds.QueryFunc = func(ctx context.Context, id uint) (*fleet.Query, error) { + require.Equal(t, queryID, id) + return &fleet.Query{ + ID: id, + TeamID: tc.queryTeamID, + Name: "referenced query", + Query: secretSQL, + }, nil + } + ds.NewTeamPolicyFunc = func(ctx context.Context, tID uint, authorID *uint, args fleet.PolicyPayload) (*fleet.Policy, error) { + return &fleet.Policy{ + PolicyData: fleet.PolicyData{ID: 1, TeamID: &callerTeam, Name: "referenced query", Query: secretSQL}, + }, nil + } + ds.TeamLiteFunc = func(ctx context.Context, tID uint) (*fleet.TeamLite, error) { + return &fleet.TeamLite{ID: tID}, nil + } + + ctx := viewer.NewContext(baseCtx, viewer.Viewer{User: tc.user}) + + _, err := svc.NewTeamPolicy(ctx, callerTeamID, fleet.NewTeamPolicyPayload{ + QueryID: new(queryID), + }) + + if tc.shouldFail { + require.Error(t, err) + var forbiddenError *authz.Forbidden + require.ErrorAs(t, err, &forbiddenError) + } else { + require.NoError(t, err) + } + require.True(t, ds.QueryFuncInvoked, "expected the referenced query to be loaded for a read authorization check") + }) + } +} + func checkAuthErr(t *testing.T, shouldFail bool, err error) { t.Helper() if shouldFail {