diff --git a/changes/policies-targeted-platforms-filter b/changes/policies-targeted-platforms-filter new file mode 100644 index 0000000000..cd598bd940 --- /dev/null +++ b/changes/policies-targeted-platforms-filter @@ -0,0 +1,2 @@ +* Added "Targeted platforms" column and platform filter dropdown to the Policies page. +* Added optional `platform` query parameter to `GET /api/v1/fleet/policies` and `GET /api/v1/fleet/fleets/{id}/policies` to filter policies by targeted platform. diff --git a/cmd/fleetctl/fleetctl/gitops_test.go b/cmd/fleetctl/fleetctl/gitops_test.go index 338888f2ca..d4736d7949 100644 --- a/cmd/fleetctl/fleetctl/gitops_test.go +++ b/cmd/fleetctl/fleetctl/gitops_test.go @@ -2071,11 +2071,11 @@ func TestGitOpsFullGlobal(t *testing.T) { policy.ID = 1 policy.Name = "Policy to delete" ds.ListTeamPoliciesFunc = func( - ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationFilter string, + ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationFilter string, platform string, ) (teamPolicies []*fleet.Policy, inheritedPolicies []*fleet.Policy, err error) { return nil, nil, nil } - ds.ListGlobalPoliciesFunc = func(ctx context.Context, opts fleet.ListOptions) ([]*fleet.Policy, error) { + ds.ListGlobalPoliciesFunc = func(ctx context.Context, opts fleet.ListOptions, platform string) ([]*fleet.Policy, error) { return []*fleet.Policy{&policy}, nil } ds.PoliciesByIDFunc = func(ctx context.Context, ids []uint) (map[uint]*fleet.Policy, error) { @@ -2463,7 +2463,7 @@ func TestGitOpsFullTeam(t *testing.T) { policy.Name = "Policy to delete" policy.TeamID = ptr.Uint(teamID) ds.ListTeamPoliciesFunc = func( - ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationFilter string, + ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationFilter string, platform string, ) (teamPolicies []*fleet.Policy, inheritedPolicies []*fleet.Policy, err error) { if teamID != 0 { return []*fleet.Policy{&policy}, nil, nil @@ -2788,9 +2788,11 @@ func TestGitOpsBasicGlobalAndTeam(t *testing.T) { require.ElementsMatch(t, names, []string{fleet.BuiltinLabelMacOS14Plus}) return map[string]uint{fleet.BuiltinLabelMacOS14Plus: 1}, nil } - ds.ListGlobalPoliciesFunc = func(ctx context.Context, opts fleet.ListOptions) ([]*fleet.Policy, error) { return nil, nil } + ds.ListGlobalPoliciesFunc = func(ctx context.Context, opts fleet.ListOptions, platform string) ([]*fleet.Policy, error) { + return nil, nil + } ds.ListTeamPoliciesFunc = func( - ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationFilter string, + ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationFilter string, platform string, ) (teamPolicies []*fleet.Policy, inheritedPolicies []*fleet.Policy, err error) { return nil, nil, nil } @@ -3127,9 +3129,11 @@ func TestGitOpsBasicGlobalAndNoTeam(t *testing.T) { require.ElementsMatch(t, names, []string{fleet.BuiltinLabelMacOS14Plus}) return map[string]uint{fleet.BuiltinLabelMacOS14Plus: 1}, nil } - ds.ListGlobalPoliciesFunc = func(ctx context.Context, opts fleet.ListOptions) ([]*fleet.Policy, error) { return nil, nil } + ds.ListGlobalPoliciesFunc = func(ctx context.Context, opts fleet.ListOptions, platform string) ([]*fleet.Policy, error) { + return nil, nil + } ds.ListTeamPoliciesFunc = func( - ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationFilter string, + ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationFilter string, platform string, ) (teamPolicies []*fleet.Policy, inheritedPolicies []*fleet.Policy, err error) { return nil, nil, nil } @@ -4699,7 +4703,7 @@ func TestGitOpsGlobalWebhooksAndTicketsEnabled(t *testing.T) { } return nil } - ds.ListGlobalPoliciesFunc = func(ctx context.Context, opts fleet.ListOptions) ([]*fleet.Policy, error) { + ds.ListGlobalPoliciesFunc = func(ctx context.Context, opts fleet.ListOptions, platform string) ([]*fleet.Policy, error) { return appliedPolicies, nil } @@ -4821,7 +4825,7 @@ func TestGitOpsFleetWebhooksAndTicketsEnabled(t *testing.T) { // Override ListTeamPolicies to return the applied policies with IDs. ds.ListTeamPoliciesFunc = func( - ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationFilter string, + ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationFilter string, platform string, ) (fleetPolicies []*fleet.Policy, inheritedPolicies []*fleet.Policy, err error) { return appliedPolicies, nil, nil } @@ -4914,7 +4918,7 @@ agent_options: // Track how many times ListTeamPolicies is called. listTeamPoliciesCalls := 0 ds.ListTeamPoliciesFunc = func( - ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationFilter string, + ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationFilter string, platform string, ) ([]*fleet.Policy, []*fleet.Policy, error) { listTeamPoliciesCalls++ return appliedPolicies, nil, nil diff --git a/cmd/fleetctl/fleetctl/testing_utils/testing_utils.go b/cmd/fleetctl/fleetctl/testing_utils/testing_utils.go index 5e2f5a6680..c8572d7abb 100644 --- a/cmd/fleetctl/fleetctl/testing_utils/testing_utils.go +++ b/cmd/fleetctl/fleetctl/testing_utils/testing_utils.go @@ -391,9 +391,11 @@ func SetupFullGitOpsPremiumServer(t *testing.T) (*mock.Store, **fleet.AppConfig, require.ElementsMatch(t, names, []string{fleet.BuiltinLabelMacOS14Plus}) return map[string]uint{fleet.BuiltinLabelMacOS14Plus: 1}, nil } - ds.ListGlobalPoliciesFunc = func(ctx context.Context, opts fleet.ListOptions) ([]*fleet.Policy, error) { return nil, nil } + ds.ListGlobalPoliciesFunc = func(ctx context.Context, opts fleet.ListOptions, platform string) ([]*fleet.Policy, error) { + return nil, nil + } ds.ListTeamPoliciesFunc = func( - ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationFilter string, + ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationFilter string, platform string, ) (teamPolicies []*fleet.Policy, inheritedPolicies []*fleet.Policy, err error) { return nil, nil, nil } diff --git a/cmd/fleetctl/fleetctl/testing_utils_test.go b/cmd/fleetctl/fleetctl/testing_utils_test.go index 2810243002..800c94ce04 100644 --- a/cmd/fleetctl/fleetctl/testing_utils_test.go +++ b/cmd/fleetctl/fleetctl/testing_utils_test.go @@ -152,11 +152,11 @@ func setupEmptyGitOpsMocks(ds *mock.Store) { } // Policies and queries - ds.ListGlobalPoliciesFunc = func(ctx context.Context, opts fleet.ListOptions) ([]*fleet.Policy, error) { + ds.ListGlobalPoliciesFunc = func(ctx context.Context, opts fleet.ListOptions, platform string) ([]*fleet.Policy, error) { return nil, nil } ds.ListTeamPoliciesFunc = func( - ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationFilter string, + ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationFilter string, platform string, ) ([]*fleet.Policy, []*fleet.Policy, error) { return nil, nil, nil } diff --git a/cmd/fleetctl/integrationtest/gitops/gitops_enterprise_integration_test.go b/cmd/fleetctl/integrationtest/gitops/gitops_enterprise_integration_test.go index 41c9580aec..19a3a0a42e 100644 --- a/cmd/fleetctl/integrationtest/gitops/gitops_enterprise_integration_test.go +++ b/cmd/fleetctl/integrationtest/gitops/gitops_enterprise_integration_test.go @@ -3970,7 +3970,7 @@ reports: team, err := s.DS.TeamByName(ctx, teamName) require.NoError(t, err) - pols, err := s.DS.ListMergedTeamPolicies(ctx, team.ID, fleet.ListOptions{}, "") + pols, err := s.DS.ListMergedTeamPolicies(ctx, team.ID, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, pols, 3) policyIDsByName := map[string]uint{} @@ -3987,7 +3987,7 @@ reports: "gitops", "--config", fleetctlConfig.Name(), "-f", globalFile, "-f", teamFile, })) - pols, err = s.DS.ListMergedTeamPolicies(ctx, team.ID, fleet.ListOptions{}, "") + pols, err = s.DS.ListMergedTeamPolicies(ctx, team.ID, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Empty(t, pols, "all policies should be removed after FMA installer is removed") @@ -4860,7 +4860,7 @@ settings: installer, err := s.DS.GetSoftwareInstallerMetadataByTeamAndTitleID(ctx, nil, titles[0].ID, false) require.NoError(t, err) - tmPols, err := s.DS.ListMergedTeamPolicies(ctx, 0, fleet.ListOptions{}, "") + tmPols, err := s.DS.ListMergedTeamPolicies(ctx, 0, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, tmPols, 1) require.Equal(t, "Install ruby", tmPols[0].Name) @@ -4880,7 +4880,7 @@ settings: installer, err = s.DS.GetSoftwareInstallerMetadataByTeamAndTitleID(ctx, &tm.ID, titles[0].ID, false) require.NoError(t, err) - tmPols, err = s.DS.ListMergedTeamPolicies(ctx, tm.ID, fleet.ListOptions{}, "") + tmPols, err = s.DS.ListMergedTeamPolicies(ctx, tm.ID, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, tmPols, 1) require.Equal(t, "Install team ruby", tmPols[0].Name) @@ -4934,7 +4934,7 @@ labels: s.assertRealRunOutput(t, fleetctltest.RunAppForTest(t, []string{"gitops", "--config", fleetctlConfig.Name(), "-f", fullFile.Name()})) // Verify policy, agent_options, controls, and reports were applied. - policies, err := s.DS.ListGlobalPolicies(ctx, fleet.ListOptions{}) + policies, err := s.DS.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, policies, 1) require.Equal(t, "Test Global Policy", policies[0].Name) @@ -4977,7 +4977,7 @@ org_settings: s.assertRealRunOutput(t, fleetctltest.RunAppForTest(t, []string{"gitops", "--config", fleetctlConfig.Name(), "-f", minimalFile.Name()})) // Verify policies were cleared. - policies, err = s.DS.ListGlobalPolicies(ctx, fleet.ListOptions{}) + policies, err = s.DS.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, policies, 0) @@ -5082,7 +5082,7 @@ policies: s.assertRealRunOutput(t, fleetctltest.RunAppForTest(t, []string{"gitops", "--config", fleetctlConfig.Name(), "-f", globalFile, "-f", teamFile})) // The global policy persisted both its include_any and exclude_all scopes. - globalPolicies, err := s.DS.ListGlobalPolicies(ctx, fleet.ListOptions{}) + globalPolicies, err := s.DS.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, globalPolicies, 1) gp := globalPolicies[0] @@ -5097,7 +5097,7 @@ policies: // The team policy persisted both its include_all and exclude_any scopes. tm, err := s.DS.TeamByName(ctx, fleetName) require.NoError(t, err) - teamPolicies, _, err := s.DS.ListTeamPolicies(ctx, tm.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + teamPolicies, _, err := s.DS.ListTeamPolicies(ctx, tm.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, teamPolicies, 1) tp := teamPolicies[0] @@ -5200,7 +5200,7 @@ labels: fl, err := s.DS.TeamByName(ctx, fleetName) require.NoError(t, err) - flPols, err := s.DS.ListMergedTeamPolicies(ctx, fl.ID, fleet.ListOptions{}, "") + flPols, err := s.DS.ListMergedTeamPolicies(ctx, fl.ID, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, flPols, 1) require.Equal(t, "Test Fleet Policy", flPols[0].Name) @@ -5241,7 +5241,7 @@ name: %s })) // Verify policies were cleared. - flPols, err = s.DS.ListMergedTeamPolicies(ctx, fl.ID, fleet.ListOptions{}, "") + flPols, err = s.DS.ListMergedTeamPolicies(ctx, fl.ID, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, flPols, 0) diff --git a/docs/REST API/rest-api.md b/docs/REST API/rest-api.md index bce4532966..07a13b8721 100644 --- a/docs/REST API/rest-api.md +++ b/docs/REST API/rest-api.md @@ -9227,6 +9227,7 @@ For example, a policy might ask “Is Gatekeeper enabled on macOS devices?“ Th | order_key | string | query | What to order results by. Allowed fields are `id`, `name`, `team_id`, `created_at`, `updated_at`, `failing_host_count`, and `passing_host_count`. | | order_direction | string | query | **Requires `order_key`**. The direction of the order given the order key. Options include `"asc"` and `"desc"`. Default is `"asc"`. | | after | string | query | The value to get results after. This needs `order_key` defined, as that's the column that would be used. | +| platform | string | query | Filters policies by targeted platform. Accepts `"darwin"`, `"windows"`, `"linux"`, or `"chrome"`. Policies that target all platforms (empty `platform` field) are always included. | #### Example @@ -9302,6 +9303,7 @@ _Available in Fleet Premium_ | order_direction | string | query | **Requires `order_key`**. The direction of the order given the order key. Options include `"asc"` and `"desc"`. Default is `"asc"`. | | after | string | query | The value to get results after. This needs `order_key` defined, as that's the column that would be used. | | automation_type | string | query | Filters by automation type. Supported values are "software", "scripts", "calendar", "conditional_access", and "other". | +| platform | string | query | Filters policies by targeted platform. Accepts `"darwin"`, `"windows"`, `"linux"`, or `"chrome"`. Policies that target all platforms (empty `platform` field) are always included. | #### Example (default usage) @@ -9500,6 +9502,7 @@ _Available in Fleet Premium_ | Name | Type | In | Description | | ------------------ | ------- | ---- | ------------------------------------------------------------------------------------------------------------- | | query | string | query | Search query keywords. Searchable fields include `name`. | +| platform | string | query | Filters policies by targeted platform. Accepts `"darwin"`, `"windows"`, `"linux"`, or `"chrome"`. Policies that target all platforms (empty `platform` field) are always included. | #### Example @@ -9530,6 +9533,7 @@ _Available in Fleet Premium_ | query | string | query | Search query keywords. Searchable fields include `name`. | | merge_inherited | boolean | query | If `true`, will include inherited ("All fleets") policies in the count. | | automation_type | string | query | Filters by automation type. Supported values are "software", "scripts", "calendar", "conditional_access", and "other". | +| platform | string | query | Filters policies by targeted platform. Accepts `"darwin"`, `"windows"`, `"linux"`, or `"chrome"`. Policies that target all platforms (empty `platform` field) are always included. | #### Example diff --git a/frontend/pages/policies/ManagePoliciesPage/ManagePoliciesPage.tsx b/frontend/pages/policies/ManagePoliciesPage/ManagePoliciesPage.tsx index 4c627d080a..86250fd1f5 100644 --- a/frontend/pages/policies/ManagePoliciesPage/ManagePoliciesPage.tsx +++ b/frontend/pages/policies/ManagePoliciesPage/ManagePoliciesPage.tsx @@ -31,6 +31,7 @@ import { APP_CONTEXT_ALL_TEAMS_ID, ITeamConfig, } from "interfaces/team"; +import { isQueryablePlatform } from "interfaces/platform"; import configAPI from "services/entities/config"; import globalPoliciesAPI, { @@ -83,6 +84,7 @@ interface IManagePoliciesPageProps { order_direction?: "asc" | "desc"; page?: string; automation_type?: AutomationType; + platform?: string; manage_automations?: string; }; search: string; @@ -179,6 +181,9 @@ const ManagePolicyPage = ({ DEFAULT_SORT_DIRECTION)(); const page = queryParams && queryParams.page ? parseInt(queryParams?.page, 10) : 0; + const targetedPlatformParam = isQueryablePlatform(queryParams?.platform) + ? queryParams?.platform + : undefined; const initialAutomationFilter = (() => { const automationQueryParam = queryParams.automation_type; @@ -269,6 +274,7 @@ const ManagePolicyPage = ({ orderDirection: sortDirection, orderKey: sortHeader, automationType: automationFilter as GlobalPoliciesAutomationType, + platform: targetedPlatformParam, }, ], ({ queryKey }) => { @@ -292,6 +298,7 @@ const ManagePolicyPage = ({ scope: "policiesCount", query: !isAllTeamsSelected ? "" : searchQuery, automationType: automationFilter as GlobalPoliciesAutomationType, + platform: targetedPlatformParam, }, ], ({ queryKey }) => globalPoliciesAPI.getCount(queryKey[0]), @@ -328,6 +335,7 @@ const ManagePolicyPage = ({ // no teams does inherit mergeInherited: true, automationType: automationFilter as AutomationType, + platform: targetedPlatformParam, }, ], ({ queryKey }) => { @@ -357,6 +365,7 @@ const ManagePolicyPage = ({ teamId: teamIdForApi || 0, // TODO: Fix number/undefined type mergeInherited: true, automationType: automationFilter as AutomationType, + platform: targetedPlatformParam, }, ], ({ queryKey }) => teamPoliciesAPI.getCount(queryKey[0]), @@ -683,7 +692,10 @@ const ManagePolicyPage = ({ const hide = isFetchingCount || policiesErrors || - (!policyResults && searchQuery === "" && !automationFilter); + (!policyResults && + searchQuery === "" && + !automationFilter && + !targetedPlatformParam); if (hide) { return null; @@ -765,7 +777,10 @@ const ManagePolicyPage = ({ ? globalPoliciesCount : teamPoliciesCountMergeInherited; const isTrulyEmpty = - (policiesCount ?? 0) === 0 && searchQuery === "" && !automationFilter; + (policiesCount ?? 0) === 0 && + searchQuery === "" && + !automationFilter && + !targetedPlatformParam; // No team ID = All fleets → only show "all" and "other" options const optionsForTeam = teamIdForApi @@ -823,7 +838,10 @@ const ManagePolicyPage = ({ page={page} onQueryChange={onQueryChange} customControl={renderAutomationFilter} - isFiltered={!!automationFilter} + isFiltered={!!automationFilter || !!targetedPlatformParam} + router={router} + queryParams={queryParams} + platform={targetedPlatformParam} otherAutomationType={otherAutomationType} onOpenManageAutomationsModal={ canAddOrDeletePolicies ? onOpenManageAutomationsModal : undefined @@ -867,7 +885,10 @@ const ManagePolicyPage = ({ page={page} onQueryChange={onQueryChange} customControl={renderAutomationFilter} - isFiltered={!!automationFilter} + isFiltered={!!automationFilter || !!targetedPlatformParam} + router={router} + queryParams={queryParams} + platform={targetedPlatformParam} otherAutomationType={otherAutomationType} onOpenManageAutomationsModal={ canAddOrDeletePolicies ? onOpenManageAutomationsModal : undefined diff --git a/frontend/pages/policies/ManagePoliciesPage/components/PoliciesTable/PoliciesTable.tests.tsx b/frontend/pages/policies/ManagePoliciesPage/components/PoliciesTable/PoliciesTable.tests.tsx index 08c7aa6db0..5aee38f8a1 100644 --- a/frontend/pages/policies/ManagePoliciesPage/components/PoliciesTable/PoliciesTable.tests.tsx +++ b/frontend/pages/policies/ManagePoliciesPage/components/PoliciesTable/PoliciesTable.tests.tsx @@ -1,12 +1,14 @@ import React from "react"; import { screen, waitFor } from "@testing-library/react"; import { noop } from "lodash"; -import { createCustomRenderer } from "test/test-utils"; +import { createCustomRenderer, createMockRouter } from "test/test-utils"; import createMockUser from "__mocks__/userMock"; import createMockPolicy from "__mocks__/policyMock"; import PoliciesTable from "./PoliciesTable"; +const mockRouter = createMockRouter(); + describe("Policies table", () => { it("Renders the page-wide empty state when no policies are present (free tier)", async () => { const render = createCustomRenderer({ @@ -28,6 +30,7 @@ describe("Policies table", () => { searchQuery="" page={0} onQueryChange={noop} + router={mockRouter} renderPoliciesCount={() => null} count={0} /> @@ -59,6 +62,7 @@ describe("Policies table", () => { searchQuery="" page={0} onQueryChange={noop} + router={mockRouter} renderPoliciesCount={() => null} count={0} /> @@ -92,6 +96,7 @@ describe("Policies table", () => { searchQuery="" page={0} onQueryChange={noop} + router={mockRouter} renderPoliciesCount={() => null} count={0} /> @@ -125,6 +130,7 @@ describe("Policies table", () => { onQueryChange={noop} renderPoliciesCount={() => null} count={0} + router={mockRouter} /> ); @@ -158,6 +164,7 @@ describe("Policies table", () => { onQueryChange={noop} renderPoliciesCount={() => null} count={0} + router={mockRouter} /> ); @@ -188,6 +195,7 @@ describe("Policies table", () => { searchQuery="shouldn't match anything" page={0} onQueryChange={noop} + router={mockRouter} renderPoliciesCount={() => null} count={0} /> @@ -221,6 +229,7 @@ describe("Policies table", () => { searchQuery="" page={0} onQueryChange={noop} + router={mockRouter} renderPoliciesCount={() => null} count={[testCriticalPolicy].length} /> @@ -260,6 +269,7 @@ describe("Policies table", () => { searchQuery="" page={0} onQueryChange={noop} + router={mockRouter} renderPoliciesCount={() => null} count={[testInheritedPolicy].length} /> @@ -299,6 +309,7 @@ describe("Policies table", () => { searchQuery="" page={0} onQueryChange={noop} + router={mockRouter} renderPoliciesCount={() => null} count={[testGlobalPolicy].length} /> @@ -341,6 +352,7 @@ describe("Policies table", () => { searchQuery="" page={0} onQueryChange={noop} + router={mockRouter} renderPoliciesCount={() => null} canAddOrDeletePolicies hasPoliciesToDelete @@ -389,6 +401,7 @@ describe("Policies table", () => { searchQuery="" page={0} onQueryChange={noop} + router={mockRouter} renderPoliciesCount={() => null} count={1} /> @@ -420,6 +433,7 @@ describe("Policies table", () => { searchQuery="" page={0} onQueryChange={noop} + router={mockRouter} renderPoliciesCount={() => null} count={1} /> @@ -428,6 +442,81 @@ describe("Policies table", () => { expect(screen.queryByText("Patch")).not.toBeInTheDocument(); }); + it("Renders the Targeted platforms column using the policy's platform field", () => { + const render = createCustomRenderer({ + context: { + app: { + isGlobalAdmin: true, + currentUser: createMockUser(), + }, + }, + }); + + const policyWithAllPlatforms = createMockPolicy({ + id: 100, + name: "cross-platform policy", + platform: "", + }); + const policyWithDarwin = createMockPolicy({ + id: 101, + name: "macOS policy", + platform: "darwin", + }); + + render( + null} + count={2} + /> + ); + + expect(screen.getByText("Targeted platforms")).toBeInTheDocument(); + expect(screen.getByTestId("darwin-icon")).toBeInTheDocument(); + expect(screen.queryByTestId("windows-icon")).not.toBeInTheDocument(); + expect(screen.queryByTestId("linux-icon")).not.toBeInTheDocument(); + expect(screen.queryByTestId("chrome-icon")).not.toBeInTheDocument(); + }); + + it("Renders the platform filter dropdown when the table is searchable", () => { + const render = createCustomRenderer({ + context: { + app: { + isGlobalAdmin: true, + currentUser: createMockUser(), + }, + }, + }); + + render( + null} + count={1} + /> + ); + + expect(screen.getByText("All platforms")).toBeInTheDocument(); + }); + it("Renders the Automations column with correct values", () => { const render = createCustomRenderer({ context: { @@ -464,6 +553,7 @@ describe("Policies table", () => { searchQuery="" page={0} onQueryChange={noop} + router={mockRouter} renderPoliciesCount={() => null} count={2} /> diff --git a/frontend/pages/policies/ManagePoliciesPage/components/PoliciesTable/PoliciesTable.tsx b/frontend/pages/policies/ManagePoliciesPage/components/PoliciesTable/PoliciesTable.tsx index 6f2d9eae2c..bea62a7246 100644 --- a/frontend/pages/policies/ManagePoliciesPage/components/PoliciesTable/PoliciesTable.tsx +++ b/frontend/pages/policies/ManagePoliciesPage/components/PoliciesTable/PoliciesTable.tsx @@ -1,13 +1,21 @@ -import React, { useContext } from "react"; +import React, { useCallback, useContext } from "react"; +import { InjectedRouter } from "react-router"; +import { SingleValue } from "react-select-5"; +import PATHS from "router/paths"; import { AppContext } from "context/app"; import { IPolicyStats, OtherAutomationType } from "interfaces/policy"; import { ITeamSummary, APP_CONTEXT_ALL_TEAMS_ID } from "interfaces/team"; import { IEmptyStateProps } from "interfaces/empty_state"; +import { SelectedPlatform } from "interfaces/platform"; +import { getNextLocationPath } from "utilities/helpers"; import Button from "components/buttons/Button"; import TableContainer from "components/TableContainer"; import { ITableQueryData } from "components/TableContainer/TableContainer"; +import DropdownWrapper from "components/forms/fields/DropdownWrapper"; +import { CustomOptionType } from "components/forms/fields/DropdownWrapper/DropdownWrapper"; import EmptyState from "components/EmptyState"; +import { AutomationType } from "services/entities/team_policies"; import { generateTableHeaders, generateDataSet } from "./PoliciesTableConfig"; import { DEFAULT_SORT_COLUMN, @@ -22,6 +30,34 @@ const isLastPage = (count: number, pageSize: number, page: number) => { const baseClass = "policies-table"; +const PLATFORM_FILTER_OPTIONS = [ + { + disabled: false, + label: "All platforms", + value: "all", + }, + { + disabled: false, + label: "macOS", + value: "darwin", + }, + { + disabled: false, + label: "Windows", + value: "windows", + }, + { + disabled: false, + label: "Linux", + value: "linux", + }, + { + disabled: false, + label: "ChromeOS", + value: "chrome", + }, +]; + interface IPoliciesTableProps { policiesList: IPolicyStats[]; isLoading: boolean; @@ -41,6 +77,17 @@ interface IPoliciesTableProps { count: number; customControl?: () => JSX.Element | null; isFiltered?: boolean; + router: InjectedRouter; + queryParams?: { + fleet_id?: string; + query?: string; + order_key?: string; + order_direction?: "asc" | "desc"; + page?: string; + automation_type?: AutomationType; + platform?: string; + }; + platform?: SelectedPlatform; otherAutomationType?: OtherAutomationType; onOpenManageAutomationsModal?: (policy: IPolicyStats) => void; } @@ -64,11 +111,47 @@ const PoliciesTable = ({ count, customControl, isFiltered, + router, + queryParams, + platform = "all", otherAutomationType, onOpenManageAutomationsModal, }: IPoliciesTableProps): JSX.Element => { const { config } = useContext(AppContext); + const handlePlatformFilterDropdownChange = useCallback( + (selectedTargetedPlatform: SingleValue) => { + router.push( + getNextLocationPath({ + pathPrefix: PATHS.MANAGE_POLICIES, + queryParams: { + ...queryParams, + page: 0, + platform: + selectedTargetedPlatform?.value === "all" + ? undefined + : selectedTargetedPlatform?.value, + }, + }) + ); + }, + [queryParams, router] + ); + + const renderPlatformDropdown = useCallback(() => { + return ( + + ); + }, [platform, handlePlatformFilterDropdownChange]); + const isAllFleets = isPremiumTier && (currentTeam?.id === null || currentTeam?.id === APP_CONTEXT_ALL_TEAMS_ID); @@ -101,6 +184,15 @@ const PoliciesTable = ({ const isTrulyEmpty = policiesList?.length === 0 && searchQuery === "" && !isFiltered; + const combinedCustomControl = () => { + return ( +
+ {customControl?.()} + {renderPlatformDropdown()} +
+ ); + }; + const isPrimoMode = config?.partnerships?.enable_primo || false; const viewingTeamPolicies = currentTeam?.id !== undefined && @@ -164,7 +256,8 @@ const PoliciesTable = ({ inputPlaceHolder="Search by name" searchable disableSearch={isTrulyEmpty} - customControl={customControl} + customControl={combinedCustomControl} + selectedDropdownFilter={platform} /> ); diff --git a/frontend/pages/policies/ManagePoliciesPage/components/PoliciesTable/PoliciesTableConfig.tsx b/frontend/pages/policies/ManagePoliciesPage/components/PoliciesTable/PoliciesTableConfig.tsx index 02fccff731..3e28ce7d67 100644 --- a/frontend/pages/policies/ManagePoliciesPage/components/PoliciesTable/PoliciesTableConfig.tsx +++ b/frontend/pages/policies/ManagePoliciesPage/components/PoliciesTable/PoliciesTableConfig.tsx @@ -8,11 +8,16 @@ import classnames from "classnames"; import Checkbox from "components/forms/fields/Checkbox"; import HeaderCell from "components/TableContainer/DataTable/HeaderCell"; import LinkCell from "components/TableContainer/DataTable/LinkCell/LinkCell"; +import PlatformCell from "components/TableContainer/DataTable/PlatformCell"; import TooltipTruncatedTextCell from "components/TableContainer/DataTable/TooltipTruncatedTextCell"; import TooltipWrapper from "components/TooltipWrapper"; import Icon from "components/Icon"; import Graphic from "components/Graphic"; import SoftwareIcon from "pages/SoftwarePage/components/icons/SoftwareIcon"; +import { + CommaSeparatedPlatformString, + isQueryablePlatform, +} from "interfaces/platform"; import { IPolicyStats, OtherAutomationType } from "interfaces/policy"; import PATHS from "router/paths"; @@ -56,9 +61,20 @@ interface ICellProps { }; } +interface IPlatformCellProps { + cell: { + value: CommaSeparatedPlatformString; + }; + row: { + original: IPolicyStats; + }; +} + interface IDataColumn { Header: ((props: IHeaderProps) => JSX.Element) | string; - Cell: (props: ICellProps) => JSX.Element; + Cell: + | ((props: ICellProps) => JSX.Element) + | ((props: IPlatformCellProps) => JSX.Element); id?: string; title?: string; accessor?: string; @@ -300,6 +316,19 @@ const generateTableHeaders = ( }, sortType: "caseInsensitive", }, + { + title: "Targeted platforms", + Header: "Targeted platforms", + disableSortBy: true, + accessor: "platform", + Cell: (cellProps: IPlatformCellProps): JSX.Element => { + const platforms = cellProps.cell.value + .split(",") + .map((s) => s.trim()) + .filter(isQueryablePlatform); + return ; + }, + }, { title: "Automations", Header: "Automations", diff --git a/frontend/pages/policies/ManagePoliciesPage/components/PoliciesTable/_styles.scss b/frontend/pages/policies/ManagePoliciesPage/components/PoliciesTable/_styles.scss index 988e517ab1..502bd176ae 100644 --- a/frontend/pages/policies/ManagePoliciesPage/components/PoliciesTable/_styles.scss +++ b/frontend/pages/policies/ManagePoliciesPage/components/PoliciesTable/_styles.scss @@ -28,6 +28,58 @@ display: flex; justify-content: center; } + + &__filter-dropdowns { + display: flex; + align-items: center; + gap: $pad-medium; + } + + &__platform-dropdown { + flex-shrink: 0; + width: 200px; + } + + // Allow horizontal scrolling when the table is wider than its container so + // body cells don't overflow past the table boundary on smaller screens. + .table-container__data-table-block .data-table-block .data-table__wrapper { + overflow-x: auto; + } + + // Stack the search bar above the filter dropdowns on smaller screens so + // they don't get scrunched together (matches the Hosts page pattern). + .table-container { + @media (max-width: $table-controls-break) { + &__header { + flex-direction: column; + } + + &__header-left { + order: 2; + flex-direction: column; + align-items: stretch; + + .results-count { + order: 2; + } + + .controls { + order: -2; + + .policies-table__filter-dropdowns { + .form-field--dropdown, + .policies-table__platform-dropdown { + flex: 1; + } + } + } + } + + &__search { + align-self: start; + } + } + } } .automations__cell-content { diff --git a/frontend/pages/queries/ManageQueriesPage/components/QueriesTable/QueriesTable.tsx b/frontend/pages/queries/ManageQueriesPage/components/QueriesTable/QueriesTable.tsx index 36f2ad7a94..fa841e05cf 100644 --- a/frontend/pages/queries/ManageQueriesPage/components/QueriesTable/QueriesTable.tsx +++ b/frontend/pages/queries/ManageQueriesPage/components/QueriesTable/QueriesTable.tsx @@ -266,6 +266,7 @@ const QueriesTable = ({ options={PLATFORM_FILTER_OPTIONS} onChange={handlePlatformFilterDropdownChange} variant="table-filter" + iconName="filter-alt" isDisabled={isTrulyEmpty} /> ); diff --git a/frontend/services/entities/global_policies.ts b/frontend/services/entities/global_policies.ts index 25e414b436..0ae9a5c07b 100644 --- a/frontend/services/entities/global_policies.ts +++ b/frontend/services/entities/global_policies.ts @@ -7,6 +7,7 @@ import { ILoadAllPoliciesResponse, IPoliciesCountResponse, } from "interfaces/policy"; +import { QueryablePlatform } from "interfaces/platform"; import { buildQueryStringFromParams, convertParamsToSnakeCase, @@ -25,6 +26,8 @@ export interface IGlobalPoliciesApiQueryParams { orderDirection?: "asc" | "desc"; query?: string; automationType?: GlobalPoliciesAutomationType; + /** Targeted platform to filter policies by. */ + platform?: QueryablePlatform; } export interface IPoliciesQueryKey extends IGlobalPoliciesApiQueryParams { @@ -32,7 +35,10 @@ export interface IPoliciesQueryKey extends IGlobalPoliciesApiQueryParams { } export interface IPoliciesCountQueryKey - extends Pick { + extends Pick< + IGlobalPoliciesApiQueryParams, + "query" | "automationType" | "platform" + > { scope: "policiesCount"; } @@ -76,6 +82,7 @@ export default { orderDirection: orderDir = ORDER_DIRECTION, query, automationType, + platform, }: IGlobalPoliciesApiQueryParams): Promise => { const { GLOBAL_POLICIES } = endpoints; @@ -86,6 +93,7 @@ export default { orderDirection: orderDir, query, automationType, + platform, }; const snakeCaseParams = convertParamsToSnakeCase(queryParams); @@ -97,15 +105,17 @@ export default { getCount: ({ query, automationType, + platform, }: Pick< IGlobalPoliciesApiQueryParams, - "query" | "automationType" + "query" | "automationType" | "platform" >): Promise => { const { GLOBAL_POLICIES } = endpoints; const path = `${GLOBAL_POLICIES}/count`; const queryParams = { query, automationType, + platform, }; const snakeCaseParams = convertParamsToSnakeCase(queryParams); const queryString = buildQueryStringFromParams(snakeCaseParams); diff --git a/frontend/services/entities/team_policies.ts b/frontend/services/entities/team_policies.ts index 06d7b8077a..c739253105 100644 --- a/frontend/services/entities/team_policies.ts +++ b/frontend/services/entities/team_policies.ts @@ -9,6 +9,7 @@ import { IPoliciesCountResponse, ILoadTeamPolicyResponse, } from "interfaces/policy"; +import { QueryablePlatform } from "interfaces/platform"; import { API_NO_TEAM_ID } from "interfaces/team"; import { buildQueryStringFromParams, QueryParams } from "utilities/url"; import { GlobalPoliciesAutomationType } from "./global_policies"; @@ -27,6 +28,8 @@ interface IPoliciesApiQueryParams { orderDirection?: "asc" | "desc"; query?: string; automationType?: AutomationType | GlobalPoliciesAutomationType; + /** Targeted platform to filter policies by. */ + platform?: QueryablePlatform; } export interface IPoliciesApiParams extends IPoliciesApiQueryParams { @@ -41,7 +44,7 @@ export interface ITeamPoliciesQueryKey extends IPoliciesApiParams { export interface ITeamPoliciesCountQueryKey extends Pick< IPoliciesApiParams, - "query" | "teamId" | "mergeInherited" | "automationType" + "query" | "teamId" | "mergeInherited" | "automationType" | "platform" > { scope: "teamPoliciesCountMergeInherited" | "teamPoliciesCount"; } @@ -51,6 +54,7 @@ export interface IPoliciesCountApiParams { query?: string; mergeInherited?: boolean; automationType?: AutomationType; + platform?: QueryablePlatform; } const ORDER_KEY = "name"; @@ -179,6 +183,7 @@ export default { query, mergeInherited, automationType, + platform, }: IPoliciesApiParams): Promise => { const { TEAMS } = endpoints; @@ -190,6 +195,7 @@ export default { query, mergeInherited, automationType, + platform, }; const snakeCaseParams = convertParamsToSnakeCase(queryParams); @@ -202,9 +208,10 @@ export default { teamId, mergeInherited = true, automationType, + platform, }: Pick< IPoliciesCountApiParams, - "query" | "teamId" | "mergeInherited" | "automationType" + "query" | "teamId" | "mergeInherited" | "automationType" | "platform" >): Promise => { const { TEAM_POLICIES } = endpoints; const path = `${TEAM_POLICIES(teamId)}/count`; @@ -212,6 +219,7 @@ export default { query, mergeInherited, automationType, + platform, }; const snakeCaseParams = convertParamsToSnakeCase(queryParams); const queryString = buildQueryStringFromParams(snakeCaseParams); diff --git a/server/datastore/mysql/policies.go b/server/datastore/mysql/policies.go index f71ea7073a..8493e583c4 100644 --- a/server/datastore/mysql/policies.go +++ b/server/datastore/mysql/policies.go @@ -929,8 +929,17 @@ WHERE return nil } -func (ds *Datastore) ListGlobalPolicies(ctx context.Context, opts fleet.ListOptions) ([]*fleet.Policy, error) { - return listPoliciesDB(ctx, ds.reader(ctx), nil, opts, "", nil) +func (ds *Datastore) ListGlobalPolicies(ctx context.Context, opts fleet.ListOptions, platform string) ([]*fleet.Policy, error) { + filterClause, filterArgs := platformFilterClause(platform) + return listPoliciesDB(ctx, ds.reader(ctx), nil, opts, filterClause, filterArgs) +} + +func platformFilterClause(platform string) (string, []any) { + platform = strings.ReplaceAll(platform, " ", "") + if platform == "" { + return "", nil + } + return " AND (p.platforms = '' OR FIND_IN_SET(?, REPLACE(p.platforms, ' ', '')))", []any{platform} } // returns the list of policies associated with the provided teamID, or the @@ -987,7 +996,7 @@ func listPoliciesDB(ctx context.Context, q sqlx.QueryerContext, teamID *uint, op // getInheritedPoliciesForTeam returns the list of global policies with the // passing and failing host counts for the provided teamID -func getInheritedPoliciesForTeam(ctx context.Context, q sqlx.QueryerContext, teamID uint, opts fleet.ListOptions) ([]*fleet.Policy, error) { +func getInheritedPoliciesForTeam(ctx context.Context, q sqlx.QueryerContext, teamID uint, opts fleet.ListOptions, platform string) ([]*fleet.Policy, error) { var args []interface{} query := ` @@ -1006,6 +1015,11 @@ func getInheritedPoliciesForTeam(ctx context.Context, q sqlx.QueryerContext, tea args = append(args, teamID) + if platformClause, platformArgs := platformFilterClause(platform); platformClause != "" { + query += platformClause + args = append(args, platformArgs...) + } + // We must normalize the name for full Unicode support (Unicode equivalence). match := norm.NFC.String(opts.MatchQuery) query, args = searchLike(query, args, match, policySearchColumns...) @@ -1029,7 +1043,7 @@ func getInheritedPoliciesForTeam(ctx context.Context, q sqlx.QueryerContext, tea // CountPolicies returns the total number of team policies. // If teamID is nil, it returns the total number of global policies. -func (ds *Datastore) CountPolicies(ctx context.Context, teamID *uint, matchQuery string, automationType string) (int, error) { +func (ds *Datastore) CountPolicies(ctx context.Context, teamID *uint, matchQuery string, automationType string, platform string) (int, error) { var ( query string args []interface{} @@ -1055,6 +1069,11 @@ func (ds *Datastore) CountPolicies(ctx context.Context, teamID *uint, matchQuery } } + if platformClause, platformArgs := platformFilterClause(platform); platformClause != "" { + query += platformClause + args = append(args, platformArgs...) + } + // We must normalize the name for full Unicode support (Unicode equivalence). match := norm.NFC.String(matchQuery) query, args = searchLike(query, args, match, policySearchColumns...) @@ -1067,7 +1086,7 @@ func (ds *Datastore) CountPolicies(ctx context.Context, teamID *uint, matchQuery return count, nil } -func (ds *Datastore) CountMergedTeamPolicies(ctx context.Context, teamID uint, matchQuery string, automationType string) (int, error) { +func (ds *Datastore) CountMergedTeamPolicies(ctx context.Context, teamID uint, matchQuery string, automationType string, platform string) (int, error) { var args []interface{} query := `SELECT count(*) FROM policies p WHERE (p.team_id = ? OR p.team_id IS NULL)` @@ -1083,6 +1102,11 @@ func (ds *Datastore) CountMergedTeamPolicies(ctx context.Context, teamID uint, m args = append(args, filterArgs...) } + if platformClause, platformArgs := platformFilterClause(platform); platformClause != "" { + query += platformClause + args = append(args, platformArgs...) + } + // We must normalize the name for full Unicode support (Unicode equivalence). match := norm.NFC.String(matchQuery) query, args = searchLike(query, args, match, policySearchColumns...) @@ -1430,25 +1454,30 @@ func newTeamPolicy(ctx context.Context, db sqlx.ExtContext, teamID uint, authorI return policyDB(ctx, db, policyID, &teamID) } -func (ds *Datastore) ListTeamPolicies(ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationType string) (teamPolicies, inheritedPolicies []*fleet.Policy, err error) { +func (ds *Datastore) ListTeamPolicies(ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationType string, platform string) (teamPolicies, inheritedPolicies []*fleet.Policy, err error) { filterClause, filterArgs, err := ds.createAutomationClause(ctx, automationType, teamID) if err != nil { return nil, nil, ctxerr.Wrap(ctx, err, "build automation filter clause") } + if platformClause, platformArgs := platformFilterClause(platform); platformClause != "" { + filterClause += platformClause + filterArgs = append(filterArgs, platformArgs...) + } + teamPolicies, err = listPoliciesDB(ctx, ds.reader(ctx), &teamID, opts, filterClause, filterArgs) if err != nil { return nil, nil, err } // get inherited (global) policies with counts of hosts for that team - inheritedPolicies, err = getInheritedPoliciesForTeam(ctx, ds.reader(ctx), teamID, iopts) + inheritedPolicies, err = getInheritedPoliciesForTeam(ctx, ds.reader(ctx), teamID, iopts, platform) if err != nil { return nil, nil, err } return teamPolicies, inheritedPolicies, err } -func (ds *Datastore) ListMergedTeamPolicies(ctx context.Context, teamID uint, opts fleet.ListOptions, automationType string) ([]*fleet.Policy, error) { +func (ds *Datastore) ListMergedTeamPolicies(ctx context.Context, teamID uint, opts fleet.ListOptions, automationType string, platform string) ([]*fleet.Policy, error) { var args []interface{} automationFilter, filterArgs, err := ds.createAutomationClause(ctx, automationType, teamID) @@ -1456,6 +1485,8 @@ func (ds *Datastore) ListMergedTeamPolicies(ctx context.Context, teamID uint, op return nil, ctxerr.Wrap(ctx, err, "build automation filter clause") } + platformClause, platformArgs := platformFilterClause(platform) + query := fmt.Sprintf(` SELECT `+policyCols+`, @@ -1470,12 +1501,16 @@ func (ds *Datastore) ListMergedTeamPolicies(ctx context.Context, teamID uint, op AND (p.team_id IS NOT NULL OR ps.inherited_team_id = ?) WHERE (p.team_id = ? OR p.team_id IS NULL) %s - `, automationFilter) + %s + `, automationFilter, platformClause) args = append(args, teamID, teamID) if len(filterArgs) > 0 { args = append(args, filterArgs...) } + if len(platformArgs) > 0 { + args = append(args, platformArgs...) + } // We must normalize the name for full Unicode support (Unicode equivalence). match := norm.NFC.String(opts.MatchQuery) @@ -2619,7 +2654,7 @@ func (ds *Datastore) UpdateHostPolicyCounts(ctx context.Context) error { } if hasTeams { - globalPolicies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + globalPolicies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") if err != nil { return ctxerr.Wrap(ctx, err, "list global policies") } diff --git a/server/datastore/mysql/policies_test.go b/server/datastore/mysql/policies_test.go index 950338414d..3450d2531d 100644 --- a/server/datastore/mysql/policies_test.go +++ b/server/datastore/mysql/policies_test.go @@ -65,6 +65,7 @@ func TestPolicies(t *testing.T) { {"TestUpdatePolicyHostCounts", testUpdatePolicyHostCounts}, {"TestCachedPolicyCountDeletesOnPolicyChange", testCachedPolicyCountDeletesOnPolicyChange}, {"TestPoliciesListOptions", testPoliciesListOptions}, + {"TestPoliciesPlatformFilter", testPoliciesPlatformFilter}, {"TestPoliciesNameUnicode", testPoliciesNameUnicode}, {"TestPoliciesNameEmoji", testPoliciesNameEmoji}, {"TestPoliciesNameSort", testPoliciesNameSort}, @@ -145,7 +146,7 @@ func testPoliciesNewGlobalPolicyLegacy(t *testing.T, ds *Datastore) { }) require.NoError(t, err) - policies, err := ds.ListGlobalPolicies(context.Background(), fleet.ListOptions{}) + policies, err := ds.ListGlobalPolicies(context.Background(), fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, policies, 2) assert.Equal(t, q.Name, policies[0].Name) @@ -163,7 +164,7 @@ func testPoliciesNewGlobalPolicyLegacy(t *testing.T, ds *Datastore) { _, err = ds.DeleteGlobalPolicies(context.Background(), []uint{policies[0].ID, policies[1].ID}) require.NoError(t, err) - policies, err = ds.ListGlobalPolicies(context.Background(), fleet.ListOptions{}) + policies, err = ds.ListGlobalPolicies(context.Background(), fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, policies, 0) } @@ -195,7 +196,7 @@ func testPoliciesNewGlobalPolicyProprietary(t *testing.T, ds *Datastore) { }) require.NoError(t, err) - policies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + policies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, policies, 2) assert.Equal(t, "query1", policies[0].Name) @@ -230,7 +231,7 @@ func testPoliciesNewGlobalPolicyProprietary(t *testing.T, ds *Datastore) { _, err = ds.DeleteGlobalPolicies(ctx, []uint{policies[0].ID, policies[1].ID}) require.NoError(t, err) - policies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + policies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, policies, 0) @@ -275,7 +276,7 @@ func testGlobalPolicyPendingScriptsAndInstalls(t *testing.T, ds *Datastore) { _, err := q.ExecContext(ctx, "UPDATE policies SET script_id = ?", script.ID) return err }) - policies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + policies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, policies, 1) @@ -331,7 +332,7 @@ func testGlobalPolicyPendingScriptsAndInstalls(t *testing.T, ds *Datastore) { _, err := q.ExecContext(ctx, "UPDATE policies SET software_installer_id = ?", installerID) return err }) - policies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + policies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, policies, 1) @@ -396,7 +397,7 @@ func testPoliciesListOptions(t *testing.T, ds *Datastore) { }) require.NoError(t, err) - policies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{MatchQuery: "apple", OrderKey: "name", OrderDirection: fleet.OrderAscending}) + policies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{MatchQuery: "apple", OrderKey: "name", OrderDirection: fleet.OrderAscending}, "") require.NoError(t, err) require.Len(t, policies, 3) assert.Equal(t, "apple", policies[0].Name) @@ -404,6 +405,113 @@ func testPoliciesListOptions(t *testing.T, ds *Datastore) { assert.Equal(t, "rotten apple", policies[2].Name) } +func testPoliciesPlatformFilter(t *testing.T, ds *Datastore) { + user1 := test.NewUser(t, ds, "Alice", "alice@example.com", true) + ctx := t.Context() + + // Cross-platform policy (empty platform string targets all platforms) + _, err := ds.NewGlobalPolicy(ctx, &user1.ID, fleet.PolicyPayload{ + Name: "cross-platform", + Query: "select 1;", + }) + require.NoError(t, err) + + _, err = ds.NewGlobalPolicy(ctx, &user1.ID, fleet.PolicyPayload{ + Name: "macos-only", + Query: "select 1;", + Platform: "darwin", + }) + require.NoError(t, err) + + _, err = ds.NewGlobalPolicy(ctx, &user1.ID, fleet.PolicyPayload{ + Name: "win-linux", + Query: "select 1;", + Platform: "windows,linux", + }) + require.NoError(t, err) + + // Empty platform = no filter, returns all policies + policies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{OrderKey: "name"}, "") + require.NoError(t, err) + require.Len(t, policies, 3) + + // darwin matches cross-platform (empty) + macos-only + policies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{OrderKey: "name"}, "darwin") + require.NoError(t, err) + require.Len(t, policies, 2) + names := []string{policies[0].Name, policies[1].Name} + assert.ElementsMatch(t, []string{"cross-platform", "macos-only"}, names) + + // windows matches cross-platform + win-linux + policies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{OrderKey: "name"}, "windows") + require.NoError(t, err) + require.Len(t, policies, 2) + names = []string{policies[0].Name, policies[1].Name} + assert.ElementsMatch(t, []string{"cross-platform", "win-linux"}, names) + + // linux matches cross-platform + win-linux + policies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{OrderKey: "name"}, "linux") + require.NoError(t, err) + require.Len(t, policies, 2) + names = []string{policies[0].Name, policies[1].Name} + assert.ElementsMatch(t, []string{"cross-platform", "win-linux"}, names) + + // chrome matches only cross-platform (no chrome-targeted policies exist) + policies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{OrderKey: "name"}, "chrome") + require.NoError(t, err) + require.Len(t, policies, 1) + assert.Equal(t, "cross-platform", policies[0].Name) + + // CountPolicies is filtered consistently. + count, err := ds.CountPolicies(ctx, nil, "", "", "darwin") + require.NoError(t, err) + assert.Equal(t, 2, count) + + count, err = ds.CountPolicies(ctx, nil, "", "", "chrome") + require.NoError(t, err) + assert.Equal(t, 1, count) + + count, err = ds.CountPolicies(ctx, nil, "", "", "") + require.NoError(t, err) + assert.Equal(t, 3, count) + + // Test team-scoped list / count with platform. + team, err := ds.NewTeam(ctx, &fleet.Team{Name: "team1"}) + require.NoError(t, err) + _, err = ds.NewTeamPolicy(ctx, team.ID, &user1.ID, fleet.PolicyPayload{ + Name: "team-darwin", + Query: "select 1;", + Platform: "darwin", + }) + require.NoError(t, err) + _, err = ds.NewTeamPolicy(ctx, team.ID, &user1.ID, fleet.PolicyPayload{ + Name: "team-all", + Query: "select 1;", + }) + require.NoError(t, err) + + // ListMergedTeamPolicies with darwin: team-darwin, team-all, and the 2 + // matching global policies (cross-platform, macos-only) + merged, err := ds.ListMergedTeamPolicies(ctx, team.ID, fleet.ListOptions{OrderKey: "name"}, "", "darwin") + require.NoError(t, err) + require.Len(t, merged, 4) + + // ListTeamPolicies with windows: the team has no darwin-only policy, so + // we expect team-all (empty platform) + cross-platform + win-linux inherited. + teamPols, inherited, err := ds.ListTeamPolicies(ctx, team.ID, fleet.ListOptions{OrderKey: "name"}, fleet.ListOptions{OrderKey: "name"}, "", "windows") + require.NoError(t, err) + require.Len(t, teamPols, 1) + assert.Equal(t, "team-all", teamPols[0].Name) + require.Len(t, inherited, 2) + inheritedNames := []string{inherited[0].Name, inherited[1].Name} + assert.ElementsMatch(t, []string{"cross-platform", "win-linux"}, inheritedNames) + + // CountMergedTeamPolicies with platform filter + mergedCount, err := ds.CountMergedTeamPolicies(ctx, team.ID, "", "", "darwin") + require.NoError(t, err) + assert.Equal(t, 4, mergedCount) +} + func testPoliciesMembershipView(deferred bool, t *testing.T, ds *Datastore) { ctx := context.Background() @@ -481,7 +589,7 @@ func testPoliciesMembershipView(deferred bool, t *testing.T, ds *Datastore) { require.NoError(t, ds.UpdateHostPolicyCounts(ctx)) - policies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + policies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, policies, 2) @@ -498,7 +606,7 @@ func testPoliciesMembershipView(deferred bool, t *testing.T, ds *Datastore) { require.NoError(t, ds.UpdateHostPolicyCounts(ctx)) - policies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + policies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, policies, 2) @@ -522,7 +630,7 @@ func testPoliciesMembershipView(deferred bool, t *testing.T, ds *Datastore) { ))) require.NoError(t, ds.UpdateHostPolicyCounts(ctx)) - policies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + policies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, policies, 2) // host1: p now passing again (was failing), host2: p still passing @@ -593,7 +701,7 @@ func testPoliciesMembershipView(deferred bool, t *testing.T, ds *Datastore) { require.NoError(t, ds.UpdateHostPolicyCounts(ctx)) - t1Pols, t1Inherited, err := ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + t1Pols, t1Inherited, err := ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, t1Pols, 1) assert.Equal(t, uint(1), t1Pols[0].PassingHostCount) @@ -607,7 +715,7 @@ func testPoliciesMembershipView(deferred bool, t *testing.T, ds *Datastore) { assert.Equal(t, uint(0), t1Inherited[1].PassingHostCount) assert.Equal(t, uint(1), t1Inherited[1].FailingHostCount) - t2Pols, t2Inherited, err := ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + t2Pols, t2Inherited, err := ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, t2Pols, 2) require.Equal(t, t2pol.ID, t2Pols[0].ID) @@ -654,7 +762,7 @@ func testTeamPolicyLegacy(t *testing.T, ds *Datastore) { }) require.NoError(t, err) - prevPolicies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + prevPolicies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, prevPolicies, 0) @@ -684,7 +792,7 @@ func testTeamPolicyLegacy(t *testing.T, ds *Datastore) { }) require.NoError(t, err) - globalPolicies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + globalPolicies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, globalPolicies, 1) @@ -699,7 +807,7 @@ func testTeamPolicyLegacy(t *testing.T, ds *Datastore) { require.NotNil(t, p2.AuthorID) assert.Equal(t, user1.ID, *p2.AuthorID) - teamPolicies, inherited1, err := ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + teamPolicies, inherited1, err := ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, teamPolicies, 1) assert.Equal(t, q.Name, teamPolicies[0].Name) @@ -711,7 +819,7 @@ func testTeamPolicyLegacy(t *testing.T, ds *Datastore) { require.Len(t, inherited1, 1) require.Equal(t, gpol, inherited1[0]) - team2Policies, inherited2, err := ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + team2Policies, inherited2, err := ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, team2Policies, 1) assert.Equal(t, q2.Name, team2Policies[0].Name) @@ -726,7 +834,7 @@ func testTeamPolicyLegacy(t *testing.T, ds *Datastore) { _, err = ds.DeleteTeamPolicies(ctx, team1.ID, []uint{teamPolicies[0].ID}) require.NoError(t, err) - teamPolicies, inherited1, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + teamPolicies, inherited1, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, teamPolicies, 0) require.Len(t, inherited1, 1) @@ -790,7 +898,7 @@ func testTeamPolicyProprietary(t *testing.T, ds *Datastore) { require.Error(t, err) require.Nil(t, gpol1) - prevPolicies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + prevPolicies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, prevPolicies, 1) requireLabels(t, []string{label1.Name, label2.Name}, prevPolicies[0].LabelsIncludeAny) @@ -844,7 +952,7 @@ func testTeamPolicyProprietary(t *testing.T, ds *Datastore) { assert.True(t, p.CalendarEventsEnabled) requireLabels(t, []string{label1.Name, label2.Name}, p.LabelsExcludeAny) - globalPolicies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + globalPolicies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, globalPolicies, len(prevPolicies)) @@ -864,7 +972,7 @@ func testTeamPolicyProprietary(t *testing.T, ds *Datastore) { require.NotNil(t, p2.AuthorID) assert.Equal(t, user1.ID, *p2.AuthorID) - teamPolicies, inherited1, err := ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + teamPolicies, inherited1, err := ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, teamPolicies, 1) assert.Equal(t, "query1", teamPolicies[0].Name) @@ -879,7 +987,7 @@ func testTeamPolicyProprietary(t *testing.T, ds *Datastore) { require.Len(t, inherited1, 1) require.Equal(t, gpol, inherited1[0]) - team2Policies, inherited2, err := ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + team2Policies, inherited2, err := ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, team2Policies, 1) assert.Equal(t, "query2", team2Policies[0].Name) @@ -991,7 +1099,7 @@ func testTeamPolicyProprietary(t *testing.T, ds *Datastore) { _, err = ds.DeleteTeamPolicies(ctx, team1.ID, []uint{teamPolicies[0].ID}) require.NoError(t, err) - teamPolicies, inherited1, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + teamPolicies, inherited1, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, teamPolicies, 0) require.Len(t, inherited1, 1) @@ -1006,7 +1114,7 @@ func testTeamPolicyProprietary(t *testing.T, ds *Datastore) { }) require.NoError(t, err) - teamPolicies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + teamPolicies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, teamPolicies, 1) assert.Equal(t, "query1", teamPolicies[0].Name) @@ -1148,7 +1256,7 @@ func testListMergedTeamPolicies(t *testing.T, ds *Datastore) { merged, err := ds.ListMergedTeamPolicies(ctx, team1.ID, fleet.ListOptions{ OrderKey: "name", OrderDirection: fleet.OrderDescending, - }, "") + }, "", "") require.NoError(t, err) require.Len(t, merged, 2) @@ -1158,14 +1266,14 @@ func testListMergedTeamPolicies(t *testing.T, ds *Datastore) { // Test filter merged, err = ds.ListMergedTeamPolicies(ctx, team1.ID, fleet.ListOptions{ MatchQuery: "query1", - }, "") + }, "", "") require.NoError(t, err) require.Len(t, merged, 1) assert.Equal(t, gpol.ID, merged[0].ID) merged, err = ds.ListMergedTeamPolicies(ctx, team1.ID, fleet.ListOptions{ MatchQuery: "query2", - }, "") + }, "", "") require.NoError(t, err) require.Len(t, merged, 1) assert.Equal(t, team1policy.ID, merged[0].ID) @@ -1186,7 +1294,7 @@ func testListMergedTeamPolicies(t *testing.T, ds *Datastore) { // team 1 shows no host counts merged, err = ds.ListMergedTeamPolicies(ctx, team1.ID, fleet.ListOptions{ OrderKey: "name", - }, "") + }, "", "") require.NoError(t, err) require.Len(t, merged, 2) assert.Equal(t, gpol.ID, merged[0].ID) @@ -1209,7 +1317,7 @@ func testListMergedTeamPolicies(t *testing.T, ds *Datastore) { // team 1 shows host counts merged, err = ds.ListMergedTeamPolicies(ctx, team1.ID, fleet.ListOptions{ OrderKey: "name", - }, "") + }, "", "") require.NoError(t, err) require.Len(t, merged, 2) assert.Equal(t, gpol.ID, merged[0].ID) @@ -1222,7 +1330,7 @@ func testListMergedTeamPolicies(t *testing.T, ds *Datastore) { // team2 shows no host counts merged, err = ds.ListMergedTeamPolicies(ctx, team2.ID, fleet.ListOptions{ OrderKey: "name", - }, "") + }, "", "") require.NoError(t, err) require.Len(t, merged, 2) assert.Equal(t, gpol.ID, merged[0].ID) @@ -1744,20 +1852,20 @@ func testTeamPolicyTransfer(t *testing.T, ds *Datastore) { checkPassingCount := func(tm1, tm1Inherited, tm2Inherited, global uint) { t.Helper() require.NoError(t, ds.UpdateHostPolicyCounts(ctx)) - policies, inherited, err := ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + policies, inherited, err := ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, policies, 1) assert.Equal(t, tm1, policies[0].PassingHostCount) require.Len(t, inherited, 1) assert.Equal(t, tm1Inherited, inherited[0].PassingHostCount) - policies, inherited, err = ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + policies, inherited, err = ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, policies, 0) // team 2 has no policies of its own require.Len(t, inherited, 1) assert.Equal(t, tm2Inherited, inherited[0].PassingHostCount) - policies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + policies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, policies, 1) assert.Equal(t, global, policies[0].PassingHostCount) @@ -1880,7 +1988,7 @@ func testApplyPolicySpec(t *testing.T, ds *Datastore) { }, })) - policies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + policies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, policies, 1) assert.Equal(t, "query1"+unicode, policies[0].Name) @@ -1897,7 +2005,7 @@ func testApplyPolicySpec(t *testing.T, ds *Datastore) { }}, policies[0].LabelsIncludeAny) assert.Equal(t, policies[0].Type, fleet.PolicyTypeDynamic) - teamPolicies, _, err := ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + teamPolicies, _, err := ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, teamPolicies, 2) assert.Equal(t, "query2", teamPolicies[0].Name) @@ -1924,7 +2032,7 @@ func testApplyPolicySpec(t *testing.T, ds *Datastore) { assert.Equal(t, "windows,linux", teamPolicies[1].Platform) assert.False(t, teamPolicies[1].CalendarEventsEnabled) - noTeamPolicies, _, err := ds.ListTeamPolicies(ctx, fleet.PolicyNoTeamID, fleet.ListOptions{}, fleet.ListOptions{}, "") + noTeamPolicies, _, err := ds.ListTeamPolicies(ctx, fleet.PolicyNoTeamID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, noTeamPolicies, 1) assert.Equal(t, "query4", noTeamPolicies[0].Name) @@ -1982,13 +2090,13 @@ func testApplyPolicySpec(t *testing.T, ds *Datastore) { }, })) - policies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + policies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, policies, 1) - teamPolicies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + teamPolicies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, teamPolicies, 2) - noTeamPolicies, _, err = ds.ListTeamPolicies(ctx, fleet.PolicyNoTeamID, fleet.ListOptions{}, fleet.ListOptions{}, "") + noTeamPolicies, _, err = ds.ListTeamPolicies(ctx, fleet.PolicyNoTeamID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, noTeamPolicies, 1) @@ -2015,7 +2123,7 @@ func testApplyPolicySpec(t *testing.T, ds *Datastore) { Type: fleet.PolicyTypeDynamic, }, })) - policies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + policies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, policies, 1) @@ -2037,7 +2145,7 @@ func testApplyPolicySpec(t *testing.T, ds *Datastore) { LabelID: barLabel.ID, }) - teamPolicies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + teamPolicies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, teamPolicies, 2) @@ -2135,7 +2243,7 @@ func testApplyPolicySpecDefaultType(t *testing.T, ds *Datastore) { }, })) - policies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + policies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, policies, 1) assert.Equal(t, "no-type-policy", policies[0].Name) @@ -2248,11 +2356,11 @@ func testApplyPolicySpecWithQueryPlatformChanges(t *testing.T, ds *Datastore) { } // load the global policies - gPolicies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + gPolicies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, gPolicies, 3) // load the team policies - tPolicies, _, err := ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + tPolicies, _, err := ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, tPolicies, 3) @@ -2303,10 +2411,10 @@ func testApplyPolicySpecWithQueryPlatformChanges(t *testing.T, ds *Datastore) { assert.Equal(t, uint64(3), globalHosts[hostLin].FailingPoliciesCount) // Ensure policy passing and failing counts are correct - gPolicies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + gPolicies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, gPolicies, 3) - tPolicies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + tPolicies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, tPolicies, 3) @@ -2404,10 +2512,10 @@ func testApplyPolicySpecWithQueryPlatformChanges(t *testing.T, ds *Datastore) { assert.Equal(t, uint64(1), globalHosts[hostLin].FailingPoliciesCount) // Ensure policy passing and failing counts are correct - gPolicies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + gPolicies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, gPolicies, 4) - tPolicies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + tPolicies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, tPolicies, 4) @@ -2426,10 +2534,10 @@ func testApplyPolicySpecWithQueryPlatformChanges(t *testing.T, ds *Datastore) { err = ds.UpdateHostPolicyCounts(ctx) require.NoError(t, err) - gPolicies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + gPolicies, err = ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, gPolicies, 4) - tPolicies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + tPolicies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, tPolicies, 4) @@ -2682,7 +2790,7 @@ func testCachedPolicyCountDeletesOnPolicyChange(t *testing.T, ds *Datastore) { globalPolicy, err = ds.Policy(ctx, globalPolicy.ID) require.NoError(t, err) assert.Equal(t, uint(2), globalPolicy.FailingHostCount) - teamPolicies, inheritedPolicies, err := ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + teamPolicies, inheritedPolicies, err := ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, teamPolicies, 1) require.Len(t, inheritedPolicies, 1) @@ -2702,7 +2810,7 @@ func testCachedPolicyCountDeletesOnPolicyChange(t *testing.T, ds *Datastore) { globalPolicy, err = ds.Policy(ctx, globalPolicy.ID) require.NoError(t, err) assert.Equal(t, uint(0), globalPolicy.FailingHostCount) - teamPolicies, inheritedPolicies, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + teamPolicies, inheritedPolicies, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, teamPolicies, 1) require.Len(t, inheritedPolicies, 1) @@ -2720,7 +2828,7 @@ func testCachedPolicyCountDeletesOnPolicyChange(t *testing.T, ds *Datastore) { err = ds.SavePolicy(ctx, teamPolicy, false, true) require.NoError(t, err) - teamPolicies, inheritedPolicies, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + teamPolicies, inheritedPolicies, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, teamPolicies, 1) require.Len(t, inheritedPolicies, 1) @@ -3115,11 +3223,11 @@ func testPolicyPlatformUpdate(t *testing.T, ds *Datastore) { require.NoError(t, err) // load the global policies - gpols, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + gpols, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, gpols, 2) // load the team policies - tpols, _, err := ds.ListTeamPolicies(ctx, tm.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + tpols, _, err := ds.ListTeamPolicies(ctx, tm.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, tpols, 2) @@ -3684,7 +3792,7 @@ func testListGlobalPoliciesCanPaginate(t *testing.T, ds *Datastore) { policies, err := ds.ListGlobalPolicies(context.Background(), fleet.ListOptions{ Page: 0, PerPage: 20, - }) + }, "") assert.Equal(t, "global policy 0", policies[0].Name) assert.Len(t, policies, 20) @@ -3694,14 +3802,14 @@ func testListGlobalPoliciesCanPaginate(t *testing.T, ds *Datastore) { policies, err = ds.ListGlobalPolicies(context.Background(), fleet.ListOptions{ Page: 1, PerPage: 20, - }) + }, "") assert.Equal(t, "global policy 20", policies[0].Name) assert.Len(t, policies, 10) require.NoError(t, err) // No list options returns all policies - policies, err = ds.ListGlobalPolicies(context.Background(), fleet.ListOptions{}) + policies, err = ds.ListGlobalPolicies(context.Background(), fleet.ListOptions{}, "") assert.Len(t, policies, 30) require.NoError(t, err) } @@ -3726,7 +3834,7 @@ func testListTeamPoliciesCanPaginate(t *testing.T, ds *Datastore) { policies, _, err := ds.ListTeamPolicies(context.Background(), tm.ID, fleet.ListOptions{ Page: 0, PerPage: 20, - }, fleet.ListOptions{}, "") + }, fleet.ListOptions{}, "", "") assert.Equal(t, "team policy 0", policies[0].Name) assert.Len(t, policies, 20) @@ -3736,14 +3844,14 @@ func testListTeamPoliciesCanPaginate(t *testing.T, ds *Datastore) { policies, _, err = ds.ListTeamPolicies(context.Background(), tm.ID, fleet.ListOptions{ Page: 1, PerPage: 20, - }, fleet.ListOptions{}, "") + }, fleet.ListOptions{}, "", "") assert.Equal(t, "team policy 20", policies[0].Name) assert.Len(t, policies, 10) require.NoError(t, err) // No list options returns all policies - policies, _, err = ds.ListTeamPolicies(context.Background(), 1, fleet.ListOptions{}, fleet.ListOptions{}, "") + policies, _, err = ds.ListTeamPolicies(context.Background(), 1, fleet.ListOptions{}, fleet.ListOptions{}, "", "") assert.Len(t, policies, 30) require.NoError(t, err) } @@ -3755,15 +3863,15 @@ func testCountPolicies(t *testing.T, ds *Datastore) { require.NoError(t, err) // no policies - globalCount, err := ds.CountPolicies(ctx, nil, "", "") + globalCount, err := ds.CountPolicies(ctx, nil, "", "", "") require.NoError(t, err) assert.Equal(t, 0, globalCount) - teamCount, err := ds.CountPolicies(ctx, &tm.ID, "", "") + teamCount, err := ds.CountPolicies(ctx, &tm.ID, "", "", "") require.NoError(t, err) assert.Equal(t, 0, teamCount) - mergedCount, err := ds.CountMergedTeamPolicies(ctx, tm.ID, "", "") + mergedCount, err := ds.CountMergedTeamPolicies(ctx, tm.ID, "", "", "") require.NoError(t, err) assert.Equal(t, 0, mergedCount) @@ -3773,15 +3881,15 @@ func testCountPolicies(t *testing.T, ds *Datastore) { require.NoError(t, err) } - globalCount, err = ds.CountPolicies(ctx, nil, "", "") + globalCount, err = ds.CountPolicies(ctx, nil, "", "", "") require.NoError(t, err) assert.Equal(t, 10, globalCount) - teamCount, err = ds.CountPolicies(ctx, &tm.ID, "", "") + teamCount, err = ds.CountPolicies(ctx, &tm.ID, "", "", "") require.NoError(t, err) assert.Equal(t, 0, teamCount) - mergedCount, err = ds.CountMergedTeamPolicies(ctx, tm.ID, "", "") + mergedCount, err = ds.CountMergedTeamPolicies(ctx, tm.ID, "", "", "") require.NoError(t, err) assert.Equal(t, 10, mergedCount) @@ -3791,33 +3899,33 @@ func testCountPolicies(t *testing.T, ds *Datastore) { require.NoError(t, err) } - teamCount, err = ds.CountPolicies(ctx, &tm.ID, "", "") + teamCount, err = ds.CountPolicies(ctx, &tm.ID, "", "", "") require.NoError(t, err) assert.Equal(t, 5, teamCount) - globalCount, err = ds.CountPolicies(ctx, nil, "", "") + globalCount, err = ds.CountPolicies(ctx, nil, "", "", "") require.NoError(t, err) assert.Equal(t, 10, globalCount) - mergedCount, err = ds.CountMergedTeamPolicies(ctx, tm.ID, "", "") + mergedCount, err = ds.CountMergedTeamPolicies(ctx, tm.ID, "", "", "") require.NoError(t, err) assert.Equal(t, 15, mergedCount) // test filter - globalCount, err = ds.CountPolicies(ctx, nil, "global policy 1", "") + globalCount, err = ds.CountPolicies(ctx, nil, "global policy 1", "", "") require.NoError(t, err) assert.Equal(t, 1, globalCount) - teamCount, err = ds.CountPolicies(ctx, &tm.ID, "team policy 1", "") + teamCount, err = ds.CountPolicies(ctx, &tm.ID, "team policy 1", "", "") require.NoError(t, err) assert.Equal(t, 1, teamCount) - mergedCount, err = ds.CountMergedTeamPolicies(ctx, tm.ID, "policy 1", "") + mergedCount, err = ds.CountMergedTeamPolicies(ctx, tm.ID, "policy 1", "", "") require.NoError(t, err) assert.Equal(t, 2, mergedCount) // test automation filter doesn't affect global policy count - globalCount, err = ds.CountPolicies(ctx, nil, "", "scripts") + globalCount, err = ds.CountPolicies(ctx, nil, "", "scripts", "") require.NoError(t, err) assert.Equal(t, 10, globalCount) } @@ -4020,7 +4128,7 @@ func testPoliciesNameUnicode(t *testing.T, ds *Datastore) { assert.True(t, IsDuplicate(err), err) // Try to find policy with equivalent name - policies, err := ds.ListGlobalPolicies(context.Background(), fleet.ListOptions{MatchQuery: equivalentNames[1]}) + policies, err := ds.ListGlobalPolicies(context.Background(), fleet.ListOptions{MatchQuery: equivalentNames[1]}, "") assert.NoError(t, err) require.Len(t, policies, 1) assert.Equal(t, equivalentNames[0], policies[0].Name) @@ -4040,7 +4148,7 @@ func testPoliciesNameUnicode(t *testing.T, ds *Datastore) { // ListTeamPolicies, including inherited policy teamPolicies, inheritedPolicies, err := ds.ListTeamPolicies( - context.Background(), team.ID, fleet.ListOptions{MatchQuery: equivalentNames[1]}, fleet.ListOptions{MatchQuery: equivalentNames[1]}, "", + context.Background(), team.ID, fleet.ListOptions{MatchQuery: equivalentNames[1]}, fleet.ListOptions{MatchQuery: equivalentNames[1]}, "", "", ) assert.NoError(t, err) require.Len(t, teamPolicies, 1) @@ -4049,10 +4157,10 @@ func testPoliciesNameUnicode(t *testing.T, ds *Datastore) { assert.Equal(t, equivalentNames[0], inheritedPolicies[0].Name) // CountPolicies - count, err := ds.CountPolicies(context.Background(), &team.ID, equivalentNames[1], "") + count, err := ds.CountPolicies(context.Background(), &team.ID, equivalentNames[1], "", "") assert.NoError(t, err) assert.Equal(t, 1, count) - count, err = ds.CountPolicies(context.Background(), nil, equivalentNames[1], "") + count, err = ds.CountPolicies(context.Background(), nil, equivalentNames[1], "", "") assert.NoError(t, err) assert.Equal(t, 1, count) } @@ -4068,13 +4176,13 @@ func testPoliciesNameEmoji(t *testing.T, ds *Datastore) { assert.Equal(t, emoji1, policyEmoji.Name) // Try to find policy with emoji0 - policies, err := ds.ListGlobalPolicies(context.Background(), fleet.ListOptions{MatchQuery: emoji0}) + policies, err := ds.ListGlobalPolicies(context.Background(), fleet.ListOptions{MatchQuery: emoji0}, "") assert.NoError(t, err) require.Len(t, policies, 1) assert.Equal(t, emoji0, policies[0].Name) // Try to find policy with emoji1 - policies, err = ds.ListGlobalPolicies(context.Background(), fleet.ListOptions{MatchQuery: emoji1}) + policies, err = ds.ListGlobalPolicies(context.Background(), fleet.ListOptions{MatchQuery: emoji1}, "") assert.NoError(t, err) require.Len(t, policies, 1) assert.Equal(t, emoji1, policies[0].Name) @@ -4092,7 +4200,7 @@ func testPoliciesNameSort(t *testing.T, ds *Datastore) { policies[0], err = ds.NewGlobalPolicy(context.Background(), nil, fleet.PolicyPayload{Name: "а"}) require.NoError(t, err) - policiesResult, err := ds.ListGlobalPolicies(context.Background(), fleet.ListOptions{OrderKey: "name"}) + policiesResult, err := ds.ListGlobalPolicies(context.Background(), fleet.ListOptions{OrderKey: "name"}, "") assert.NoError(t, err) require.Len(t, policies, 3) for i, policy := range policies { @@ -5349,7 +5457,7 @@ func testApplyPolicySpecWithInstallers(t *testing.T, ds *Datastore) { }, }) require.NoError(t, err) - team1Policies, _, err := ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + team1Policies, _, err := ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, team1Policies, 2) @@ -5362,7 +5470,7 @@ func testApplyPolicySpecWithInstallers(t *testing.T, ds *Datastore) { require.NoError(t, err) require.Equal(t, va1Meta.VPPAppsTeamsID, *vppPolicy1Team1.VPPAppsTeamsID) - team2Policies, _, err := ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + team2Policies, _, err := ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, team2Policies, 2) require.NotNil(t, team2Policies[0].SoftwareInstallerID) @@ -5372,7 +5480,7 @@ func testApplyPolicySpecWithInstallers(t *testing.T, ds *Datastore) { require.NoError(t, err) require.Equal(t, va2Meta.VPPAppsTeamsID, *vppPolicy2Team2.VPPAppsTeamsID) - noTeamPolicies, _, err := ds.ListTeamPolicies(ctx, fleet.PolicyNoTeamID, fleet.ListOptions{}, fleet.ListOptions{}, "") + noTeamPolicies, _, err := ds.ListTeamPolicies(ctx, fleet.PolicyNoTeamID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, noTeamPolicies, 2) require.NotNil(t, noTeamPolicies[0].SoftwareInstallerID) @@ -5415,7 +5523,7 @@ func testApplyPolicySpecWithInstallers(t *testing.T, ds *Datastore) { }, }) require.NoError(t, err) - team1Policies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + team1Policies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, team1Policies, 2) require.Nil(t, team1Policies[0].SoftwareInstallerID) @@ -5535,7 +5643,7 @@ func testApplyPolicySpecWithInstallers(t *testing.T, ds *Datastore) { }, }) require.NoError(t, err) - team2Policies, _, err = ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + team2Policies, _, err = ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, team2Policies, 2) require.Nil(t, team2Policies[0].SoftwareInstallerID) @@ -5611,7 +5719,7 @@ func testApplyPolicySpecWithInstallers(t *testing.T, ds *Datastore) { }, }) require.NoError(t, err) - team1Policies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + team1Policies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, team1Policies, 2) require.NotNil(t, team1Policies[0].SoftwareInstallerID) @@ -5631,7 +5739,7 @@ func testApplyPolicySpecWithInstallers(t *testing.T, ds *Datastore) { ) }) require.False(t, countBiggerThanZero) - team2Policies, _, err = ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + team2Policies, _, err = ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, team2Policies, 2) require.NotNil(t, team2Policies[0].SoftwareInstallerID) @@ -5672,7 +5780,7 @@ func testApplyPolicySpecWithInstallers(t *testing.T, ds *Datastore) { }, }) require.NoError(t, err) - team1Policies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + team1Policies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, team1Policies, 2) require.Equal(t, uint(1), team1Policies[0].FailingHostCount) @@ -5726,7 +5834,7 @@ func testApplyPolicySpecWithInstallers(t *testing.T, ds *Datastore) { }, }) require.NoError(t, err) - team1Policies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + team1Policies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, team1Policies, 2) require.Equal(t, uint(0), team1Policies[0].FailingHostCount) @@ -5872,7 +5980,7 @@ func testTeamPoliciesNoTeam(t *testing.T, ds *Datastore) { require.NoError(t, err) // Tests on global domain. - globalPolicies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + globalPolicies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) require.Len(t, globalPolicies, 2) require.Equal(t, globalPolicy1.ID, globalPolicies[0].ID) @@ -5888,7 +5996,7 @@ func testTeamPoliciesNoTeam(t *testing.T, ds *Datastore) { require.Equal(t, p, globalPolicy) ids = append(ids, globalPolicy.ID) } - c, err := ds.CountPolicies(ctx, nil, "", "") + c, err := ds.CountPolicies(ctx, nil, "", "", "") require.NoError(t, err) require.Equal(t, 2, c) globalPoliciesByID, err := ds.PoliciesByID(ctx, ids) @@ -5898,7 +6006,7 @@ func testTeamPoliciesNoTeam(t *testing.T, ds *Datastore) { require.Equal(t, globalPoliciesByID[globalPolicies[1].ID], globalPolicies[1]) // Tests on team1 domain. - teamPolicies, inheritedPolicies, err := ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + teamPolicies, inheritedPolicies, err := ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, teamPolicies, 1) require.Equal(t, policy1Team1.ID, teamPolicies[0].ID) @@ -5922,13 +6030,13 @@ func testTeamPoliciesNoTeam(t *testing.T, ds *Datastore) { require.NoError(t, err) require.Len(t, teamPoliciesByID, 1) require.Equal(t, teamPoliciesByID[teamPolicies[0].ID], teamPolicies[0]) - c, err = ds.CountMergedTeamPolicies(ctx, team1.ID, "", "") + c, err = ds.CountMergedTeamPolicies(ctx, team1.ID, "", "", "") require.NoError(t, err) require.Equal(t, 3, c) - c, err = ds.CountPolicies(ctx, &team1.ID, "", "") + c, err = ds.CountPolicies(ctx, &team1.ID, "", "", "") require.NoError(t, err) require.Equal(t, 1, c) - mergedTeamPolicies, err := ds.ListMergedTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, "") + mergedTeamPolicies, err := ds.ListMergedTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, mergedTeamPolicies, 3) require.Equal(t, policy1Team1.ID, mergedTeamPolicies[0].ID) @@ -5942,7 +6050,7 @@ func testTeamPoliciesNoTeam(t *testing.T, ds *Datastore) { require.Equal(t, uint(1), mergedTeamPolicies[2].PassingHostCount) // Tests on team2 domain. - teamPolicies, inheritedPolicies, err = ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + teamPolicies, inheritedPolicies, err = ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, teamPolicies, 2) require.Equal(t, policy2Team2.ID, teamPolicies[0].ID) @@ -5970,13 +6078,13 @@ func testTeamPoliciesNoTeam(t *testing.T, ds *Datastore) { require.Len(t, teamPoliciesByID, 2) require.Equal(t, teamPoliciesByID[teamPolicies[0].ID], teamPolicies[0]) require.Equal(t, teamPoliciesByID[teamPolicies[1].ID], teamPolicies[1]) - c, err = ds.CountMergedTeamPolicies(ctx, team2.ID, "", "") + c, err = ds.CountMergedTeamPolicies(ctx, team2.ID, "", "", "") require.NoError(t, err) require.Equal(t, 4, c) - c, err = ds.CountPolicies(ctx, &team2.ID, "", "") + c, err = ds.CountPolicies(ctx, &team2.ID, "", "", "") require.NoError(t, err) require.Equal(t, 2, c) - mergedTeamPolicies, err = ds.ListMergedTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, "") + mergedTeamPolicies, err = ds.ListMergedTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, mergedTeamPolicies, 4) require.Equal(t, policy2Team2.ID, mergedTeamPolicies[0].ID) @@ -5993,7 +6101,7 @@ func testTeamPoliciesNoTeam(t *testing.T, ds *Datastore) { require.Equal(t, uint(0), mergedTeamPolicies[3].PassingHostCount) // Tests on "No team" domain. - teamPolicies, inheritedPolicies, err = ds.ListTeamPolicies(ctx, fleet.PolicyNoTeamID, fleet.ListOptions{}, fleet.ListOptions{}, "") + teamPolicies, inheritedPolicies, err = ds.ListTeamPolicies(ctx, fleet.PolicyNoTeamID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, teamPolicies, 2) require.Equal(t, policy0NoTeam.ID, teamPolicies[0].ID) @@ -6021,13 +6129,13 @@ func testTeamPoliciesNoTeam(t *testing.T, ds *Datastore) { require.Len(t, teamPoliciesByID, 2) require.Equal(t, teamPoliciesByID[teamPolicies[0].ID], teamPolicies[0]) require.Equal(t, teamPoliciesByID[teamPolicies[1].ID], teamPolicies[1]) - c, err = ds.CountMergedTeamPolicies(ctx, fleet.PolicyNoTeamID, "", "") + c, err = ds.CountMergedTeamPolicies(ctx, fleet.PolicyNoTeamID, "", "", "") require.NoError(t, err) require.Equal(t, 4, c) - c, err = ds.CountPolicies(ctx, ptr.Uint(fleet.PolicyNoTeamID), "", "") + c, err = ds.CountPolicies(ctx, new(fleet.PolicyNoTeamID), "", "", "") require.NoError(t, err) require.Equal(t, 2, c) - mergedTeamPolicies, err = ds.ListMergedTeamPolicies(ctx, fleet.PolicyNoTeamID, fleet.ListOptions{}, "") + mergedTeamPolicies, err = ds.ListMergedTeamPolicies(ctx, fleet.PolicyNoTeamID, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, mergedTeamPolicies, 4) require.Equal(t, policy0NoTeam.ID, mergedTeamPolicies[0].ID) @@ -6804,7 +6912,7 @@ func testPolicyLabelMembershipCleanup(t *testing.T, ds *Datastore) { }, }) require.NoError(t, err) - allPolicies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + allPolicies, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) var freshPolicy *fleet.Policy for _, p := range allPolicies { @@ -7663,7 +7771,7 @@ func testApplyPolicySpecsNeedsFullMembershipCleanupFlag(t *testing.T, ds *Datast })) // Find the policy by name so the test is not sensitive to other global policies created by concurrent tests. - pols, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + pols, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) var pol *fleet.Policy for _, p := range pols { @@ -7763,7 +7871,7 @@ func testCleanupPolicyMembershipCrashRecovery(t *testing.T, ds *Datastore) { require.NoError(t, ds.ApplyPolicySpecs(ctx, user1.ID, []*fleet.PolicySpec{ {Name: "retry recovery policy", Query: "select 1;", Type: fleet.PolicyTypeDynamic}, })) - pols, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + pols, err := ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) var pol *fleet.Policy for _, p := range pols { @@ -8075,7 +8183,7 @@ func testTeamPatchPolicy(t *testing.T, ds *Datastore) { require.NoError(t, err) // Verify the policy was created with the expected auto-generated query and platform. - policies, _, err := ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + policies, _, err := ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, policies, 1) require.Equal(t, "patch-valid-slug", policies[0].Name) @@ -8101,7 +8209,7 @@ func testTeamPatchPolicy(t *testing.T, ds *Datastore) { }) require.NoError(t, err) - policies, _, err = ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + policies, _, err = ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, policies, 1) require.Equal(t, previousID, policies[0].ID) @@ -8264,7 +8372,7 @@ func testApplyPolicySpecsRenamePatchPolicyRegression43687(t *testing.T, ds *Data }) require.NoError(t, err) - policies, _, err := ds.ListTeamPolicies(ctx, team.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + policies, _, err := ds.ListTeamPolicies(ctx, team.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, policies, 1) require.Equal(t, "Adobe Reader up to date", policies[0].Name) @@ -8286,7 +8394,7 @@ func testApplyPolicySpecsRenamePatchPolicyRegression43687(t *testing.T, ds *Data require.NoError(t, err) // The same row must now carry the new name (update, not delete + recreate). - policies, _, err = ds.ListTeamPolicies(ctx, team.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + policies, _, err = ds.ListTeamPolicies(ctx, team.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, policies, 1) require.Equal(t, originalID, policies[0].ID) @@ -8422,7 +8530,7 @@ func testTeamPolicyAutomationFilter(t *testing.T, ds *Datastore) { merged, err := ds.ListMergedTeamPolicies(ctx, 0, fleet.ListOptions{ OrderKey: "name", OrderDirection: fleet.OrderAscending, - }, "") + }, "", "") require.NoError(t, err) require.Len(t, merged, 8) @@ -8435,7 +8543,7 @@ func testTeamPolicyAutomationFilter(t *testing.T, ds *Datastore) { assert.Equal(t, teamWebhookPolicy.ID, merged[6].ID) assert.Equal(t, teamPatchPolicy.ID, merged[7].ID) - mergedCount, err := ds.CountMergedTeamPolicies(ctx, 0, "", "") + mergedCount, err := ds.CountMergedTeamPolicies(ctx, 0, "", "", "") require.NoError(t, err) assert.Equal(t, 8, mergedCount) @@ -8443,14 +8551,14 @@ func testTeamPolicyAutomationFilter(t *testing.T, ds *Datastore) { merged, err = ds.ListMergedTeamPolicies(ctx, 0, fleet.ListOptions{ OrderKey: "name", OrderDirection: fleet.OrderAscending, - }, "software") + }, "software", "") require.NoError(t, err) require.Len(t, merged, 3) assert.Equal(t, teamInstallerPolicy.ID, merged[0].ID) assert.Equal(t, teamAppStorePolicy.ID, merged[1].ID) assert.Equal(t, teamPatchPolicy.ID, merged[2].ID) - mergedCount, err = ds.CountMergedTeamPolicies(ctx, 0, "", "software") + mergedCount, err = ds.CountMergedTeamPolicies(ctx, 0, "", "software", "") require.NoError(t, err) assert.Equal(t, 3, mergedCount) @@ -8458,12 +8566,12 @@ func testTeamPolicyAutomationFilter(t *testing.T, ds *Datastore) { merged, err = ds.ListMergedTeamPolicies(ctx, 0, fleet.ListOptions{ OrderKey: "name", OrderDirection: fleet.OrderAscending, - }, "scripts") + }, "scripts", "") require.NoError(t, err) require.Len(t, merged, 1) assert.Equal(t, teamScriptPolicy.ID, merged[0].ID) - mergedCount, err = ds.CountMergedTeamPolicies(ctx, 0, "", "scripts") + mergedCount, err = ds.CountMergedTeamPolicies(ctx, 0, "", "scripts", "") require.NoError(t, err) assert.Equal(t, 1, mergedCount) @@ -8471,12 +8579,12 @@ func testTeamPolicyAutomationFilter(t *testing.T, ds *Datastore) { merged, err = ds.ListMergedTeamPolicies(ctx, 0, fleet.ListOptions{ OrderKey: "name", OrderDirection: fleet.OrderAscending, - }, "calendar") + }, "calendar", "") require.NoError(t, err) require.Len(t, merged, 1) assert.Equal(t, teamCalendarPolicy.ID, merged[0].ID) - mergedCount, err = ds.CountMergedTeamPolicies(ctx, 0, "", "calendar") + mergedCount, err = ds.CountMergedTeamPolicies(ctx, 0, "", "calendar", "") require.NoError(t, err) assert.Equal(t, 1, mergedCount) @@ -8484,12 +8592,12 @@ func testTeamPolicyAutomationFilter(t *testing.T, ds *Datastore) { merged, err = ds.ListMergedTeamPolicies(ctx, 0, fleet.ListOptions{ OrderKey: "name", OrderDirection: fleet.OrderAscending, - }, "conditional_access") + }, "conditional_access", "") require.NoError(t, err) require.Len(t, merged, 1) assert.Equal(t, teamConditionalPolicy.ID, merged[0].ID) - mergedCount, err = ds.CountMergedTeamPolicies(ctx, 0, "", "conditional_access") + mergedCount, err = ds.CountMergedTeamPolicies(ctx, 0, "", "conditional_access", "") require.NoError(t, err) assert.Equal(t, 1, mergedCount) @@ -8497,12 +8605,12 @@ func testTeamPolicyAutomationFilter(t *testing.T, ds *Datastore) { merged, err = ds.ListMergedTeamPolicies(ctx, 0, fleet.ListOptions{ OrderKey: "name", OrderDirection: fleet.OrderAscending, - }, "other") + }, "other", "") require.NoError(t, err) require.Len(t, merged, 1) assert.Equal(t, teamWebhookPolicy.ID, merged[0].ID) - mergedCount, err = ds.CountMergedTeamPolicies(ctx, 0, "", "other") + mergedCount, err = ds.CountMergedTeamPolicies(ctx, 0, "", "other", "") require.NoError(t, err) assert.Equal(t, 1, mergedCount) @@ -8510,14 +8618,14 @@ func testTeamPolicyAutomationFilter(t *testing.T, ds *Datastore) { policies, _, err := ds.ListTeamPolicies(ctx, 0, fleet.ListOptions{ OrderKey: "name", OrderDirection: fleet.OrderAscending, - }, fleet.ListOptions{}, "software") + }, fleet.ListOptions{}, "software", "") require.NoError(t, err) require.Len(t, policies, 3) assert.Equal(t, teamInstallerPolicy.ID, policies[0].ID) assert.Equal(t, teamAppStorePolicy.ID, policies[1].ID) assert.Equal(t, teamPatchPolicy.ID, policies[2].ID) - mergedCount, err = ds.CountPolicies(ctx, ptr.Uint(0), "", "software") + mergedCount, err = ds.CountPolicies(ctx, new(uint(0)), "", "software", "") require.NoError(t, err) assert.Equal(t, 3, mergedCount) } @@ -8563,7 +8671,7 @@ func testApplyPolicySpecNoSpuriousStatsReset(t *testing.T, ds *Datastore) { require.NoError(t, err) // Get the policy to find its ID. - policies, _, err := ds.ListTeamPolicies(ctx, team.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + policies, _, err := ds.ListTeamPolicies(ctx, team.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, policies, 1) pol := policies[0] @@ -8577,7 +8685,7 @@ func testApplyPolicySpecNoSpuriousStatsReset(t *testing.T, ds *Datastore) { require.NoError(t, err) // Verify the policy has a failing host count of 1. - policies, _, err = ds.ListTeamPolicies(ctx, team.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + policies, _, err = ds.ListTeamPolicies(ctx, team.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, policies, 1) require.Equal(t, uint(1), policies[0].FailingHostCount) @@ -8600,7 +8708,7 @@ func testApplyPolicySpecNoSpuriousStatsReset(t *testing.T, ds *Datastore) { })) // Verify that policy stats were NOT reset. - policies, _, err = ds.ListTeamPolicies(ctx, team.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + policies, _, err = ds.ListTeamPolicies(ctx, team.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, policies, 1) assert.Equal(t, uint(1), policies[0].FailingHostCount, "policy stats should not have been reset") diff --git a/server/datastore/mysql/software_installers_test.go b/server/datastore/mysql/software_installers_test.go index ced01a08d1..3e98a126ae 100644 --- a/server/datastore/mysql/software_installers_test.go +++ b/server/datastore/mysql/software_installers_test.go @@ -3277,7 +3277,7 @@ func testMatchOrCreateSoftwareInstallerWithAutomaticPolicies(t *testing.T, ds *D }) require.NoError(t, err) - team1Policies, _, err := ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + team1Policies, _, err := ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Empty(t, team1Policies) @@ -3298,7 +3298,7 @@ func testMatchOrCreateSoftwareInstallerWithAutomaticPolicies(t *testing.T, ds *D }) require.NoError(t, err) - team1Policies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + team1Policies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, team1Policies, 1) require.Equal(t, "[Install software] Foobar (pkg)", team1Policies[0].Name) @@ -3332,7 +3332,7 @@ func testMatchOrCreateSoftwareInstallerWithAutomaticPolicies(t *testing.T, ds *D }) require.NoError(t, err) - team1Policies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + team1Policies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, team1Policies, 2) require.Equal(t, "[Install software] FooFMA", team1Policies[1].Name) @@ -3373,7 +3373,7 @@ func testMatchOrCreateSoftwareInstallerWithAutomaticPolicies(t *testing.T, ds *D require.NoError(t, err) require.Equal(t, "upgradecode", msiThatShouldHaveUpgradeCode.UpgradeCode) - noTeamPolicies, _, err := ds.ListTeamPolicies(ctx, fleet.PolicyNoTeamID, fleet.ListOptions{}, fleet.ListOptions{}, "") + noTeamPolicies, _, err := ds.ListTeamPolicies(ctx, fleet.PolicyNoTeamID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, noTeamPolicies, 1) require.Equal(t, "[Install software] Zoobar (msi)", noTeamPolicies[0].Name) @@ -3402,7 +3402,7 @@ func testMatchOrCreateSoftwareInstallerWithAutomaticPolicies(t *testing.T, ds *D }) require.NoError(t, err) - team2Policies, _, err := ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + team2Policies, _, err := ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, team2Policies, 1) require.Equal(t, "[Install software] Barfoo (deb)", team2Policies[0].Name) @@ -3436,7 +3436,7 @@ Software won't be installed on Linux hosts with RPM-based distributions because }) require.NoError(t, err) - team2Policies, _, err = ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + team2Policies, _, err = ds.ListTeamPolicies(ctx, team2.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, team2Policies, 2) require.Equal(t, "[Install software] Barzoo (rpm)", team2Policies[1].Name) @@ -3476,7 +3476,7 @@ Software won't be installed on Linux hosts with Debian-based distributions becau }) require.NoError(t, err) - team1Policies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + team1Policies, _, err = ds.ListTeamPolicies(ctx, team1.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, team1Policies, 4) require.Equal(t, "[Install software] OtherFoobar (pkg) 2", team1Policies[3].Name) @@ -3518,7 +3518,7 @@ Software won't be installed on Linux hosts with Debian-based distributions becau }) require.NoError(t, err) - team3Policies, _, err := ds.ListTeamPolicies(ctx, team3.ID, fleet.ListOptions{}, fleet.ListOptions{}, "") + team3Policies, _, err := ds.ListTeamPolicies(ctx, team3.ID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") require.NoError(t, err) require.Len(t, team3Policies, 3) require.Equal(t, "[Install software] Something2 (msi) 3", team3Policies[2].Name) diff --git a/server/fleet/api_policies.go b/server/fleet/api_policies.go index ff1096ad5b..9e6c33feb2 100644 --- a/server/fleet/api_policies.go +++ b/server/fleet/api_policies.go @@ -30,7 +30,8 @@ func (r GlobalPolicyResponse) Error() error { return r.Err } ///////////////////////////////////////////////////////////////////////////////// type ListGlobalPoliciesRequest struct { - Opts ListOptions `url:"list_options"` + Opts ListOptions `url:"list_options"` + Platform string `query:"platform,optional"` } type ListGlobalPoliciesResponse struct { @@ -61,6 +62,7 @@ func (r GetPolicyByIDResponse) Error() error { return r.Err } type CountGlobalPoliciesRequest struct { ListOptions ListOptions `url:"list_options"` + Platform string `query:"platform,optional"` } type CountGlobalPoliciesResponse struct { @@ -206,6 +208,7 @@ type ListTeamPoliciesRequest struct { InheritedOrderKey string `query:"inherited_order_key,optional"` MergeInherited bool `query:"merge_inherited,optional"` AutomationType string `query:"automation_type,optional"` + Platform string `query:"platform,optional"` } type ListTeamPoliciesResponse struct { @@ -225,6 +228,7 @@ type CountTeamPoliciesRequest struct { TeamID uint `url:"fleet_id"` MergeInherited bool `query:"merge_inherited,optional"` AutomationType string `query:"automation_type,optional"` + Platform string `query:"platform,optional"` } type CountTeamPoliciesResponse struct { diff --git a/server/fleet/datastore.go b/server/fleet/datastore.go index a71dc70970..4a4fb2716e 100644 --- a/server/fleet/datastore.go +++ b/server/fleet/datastore.go @@ -958,11 +958,11 @@ type Datastore interface { // and resets automation retry attempts, identical to a query-change side-effect. ResetPolicy(ctx context.Context, policyID uint) error - ListGlobalPolicies(ctx context.Context, opts ListOptions) ([]*Policy, error) + ListGlobalPolicies(ctx context.Context, opts ListOptions, platform string) ([]*Policy, error) PoliciesByID(ctx context.Context, ids []uint) (map[uint]*Policy, error) DeleteGlobalPolicies(ctx context.Context, ids []uint) ([]uint, error) - CountPolicies(ctx context.Context, teamID *uint, matchQuery string, automationType string) (int, error) - CountMergedTeamPolicies(ctx context.Context, teamID uint, matchQuery string, automationType string) (int, error) + CountPolicies(ctx context.Context, teamID *uint, matchQuery string, automationType string, platform string) (int, error) + CountMergedTeamPolicies(ctx context.Context, teamID uint, matchQuery string, automationType string, platform string) (int, error) UpdateHostPolicyCounts(ctx context.Context) error PolicyQueriesForHost(ctx context.Context, host *Host) (map[string]string, error) @@ -1061,8 +1061,8 @@ type Datastore interface { // Team Policies NewTeamPolicy(ctx context.Context, teamID uint, authorID *uint, args PolicyPayload) (*Policy, error) - ListTeamPolicies(ctx context.Context, teamID uint, opts ListOptions, iopts ListOptions, automationType string) (teamPolicies, inheritedPolicies []*Policy, err error) - ListMergedTeamPolicies(ctx context.Context, teamID uint, opts ListOptions, automationType string) ([]*Policy, error) + ListTeamPolicies(ctx context.Context, teamID uint, opts ListOptions, iopts ListOptions, automationType string, platform string) (teamPolicies, inheritedPolicies []*Policy, err error) + ListMergedTeamPolicies(ctx context.Context, teamID uint, opts ListOptions, automationType string, platform string) ([]*Policy, error) DeleteTeamPolicies(ctx context.Context, teamID uint, ids []uint) ([]uint, error) TeamPolicy(ctx context.Context, teamID uint, policyID uint) (*Policy, error) diff --git a/server/fleet/policies.go b/server/fleet/policies.go index a1cd9d00f4..7ff1e63ad2 100644 --- a/server/fleet/policies.go +++ b/server/fleet/policies.go @@ -263,6 +263,21 @@ func verifyPolicyPlatforms(platforms string) error { return nil } +// ValidatePolicyPlatformFilter validates the platform query parameter used to +// filter policies on list/count endpoints. An empty string means "no filter" +// and is always valid; otherwise the value must be a single supported +// platform token. +func ValidatePolicyPlatformFilter(platform string) error { + if platform == "" { + return nil + } + switch platform { + case "windows", "linux", "darwin", "chrome": + return nil + } + return NewInvalidArgumentError("platform", `Invalid platform: must be one of "darwin", "windows", "linux", or "chrome".`) +} + func verifyPatchPolicy(team string, typ string) error { if typ == PolicyTypePatch && emptyString(team) { return errPatchPolicyRequiresTeam diff --git a/server/fleet/service.go b/server/fleet/service.go index cae50e6beb..c8b312cb98 100644 --- a/server/fleet/service.go +++ b/server/fleet/service.go @@ -741,14 +741,14 @@ type Service interface { // GlobalPolicyService NewGlobalPolicy(ctx context.Context, p PolicyPayload) (*Policy, error) - ListGlobalPolicies(ctx context.Context, opts ListOptions) ([]*Policy, error) + ListGlobalPolicies(ctx context.Context, opts ListOptions, platform string) ([]*Policy, error) DeleteGlobalPolicies(ctx context.Context, ids []uint) ([]uint, error) ModifyGlobalPolicy(ctx context.Context, id uint, p ModifyPolicyPayload) (*Policy, error) GetPolicyByID(ctx context.Context, policyID uint) (*Policy, error) ResetPolicy(ctx context.Context, policyID uint) error ListPolicyAutomationActivities(ctx context.Context, policyID uint, opts ListOptions, status string) ([]*PolicyAutomationActivity, *PaginationMetadata, error) ApplyPolicySpecs(ctx context.Context, policies []*PolicySpec) error - CountGlobalPolicies(ctx context.Context, matchQuery string) (int, error) + CountGlobalPolicies(ctx context.Context, matchQuery string, platform string) (int, error) AutofillPolicySql(ctx context.Context, sql string) (description string, resolution string, err error) // ///////////////////////////////////////////////////////////////////////////// @@ -873,11 +873,11 @@ type Service interface { // Team Policies NewTeamPolicy(ctx context.Context, teamID uint, p NewTeamPolicyPayload) (*Policy, error) - ListTeamPolicies(ctx context.Context, teamID uint, opts ListOptions, iopts ListOptions, mergeInherited bool, automationType string) (teamPolicies, inheritedPolicies []*Policy, err error) + ListTeamPolicies(ctx context.Context, teamID uint, opts ListOptions, iopts ListOptions, mergeInherited bool, automationType string, platform string) (teamPolicies, inheritedPolicies []*Policy, err error) DeleteTeamPolicies(ctx context.Context, teamID uint, ids []uint) ([]uint, error) ModifyTeamPolicy(ctx context.Context, teamID uint, id uint, p ModifyPolicyPayload) (*Policy, error) GetTeamPolicyByID(ctx context.Context, teamID uint, policyID uint) (*Policy, error) - CountTeamPolicies(ctx context.Context, teamID uint, matchQuery string, mergeInherited bool, automationType string) (int, int, error) + CountTeamPolicies(ctx context.Context, teamID uint, matchQuery string, mergeInherited bool, automationType string, platform string) (int, int, error) // ///////////////////////////////////////////////////////////////////////////// // Geolocation diff --git a/server/mock/datastore_mock.go b/server/mock/datastore_mock.go index acee47839f..5362d363f8 100644 --- a/server/mock/datastore_mock.go +++ b/server/mock/datastore_mock.go @@ -684,15 +684,15 @@ type SavePolicyFunc func(ctx context.Context, p *fleet.Policy, shouldRemoveAllPo type ResetPolicyFunc func(ctx context.Context, policyID uint) error -type ListGlobalPoliciesFunc func(ctx context.Context, opts fleet.ListOptions) ([]*fleet.Policy, error) +type ListGlobalPoliciesFunc func(ctx context.Context, opts fleet.ListOptions, platform string) ([]*fleet.Policy, error) type PoliciesByIDFunc func(ctx context.Context, ids []uint) (map[uint]*fleet.Policy, error) type DeleteGlobalPoliciesFunc func(ctx context.Context, ids []uint) ([]uint, error) -type CountPoliciesFunc func(ctx context.Context, teamID *uint, matchQuery string, automationType string) (int, error) +type CountPoliciesFunc func(ctx context.Context, teamID *uint, matchQuery string, automationType string, platform string) (int, error) -type CountMergedTeamPoliciesFunc func(ctx context.Context, teamID uint, matchQuery string, automationType string) (int, error) +type CountMergedTeamPoliciesFunc func(ctx context.Context, teamID uint, matchQuery string, automationType string, platform string) (int, error) type UpdateHostPolicyCountsFunc func(ctx context.Context) error @@ -772,9 +772,9 @@ type ListOutOfDateCalendarEventsFunc func(ctx context.Context, t time.Time) ([]* type NewTeamPolicyFunc func(ctx context.Context, teamID uint, authorID *uint, args fleet.PolicyPayload) (*fleet.Policy, error) -type ListTeamPoliciesFunc func(ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationType string) (teamPolicies []*fleet.Policy, inheritedPolicies []*fleet.Policy, err error) +type ListTeamPoliciesFunc func(ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationType string, platform string) (teamPolicies []*fleet.Policy, inheritedPolicies []*fleet.Policy, err error) -type ListMergedTeamPoliciesFunc func(ctx context.Context, teamID uint, opts fleet.ListOptions, automationType string) ([]*fleet.Policy, error) +type ListMergedTeamPoliciesFunc func(ctx context.Context, teamID uint, opts fleet.ListOptions, automationType string, platform string) ([]*fleet.Policy, error) type DeleteTeamPoliciesFunc func(ctx context.Context, teamID uint, ids []uint) ([]uint, error) @@ -7723,11 +7723,11 @@ func (s *DataStore) ResetPolicy(ctx context.Context, policyID uint) error { return s.ResetPolicyFunc(ctx, policyID) } -func (s *DataStore) ListGlobalPolicies(ctx context.Context, opts fleet.ListOptions) ([]*fleet.Policy, error) { +func (s *DataStore) ListGlobalPolicies(ctx context.Context, opts fleet.ListOptions, platform string) ([]*fleet.Policy, error) { s.mu.Lock() s.ListGlobalPoliciesFuncInvoked = true s.mu.Unlock() - return s.ListGlobalPoliciesFunc(ctx, opts) + return s.ListGlobalPoliciesFunc(ctx, opts, platform) } func (s *DataStore) PoliciesByID(ctx context.Context, ids []uint) (map[uint]*fleet.Policy, error) { @@ -7744,18 +7744,18 @@ func (s *DataStore) DeleteGlobalPolicies(ctx context.Context, ids []uint) ([]uin return s.DeleteGlobalPoliciesFunc(ctx, ids) } -func (s *DataStore) CountPolicies(ctx context.Context, teamID *uint, matchQuery string, automationType string) (int, error) { +func (s *DataStore) CountPolicies(ctx context.Context, teamID *uint, matchQuery string, automationType string, platform string) (int, error) { s.mu.Lock() s.CountPoliciesFuncInvoked = true s.mu.Unlock() - return s.CountPoliciesFunc(ctx, teamID, matchQuery, automationType) + return s.CountPoliciesFunc(ctx, teamID, matchQuery, automationType, platform) } -func (s *DataStore) CountMergedTeamPolicies(ctx context.Context, teamID uint, matchQuery string, automationType string) (int, error) { +func (s *DataStore) CountMergedTeamPolicies(ctx context.Context, teamID uint, matchQuery string, automationType string, platform string) (int, error) { s.mu.Lock() s.CountMergedTeamPoliciesFuncInvoked = true s.mu.Unlock() - return s.CountMergedTeamPoliciesFunc(ctx, teamID, matchQuery, automationType) + return s.CountMergedTeamPoliciesFunc(ctx, teamID, matchQuery, automationType, platform) } func (s *DataStore) UpdateHostPolicyCounts(ctx context.Context) error { @@ -8031,18 +8031,18 @@ func (s *DataStore) NewTeamPolicy(ctx context.Context, teamID uint, authorID *ui return s.NewTeamPolicyFunc(ctx, teamID, authorID, args) } -func (s *DataStore) ListTeamPolicies(ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationType string) (teamPolicies []*fleet.Policy, inheritedPolicies []*fleet.Policy, err error) { +func (s *DataStore) ListTeamPolicies(ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationType string, platform string) (teamPolicies []*fleet.Policy, inheritedPolicies []*fleet.Policy, err error) { s.mu.Lock() s.ListTeamPoliciesFuncInvoked = true s.mu.Unlock() - return s.ListTeamPoliciesFunc(ctx, teamID, opts, iopts, automationType) + return s.ListTeamPoliciesFunc(ctx, teamID, opts, iopts, automationType, platform) } -func (s *DataStore) ListMergedTeamPolicies(ctx context.Context, teamID uint, opts fleet.ListOptions, automationType string) ([]*fleet.Policy, error) { +func (s *DataStore) ListMergedTeamPolicies(ctx context.Context, teamID uint, opts fleet.ListOptions, automationType string, platform string) ([]*fleet.Policy, error) { s.mu.Lock() s.ListMergedTeamPoliciesFuncInvoked = true s.mu.Unlock() - return s.ListMergedTeamPoliciesFunc(ctx, teamID, opts, automationType) + return s.ListMergedTeamPoliciesFunc(ctx, teamID, opts, automationType, platform) } func (s *DataStore) DeleteTeamPolicies(ctx context.Context, teamID uint, ids []uint) ([]uint, error) { diff --git a/server/mock/service/service_mock.go b/server/mock/service/service_mock.go index d4e3b490bb..582a4196e6 100644 --- a/server/mock/service/service_mock.go +++ b/server/mock/service/service_mock.go @@ -454,7 +454,7 @@ type DeleteTeamScheduledQueriesFunc func(ctx context.Context, teamID uint, id ui type NewGlobalPolicyFunc func(ctx context.Context, p fleet.PolicyPayload) (*fleet.Policy, error) -type ListGlobalPoliciesFunc func(ctx context.Context, opts fleet.ListOptions) ([]*fleet.Policy, error) +type ListGlobalPoliciesFunc func(ctx context.Context, opts fleet.ListOptions, platform string) ([]*fleet.Policy, error) type DeleteGlobalPoliciesFunc func(ctx context.Context, ids []uint) ([]uint, error) @@ -468,7 +468,7 @@ type ListPolicyAutomationActivitiesFunc func(ctx context.Context, policyID uint, type ApplyPolicySpecsFunc func(ctx context.Context, policies []*fleet.PolicySpec) error -type CountGlobalPoliciesFunc func(ctx context.Context, matchQuery string) (int, error) +type CountGlobalPoliciesFunc func(ctx context.Context, matchQuery string, platform string) (int, error) type AutofillPolicySqlFunc func(ctx context.Context, sql string) (description string, resolution string, err error) @@ -534,7 +534,7 @@ type ListSoftwareByCVEFunc func(ctx context.Context, cve string, teamID *uint) ( type NewTeamPolicyFunc func(ctx context.Context, teamID uint, p fleet.NewTeamPolicyPayload) (*fleet.Policy, error) -type ListTeamPoliciesFunc func(ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, mergeInherited bool, automationType string) (teamPolicies []*fleet.Policy, inheritedPolicies []*fleet.Policy, err error) +type ListTeamPoliciesFunc func(ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, mergeInherited bool, automationType string, platform string) (teamPolicies []*fleet.Policy, inheritedPolicies []*fleet.Policy, err error) type DeleteTeamPoliciesFunc func(ctx context.Context, teamID uint, ids []uint) ([]uint, error) @@ -542,7 +542,7 @@ type ModifyTeamPolicyFunc func(ctx context.Context, teamID uint, id uint, p flee type GetTeamPolicyByIDFunc func(ctx context.Context, teamID uint, policyID uint) (*fleet.Policy, error) -type CountTeamPoliciesFunc func(ctx context.Context, teamID uint, matchQuery string, mergeInherited bool, automationType string) (int, int, error) +type CountTeamPoliciesFunc func(ctx context.Context, teamID uint, matchQuery string, mergeInherited bool, automationType string, platform string) (int, int, error) type LookupGeoIPFunc func(ctx context.Context, ip string) *fleet.GeoLocation @@ -3903,11 +3903,11 @@ func (s *Service) NewGlobalPolicy(ctx context.Context, p fleet.PolicyPayload) (* return s.NewGlobalPolicyFunc(ctx, p) } -func (s *Service) ListGlobalPolicies(ctx context.Context, opts fleet.ListOptions) ([]*fleet.Policy, error) { +func (s *Service) ListGlobalPolicies(ctx context.Context, opts fleet.ListOptions, platform string) ([]*fleet.Policy, error) { s.mu.Lock() s.ListGlobalPoliciesFuncInvoked = true s.mu.Unlock() - return s.ListGlobalPoliciesFunc(ctx, opts) + return s.ListGlobalPoliciesFunc(ctx, opts, platform) } func (s *Service) DeleteGlobalPolicies(ctx context.Context, ids []uint) ([]uint, error) { @@ -3952,11 +3952,11 @@ func (s *Service) ApplyPolicySpecs(ctx context.Context, policies []*fleet.Policy return s.ApplyPolicySpecsFunc(ctx, policies) } -func (s *Service) CountGlobalPolicies(ctx context.Context, matchQuery string) (int, error) { +func (s *Service) CountGlobalPolicies(ctx context.Context, matchQuery string, platform string) (int, error) { s.mu.Lock() s.CountGlobalPoliciesFuncInvoked = true s.mu.Unlock() - return s.CountGlobalPoliciesFunc(ctx, matchQuery) + return s.CountGlobalPoliciesFunc(ctx, matchQuery, platform) } func (s *Service) AutofillPolicySql(ctx context.Context, sql string) (description string, resolution string, err error) { @@ -4183,11 +4183,11 @@ func (s *Service) NewTeamPolicy(ctx context.Context, teamID uint, p fleet.NewTea return s.NewTeamPolicyFunc(ctx, teamID, p) } -func (s *Service) ListTeamPolicies(ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, mergeInherited bool, automationType string) (teamPolicies []*fleet.Policy, inheritedPolicies []*fleet.Policy, err error) { +func (s *Service) ListTeamPolicies(ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, mergeInherited bool, automationType string, platform string) (teamPolicies []*fleet.Policy, inheritedPolicies []*fleet.Policy, err error) { s.mu.Lock() s.ListTeamPoliciesFuncInvoked = true s.mu.Unlock() - return s.ListTeamPoliciesFunc(ctx, teamID, opts, iopts, mergeInherited, automationType) + return s.ListTeamPoliciesFunc(ctx, teamID, opts, iopts, mergeInherited, automationType, platform) } func (s *Service) DeleteTeamPolicies(ctx context.Context, teamID uint, ids []uint) ([]uint, error) { @@ -4211,11 +4211,11 @@ func (s *Service) GetTeamPolicyByID(ctx context.Context, teamID uint, policyID u return s.GetTeamPolicyByIDFunc(ctx, teamID, policyID) } -func (s *Service) CountTeamPolicies(ctx context.Context, teamID uint, matchQuery string, mergeInherited bool, automationType string) (int, int, error) { +func (s *Service) CountTeamPolicies(ctx context.Context, teamID uint, matchQuery string, mergeInherited bool, automationType string, platform string) (int, int, error) { s.mu.Lock() s.CountTeamPoliciesFuncInvoked = true s.mu.Unlock() - return s.CountTeamPoliciesFunc(ctx, teamID, matchQuery, mergeInherited, automationType) + return s.CountTeamPoliciesFunc(ctx, teamID, matchQuery, mergeInherited, automationType, platform) } func (s *Service) LookupGeoIP(ctx context.Context, ip string) *fleet.GeoLocation { diff --git a/server/service/global_policies.go b/server/service/global_policies.go index a1232c8291..0701c41161 100644 --- a/server/service/global_policies.go +++ b/server/service/global_policies.go @@ -108,19 +108,23 @@ func (svc Service) NewGlobalPolicy(ctx context.Context, p fleet.PolicyPayload) ( func listGlobalPoliciesEndpoint(ctx context.Context, request interface{}, svc fleet.Service) (fleet.Errorer, error) { req := request.(*fleet.ListGlobalPoliciesRequest) - resp, err := svc.ListGlobalPolicies(ctx, req.Opts) + resp, err := svc.ListGlobalPolicies(ctx, req.Opts, req.Platform) if err != nil { return fleet.ListGlobalPoliciesResponse{Err: err}, nil } return fleet.ListGlobalPoliciesResponse{Policies: resp}, nil } -func (svc Service) ListGlobalPolicies(ctx context.Context, opts fleet.ListOptions) ([]*fleet.Policy, error) { +func (svc Service) ListGlobalPolicies(ctx context.Context, opts fleet.ListOptions, platform string) ([]*fleet.Policy, error) { if err := svc.authz.Authorize(ctx, &fleet.Policy{}, fleet.ActionRead); err != nil { return nil, err } - return svc.ds.ListGlobalPolicies(ctx, opts) + if err := fleet.ValidatePolicyPlatformFilter(platform); err != nil { + return nil, ctxerr.Wrap(ctx, err) + } + + return svc.ds.ListGlobalPolicies(ctx, opts, platform) } // /////////////////////////////////////////////////////////////////////////////// @@ -129,19 +133,23 @@ func (svc Service) ListGlobalPolicies(ctx context.Context, opts fleet.ListOption func countGlobalPoliciesEndpoint(ctx context.Context, request interface{}, svc fleet.Service) (fleet.Errorer, error) { req := request.(*fleet.CountGlobalPoliciesRequest) - resp, err := svc.CountGlobalPolicies(ctx, req.ListOptions.MatchQuery) + resp, err := svc.CountGlobalPolicies(ctx, req.ListOptions.MatchQuery, req.Platform) if err != nil { return fleet.CountGlobalPoliciesResponse{Err: err}, nil } return fleet.CountGlobalPoliciesResponse{Count: resp}, nil } -func (svc Service) CountGlobalPolicies(ctx context.Context, matchQuery string) (int, error) { +func (svc Service) CountGlobalPolicies(ctx context.Context, matchQuery string, platform string) (int, error) { if err := svc.authz.Authorize(ctx, &fleet.Policy{}, fleet.ActionRead); err != nil { return 0, err } - count, err := svc.ds.CountPolicies(ctx, nil, matchQuery, "") + if err := fleet.ValidatePolicyPlatformFilter(platform); err != nil { + return 0, ctxerr.Wrap(ctx, err) + } + + count, err := svc.ds.CountPolicies(ctx, nil, matchQuery, "", platform) if err != nil { return 0, err } @@ -284,7 +292,7 @@ func (svc *Service) ResetAutomation(ctx context.Context, teamIDs, policyIDs []ui pIDs[id] = struct{}{} } for _, teamID := range teamIDs { - p1, p2, err := svc.ds.ListTeamPolicies(ctx, teamID, fleet.ListOptions{}, fleet.ListOptions{}, "") + p1, p2, err := svc.ds.ListTeamPolicies(ctx, teamID, fleet.ListOptions{}, fleet.ListOptions{}, "", "") if err != nil { return err } diff --git a/server/service/global_policies_test.go b/server/service/global_policies_test.go index 7d6396599d..91da610c50 100644 --- a/server/service/global_policies_test.go +++ b/server/service/global_policies_test.go @@ -44,7 +44,7 @@ func TestGlobalPoliciesAuth(t *testing.T) { ds.NewGlobalPolicyFunc = func(ctx context.Context, authorID *uint, args fleet.PolicyPayload) (*fleet.Policy, error) { return &fleet.Policy{}, nil } - ds.ListGlobalPoliciesFunc = func(ctx context.Context, opts fleet.ListOptions) ([]*fleet.Policy, error) { + ds.ListGlobalPoliciesFunc = func(ctx context.Context, opts fleet.ListOptions, platform string) ([]*fleet.Policy, error) { return nil, nil } ds.PoliciesByIDFunc = func(ctx context.Context, ids []uint) (map[uint]*fleet.Policy, error) { @@ -132,7 +132,7 @@ func TestGlobalPoliciesAuth(t *testing.T) { }) checkAuthErr(t, tt.shouldFailWrite, err) - _, err = svc.ListGlobalPolicies(ctx, fleet.ListOptions{}) + _, err = svc.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") checkAuthErr(t, tt.shouldFailRead, err) _, err = svc.GetPolicyByID(ctx, 1) diff --git a/server/service/integration_enterprise_test.go b/server/service/integration_enterprise_test.go index 15d6b991a6..6d720cf0d3 100644 --- a/server/service/integration_enterprise_test.go +++ b/server/service/integration_enterprise_test.go @@ -31314,7 +31314,7 @@ func (s *integrationEnterpriseTestSuite) TestApplyPolicySpecsBatchMixedScopes() invalidName := "batch-invalid-" + t.Name() assertNonePersisted := func(label string) { - policies, err := s.ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + policies, err := s.ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) for _, p := range policies { require.NotEqual(t, validAnyName, p.Name, "%s: no spec from rejected batch should persist", label) @@ -31357,7 +31357,7 @@ func (s *integrationEnterpriseTestSuite) TestApplyPolicySpecsBatchMixedScopes() }, }, http.StatusOK) - policies, err := s.ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + policies, err := s.ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) byName := make(map[string]*fleet.Policy, len(policies)) for _, p := range policies { diff --git a/server/service/team_policies.go b/server/service/team_policies.go index 0e898fca0a..a89954e9ee 100644 --- a/server/service/team_policies.go +++ b/server/service/team_policies.go @@ -334,14 +334,14 @@ func listTeamPoliciesEndpoint(ctx context.Context, request interface{}, svc flee OrderKey: req.InheritedOrderKey, } - tmPols, inheritedPols, err := svc.ListTeamPolicies(ctx, req.TeamID, req.Opts, inheritedListOptions, req.MergeInherited, req.AutomationType) + tmPols, inheritedPols, err := svc.ListTeamPolicies(ctx, req.TeamID, req.Opts, inheritedListOptions, req.MergeInherited, req.AutomationType, req.Platform) if err != nil { return fleet.ListTeamPoliciesResponse{Err: err}, nil } return fleet.ListTeamPoliciesResponse{Policies: tmPols, InheritedPolicies: inheritedPols}, nil } -func (svc *Service) ListTeamPolicies(ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, mergeInherited bool, automationFilter string) (teamPolicies, inheritedPolicies []*fleet.Policy, err error) { +func (svc *Service) ListTeamPolicies(ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, mergeInherited bool, automationFilter string, platform string) (teamPolicies, inheritedPolicies []*fleet.Policy, err error) { if err := svc.authz.Authorize(ctx, &fleet.Policy{ PolicyData: fleet.PolicyData{ TeamID: ptr.Uint(teamID), @@ -350,6 +350,10 @@ func (svc *Service) ListTeamPolicies(ctx context.Context, teamID uint, opts flee return nil, nil, err } + if err := fleet.ValidatePolicyPlatformFilter(platform); err != nil { + return nil, nil, ctxerr.Wrap(ctx, err) + } + if teamID > 0 { if _, err := svc.ds.TeamLite(ctx, teamID); err != nil { // TODO see if we can use TeamExists here instead return nil, nil, ctxerr.Wrapf(ctx, err, "loading team %d", teamID) @@ -357,7 +361,7 @@ func (svc *Service) ListTeamPolicies(ctx context.Context, teamID uint, opts flee } if mergeInherited { - policies, err := svc.ds.ListMergedTeamPolicies(ctx, teamID, opts, automationFilter) + policies, err := svc.ds.ListMergedTeamPolicies(ctx, teamID, opts, automationFilter, platform) if err != nil { return nil, nil, err } @@ -372,7 +376,7 @@ func (svc *Service) ListTeamPolicies(ctx context.Context, teamID uint, opts flee return policies, nil, nil } - teamPolicies, inheritedPolicies, err = svc.ds.ListTeamPolicies(ctx, teamID, opts, iopts, automationFilter) + teamPolicies, inheritedPolicies, err = svc.ds.ListTeamPolicies(ctx, teamID, opts, iopts, automationFilter, platform) if err != nil { return nil, nil, err } @@ -395,14 +399,14 @@ func (svc *Service) ListTeamPolicies(ctx context.Context, teamID uint, opts flee func countTeamPoliciesEndpoint(ctx context.Context, request interface{}, svc fleet.Service) (fleet.Errorer, error) { req := request.(*fleet.CountTeamPoliciesRequest) - count, inheritedCount, err := svc.CountTeamPolicies(ctx, req.TeamID, req.ListOptions.MatchQuery, req.MergeInherited, req.AutomationType) + count, inheritedCount, err := svc.CountTeamPolicies(ctx, req.TeamID, req.ListOptions.MatchQuery, req.MergeInherited, req.AutomationType, req.Platform) if err != nil { return fleet.CountTeamPoliciesResponse{Err: err}, nil } return fleet.CountTeamPoliciesResponse{Count: count, InheritedPolicyCount: inheritedCount}, nil } -func (svc *Service) CountTeamPolicies(ctx context.Context, teamID uint, matchQuery string, mergeInherited bool, automationType string) (int, int, error) { +func (svc *Service) CountTeamPolicies(ctx context.Context, teamID uint, matchQuery string, mergeInherited bool, automationType string, platform string) (int, int, error) { if err := svc.authz.Authorize(ctx, &fleet.Policy{ PolicyData: fleet.PolicyData{ TeamID: ptr.Uint(teamID), @@ -411,6 +415,10 @@ func (svc *Service) CountTeamPolicies(ctx context.Context, teamID uint, matchQue return 0, 0, err } + if err := fleet.ValidatePolicyPlatformFilter(platform); err != nil { + return 0, 0, ctxerr.Wrap(ctx, err) + } + if teamID > 0 { if _, err := svc.ds.TeamLite(ctx, teamID); err != nil { // TODO see if we can use TeamExists here instead return 0, 0, ctxerr.Wrapf(ctx, err, "loading team %d", teamID) @@ -418,18 +426,24 @@ func (svc *Service) CountTeamPolicies(ctx context.Context, teamID uint, matchQue } if mergeInherited { - count, err := svc.ds.CountMergedTeamPolicies(ctx, teamID, matchQuery, automationType) + count, err := svc.ds.CountMergedTeamPolicies(ctx, teamID, matchQuery, automationType, platform) if err != nil { return 0, 0, err } - inheritedCount, err := svc.ds.CountPolicies(ctx, nil, matchQuery, automationType) + // CountPolicies ignores automationType when teamID is nil, so the + // inherited count would be wrong (too high) when an automation filter + // is active. Short-circuit to 0 in that case. + if automationType != "" { + return count, 0, nil + } + inheritedCount, err := svc.ds.CountPolicies(ctx, nil, matchQuery, automationType, platform) if err != nil { return 0, 0, err } return count, inheritedCount, nil } - count, err := svc.ds.CountPolicies(ctx, &teamID, matchQuery, automationType) + count, err := svc.ds.CountPolicies(ctx, &teamID, matchQuery, automationType, platform) if err != nil { return 0, 0, err } diff --git a/server/service/team_policies_test.go b/server/service/team_policies_test.go index de02fef22e..24c8e07f8b 100644 --- a/server/service/team_policies_test.go +++ b/server/service/team_policies_test.go @@ -27,7 +27,7 @@ func TestTeamPoliciesAuth(t *testing.T) { }, }, nil } - ds.ListTeamPoliciesFunc = func(ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationFilter string) (tpol, ipol []*fleet.Policy, err error) { + ds.ListTeamPoliciesFunc = func(ctx context.Context, teamID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationFilter string, platform string) (tpol, ipol []*fleet.Policy, err error) { return nil, nil, nil } ds.PoliciesByIDFunc = func(ctx context.Context, ids []uint) (map[uint]*fleet.Policy, error) { @@ -155,7 +155,7 @@ func TestTeamPoliciesAuth(t *testing.T) { }) checkAuthErr(t, tt.shouldFailWrite, err) - _, _, err = svc.ListTeamPolicies(ctx, 1, fleet.ListOptions{}, fleet.ListOptions{}, false, "") + _, _, err = svc.ListTeamPolicies(ctx, 1, fleet.ListOptions{}, fleet.ListOptions{}, false, "", "") checkAuthErr(t, tt.shouldFailRead, err) _, err = svc.GetTeamPolicyByID(ctx, 1, 1) @@ -257,10 +257,10 @@ func TestTeamPolicyAutomationsPopulated(t *testing.T) { ds.TeamPolicyFunc = func(ctx context.Context, tID uint, id uint) (*fleet.Policy, error) { return freshPolicy(), nil } - ds.ListTeamPoliciesFunc = func(ctx context.Context, tID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationFilter string) ([]*fleet.Policy, []*fleet.Policy, error) { + ds.ListTeamPoliciesFunc = func(ctx context.Context, tID uint, opts fleet.ListOptions, iopts fleet.ListOptions, automationFilter string, platform string) ([]*fleet.Policy, []*fleet.Policy, error) { return []*fleet.Policy{freshPolicy()}, nil, nil } - ds.ListMergedTeamPoliciesFunc = func(ctx context.Context, tID uint, opts fleet.ListOptions, automationFilter string) ([]*fleet.Policy, error) { + ds.ListMergedTeamPoliciesFunc = func(ctx context.Context, tID uint, opts fleet.ListOptions, automationFilter string, platform string) ([]*fleet.Policy, error) { return []*fleet.Policy{freshPolicy()}, nil } ds.SavePolicyFunc = func(ctx context.Context, p *fleet.Policy, _ bool, _ bool) error { @@ -384,7 +384,7 @@ func TestTeamPolicyAutomationsPopulated(t *testing.T) { svc, baseCtx := newTestService(t, ds, nil, nil) ctx := adminCtx(baseCtx) - teamPols, _, err := svc.ListTeamPolicies(ctx, teamID, fleet.ListOptions{}, fleet.ListOptions{}, false, "") + teamPols, _, err := svc.ListTeamPolicies(ctx, teamID, fleet.ListOptions{}, fleet.ListOptions{}, false, "", "") require.NoError(t, err) require.Len(t, teamPols, 1) requireAutomationsPopulated(t, teamPols[0]) @@ -396,7 +396,7 @@ func TestTeamPolicyAutomationsPopulated(t *testing.T) { svc, baseCtx := newTestService(t, ds, nil, nil) ctx := adminCtx(baseCtx) - merged, _, err := svc.ListTeamPolicies(ctx, teamID, fleet.ListOptions{}, fleet.ListOptions{}, true, "") + merged, _, err := svc.ListTeamPolicies(ctx, teamID, fleet.ListOptions{}, fleet.ListOptions{}, true, "", "") require.NoError(t, err) require.Len(t, merged, 1) requireAutomationsPopulated(t, merged[0]) diff --git a/server/service/testing_client_test.go b/server/service/testing_client_test.go index 5f1770b09e..98e613a59d 100644 --- a/server/service/testing_client_test.go +++ b/server/service/testing_client_test.go @@ -235,7 +235,7 @@ func (ts *withServer) commonTearDownTest(t *testing.T) { return err }) - globalPolicies, err := ts.ds.ListGlobalPolicies(ctx, fleet.ListOptions{}) + globalPolicies, err := ts.ds.ListGlobalPolicies(ctx, fleet.ListOptions{}, "") require.NoError(t, err) if len(globalPolicies) > 0 { var globalPolicyIDs []uint