From 82f241628cab68fa1c61afa25d0fd10091e2e46d Mon Sep 17 00:00:00 2001 From: Bryan Frimin Date: Wed, 18 Mar 2026 16:00:40 +0100 Subject: [PATCH] Refactor MeasuresPage to use Relay fragments Replace the client-side grouped-by-category view (fetching 500 items) with a flat table using server-side filtering and cursor-based pagination. Colocate GraphQL queries, fragments, and mutations in the component file per console CLAUDE.md conventions. Backend changes add a category filter to the measure list endpoints (GraphQL, MCP) and a new measureCategories field on Organization. Signed-off-by: Bryan Frimin --- apps/console/src/hooks/graph/MeasureGraph.ts | 15 +- .../organizations/measures/MeasuresPage.tsx | 475 +++++++++++------- .../measures/MeasuresPageLoader.tsx | 28 ++ apps/console/src/routes/measureRoutes.ts | 13 +- pkg/coredata/measure.go | 35 ++ pkg/coredata/measure_filter.go | 28 +- pkg/probo/framework_service.go | 2 +- pkg/probo/measure_service.go | 37 ++ pkg/server/api/console/v1/schema.graphql | 3 + pkg/server/api/console/v1/v1_resolver.go | 29 +- pkg/server/api/mcp/v1/schema.resolvers.go | 6 +- pkg/server/api/mcp/v1/specification.yaml | 3 + 12 files changed, 459 insertions(+), 215 deletions(-) create mode 100644 apps/console/src/pages/organizations/measures/MeasuresPageLoader.tsx diff --git a/apps/console/src/hooks/graph/MeasureGraph.ts b/apps/console/src/hooks/graph/MeasureGraph.ts index ad86b6fa8..850517bd7 100644 --- a/apps/console/src/hooks/graph/MeasureGraph.ts +++ b/apps/console/src/hooks/graph/MeasureGraph.ts @@ -7,18 +7,7 @@ import { useMutationWithToasts } from "../useMutationWithToasts"; /* eslint-disable relay/unused-fields, relay/must-colocate-fragment-spreads */ -export const measuresQuery = graphql` - query MeasureGraphListQuery($organizationId: ID!) { - organization: node(id: $organizationId) @required(action: THROW) { - __typename - ... on Organization { - id - canCreateMeasure: permission(action: "core:measure:create") - ...MeasuresPageFragment - } - } - } -`; +export const MeasureConnectionKey = "MeasuresPage_measures"; const deleteMeasureMutation = graphql` mutation MeasureGraphDeleteMutation( @@ -31,8 +20,6 @@ const deleteMeasureMutation = graphql` } `; -export const MeasureConnectionKey = "MeasuresGraphListQuery__measures"; - export function useDeleteMeasureMutation() { const { __ } = useTranslate(); diff --git a/apps/console/src/pages/organizations/measures/MeasuresPage.tsx b/apps/console/src/pages/organizations/measures/MeasuresPage.tsx index 476efe376..55b8432e3 100644 --- a/apps/console/src/pages/organizations/measures/MeasuresPage.tsx +++ b/apps/console/src/pages/organizations/measures/MeasuresPage.tsx @@ -1,4 +1,9 @@ -import { groupBy, objectKeys, slugify, sprintf } from "@probo/helpers"; +import { + formatError, + getMeasureStateLabel, + sprintf, + type GraphQLError, +} from "@probo/helpers"; import { usePageTitle } from "@probo/hooks"; import { useTranslate } from "@probo/i18n"; import { @@ -7,14 +12,13 @@ import { Card, DropdownItem, FileButton, - IconChevronDown, - IconChevronUp, IconFolderUpload, IconPencil, IconPlusLarge, IconTrashCan, - MeasureImplementation, + Option, PageHeader, + Select, Table, Tbody, Td, @@ -23,52 +27,103 @@ import { Tr, useConfirm, useDialogRef, + useToast, } from "@probo/ui"; import { MeasureBadge } from "@probo/ui/src/Molecules/Badge/MeasureBadge"; -import { type ChangeEventHandler, useMemo, useRef, useState } from "react"; +import { type ChangeEventHandler, useRef, useState, useTransition } from "react"; import { + ConnectionHandler, + graphql, type PreloadedQuery, useFragment, + useMutation, + usePaginationFragment, usePreloadedQuery, } from "react-relay"; -import { Link, useParams } from "react-router"; -import { graphql } from "relay-runtime"; -import type { MeasureGraphListQuery } from "#/__generated__/core/MeasureGraphListQuery.graphql"; -import type { - MeasuresPageFragment$data, - MeasuresPageFragment$key, -} from "#/__generated__/core/MeasuresPageFragment.graphql"; +import type { MeasuresPageDeleteMutation } from "#/__generated__/core/MeasuresPageDeleteMutation.graphql"; +import type { MeasuresPageFragment$key } from "#/__generated__/core/MeasuresPageFragment.graphql"; import type { MeasuresPageImportMutation } from "#/__generated__/core/MeasuresPageImportMutation.graphql"; -import { - measuresQuery, - useDeleteMeasureMutation, -} from "#/hooks/graph/MeasureGraph"; +import type { MeasuresPageListQuery } from "#/__generated__/core/MeasuresPageListQuery.graphql"; +import type { + MeasuresPageRefetchQuery, + MeasureState, +} from "#/__generated__/core/MeasuresPageRefetchQuery.graphql"; +import type { MeasuresPageRowFragment$key } from "#/__generated__/core/MeasuresPageRowFragment.graphql"; import { useMutationWithToasts } from "#/hooks/useMutationWithToasts"; import { useOrganizationId } from "#/hooks/useOrganizationId"; -import type { NodeOf } from "#/types"; import MeasureFormDialog from "./dialog/MeasureFormDialog"; -type Props = { - queryRef: PreloadedQuery; -}; +export const MeasuresConnectionKey = "MeasuresPage_measures"; -const measuresFragment = graphql` - fragment MeasuresPageFragment on Organization { - measures(first: 500) @connection(key: "MeasuresGraphListQuery__measures") { - __id +export const measuresPageQuery = graphql` + query MeasuresPageListQuery($organizationId: ID!) { + organization: node(id: $organizationId) @required(action: THROW) { + __typename + ... on Organization { + canCreateMeasure: permission(action: "core:measure:create") + measureCategories + ...MeasuresPageFragment + } + } + } +`; + +const measureRowFragment = graphql` + fragment MeasuresPageRowFragment on Measure { + id + name + category + state + canUpdate: permission(action: "core:measure:update") + canDelete: permission(action: "core:measure:delete") + ...MeasureFormDialogMeasureFragment + } +`; + +const deleteMeasureMutation = graphql` + mutation MeasuresPageDeleteMutation( + $input: DeleteMeasureInput! + $connections: [ID!]! + ) { + deleteMeasure(input: $input) { + deletedMeasureId @deleteEdge(connections: $connections) + } + } +`; + +const measuresPageFragment = graphql` + fragment MeasuresPageFragment on Organization + @refetchable(queryName: "MeasuresPageRefetchQuery") + @argumentDefinitions( + first: { type: "Int", defaultValue: 20 } + after: { type: "CursorKey" } + state: { type: "MeasureState", defaultValue: null } + category: { type: "String", defaultValue: null } + ) { + id + measures( + first: $first + after: $after + filter: { state: $state, category: $category } + ) + @connection( + key: "MeasuresPage_measures" + filters: ["filter"] + ) { edges { node { id - name - category - state canUpdate: permission(action: "core:measure:update") canDelete: permission(action: "core:measure:delete") - ...MeasureFormDialogMeasureFragment + ...MeasuresPageRowFragment } } + pageInfo { + hasNextPage + endCursor + } } } `; @@ -91,24 +146,79 @@ const importMeasuresMutation = graphql` } `; -export default function MeasuresPage(props: Props) { +interface MeasuresPageProps { + queryRef: PreloadedQuery; +} + +export default function MeasuresPage({ queryRef }: MeasuresPageProps) { const { __ } = useTranslate(); - const organization = usePreloadedQuery( - measuresQuery, - props.queryRef, - ).organization; + const organizationId = useOrganizationId(); + + usePageTitle(__("Measures")); + + const { organization } = usePreloadedQuery(measuresPageQuery, queryRef); if (organization.__typename !== "Organization") { throw new Error("invalid node type"); } - const data = useFragment( - measuresFragment, - organization, + + const [isPending, startTransition] = useTransition(); + const [stateFilter, setStateFilter] = useState(null); + const [categoryFilter, setCategoryFilter] = useState(null); + + const { data, loadNext, hasNext, isLoadingNext, refetch } + = usePaginationFragment( + measuresPageFragment, + organization, + ); + + const refetchFilters = (overrides: Record = {}) => { + startTransition(() => { + refetch( + { + state: stateFilter, + category: categoryFilter, + ...overrides, + }, + { fetchPolicy: "network-only" }, + ); + }); + }; + + const handleStateFilterChange = (value: string) => { + const newState = value === "ALL" ? null : (value as MeasureState); + setStateFilter(newState); + refetchFilters({ state: newState }); + }; + + const handleCategoryFilterChange = (value: string) => { + const newCategory = value === "ALL" ? null : value; + setCategoryFilter(newCategory); + refetchFilters({ category: newCategory }); + }; + + const currentFilter = { + state: stateFilter, + category: categoryFilter, + }; + + const connectionId = ConnectionHandler.getConnectionID( + organizationId, + MeasuresConnectionKey, + { filter: currentFilter }, ); - const connectionId = data.measures.__id; - const measures = data.measures.edges.map(edge => edge.node); - const measuresPerCategory = useMemo(() => { - return groupBy(measures, measure => measure.category); - }, [measures]); + const allFiltersNullConnectionId = ConnectionHandler.getConnectionID( + organizationId, + MeasuresConnectionKey, + { filter: { state: null, category: null } }, + ); + const hasActiveFilter = stateFilter || categoryFilter; + const createConnectionIds = hasActiveFilter + ? [allFiltersNullConnectionId, connectionId] + : [connectionId]; + + const measures = data?.measures?.edges?.map(edge => edge.node) ?? []; + const categories = organization.measureCategories ?? []; + const [importMeasures] = useMutationWithToasts( importMeasuresMutation, { @@ -117,7 +227,6 @@ export default function MeasuresPage(props: Props) { }, ); const importFileRef = useRef(null); - usePageTitle(__("Measures")); const handleImport: ChangeEventHandler = (event) => { const file = event.target.files?.[0]; @@ -127,10 +236,10 @@ export default function MeasuresPage(props: Props) { void importMeasures({ variables: { input: { - organizationId: organization.id, + organizationId, file: null, }, - connections: [connectionId], + connections: createConnectionIds, }, uploadables: { "input.file": file, @@ -141,6 +250,10 @@ export default function MeasuresPage(props: Props) { }); }; + const hasAnyAction = measures.some( + ({ canUpdate, canDelete }) => canUpdate || canDelete, + ); + return (
)} - - {objectKeys(measuresPerCategory) - .sort((a, b) => a.localeCompare(b)) - .map(category => ( - - ))} + +
+ + +
+ +
+ {measures.length > 0 + ? ( + + + + + + + + {hasAnyAction && + + + {measures.map(measure => ( + + ))} + +
{__("Measure")}{__("Category")}{__("State")}} +
+ + {hasNext && ( +
+ +
+ )} +
+ ) + : ( + +
+

+ {__("No measures yet")} +

+

+ {__("Create your first measure to get started.")} +

+
+
+ )} +
); } -type CategoryProps = { - category: string; - measures: NodeOf[]; - connectionId: string; -}; - -function Category(props: CategoryProps) { - const params = useParams<{ categoryId?: string }>(); - const { __ } = useTranslate(); - const organizationId = useOrganizationId(); - const categoryId = slugify(props.category); - const [limit, setLimit] = useState(4); - const measures = useMemo(() => { - return limit ? props.measures.slice(0, limit) : props.measures; - }, [props.measures, limit]); - const showMoreButton = limit !== null && props.measures.length > limit; - const isExpanded = categoryId === params.categoryId; - const ExpandComponent = isExpanded ? IconChevronUp : IconChevronDown; - const completedMeasures = props.measures.filter( - m => m.state === "IMPLEMENTED", - ); - - return ( - - -

{props.category}

-
- - {__("Completion")} - : - {" "} - - {completedMeasures.length} - / - {props.measures.length} - - - | - -
- - {isExpanded && ( -
- - - - - - {measures.some( - ({ canUpdate, canDelete }) => canUpdate || canDelete, - ) && } - - - - {measures.map(measure => ( - - ))} - -
{__("Measure")}{__("State")}
- {showMoreButton && ( - - )} -
- )} -
- ); -} - type MeasureRowProps = { - measure: NodeOf; + measureKey: MeasuresPageRowFragment$key; connectionId: string; hasAnyAction: boolean; }; function MeasureRow(props: MeasureRowProps) { - const { __ } = useTranslate(); - const [deleteMeasure, isDeleting] = useDeleteMeasureMutation(); - const confirm = useConfirm(); + const measure = useFragment(measureRowFragment, props.measureKey); const organizationId = useOrganizationId(); + const { __ } = useTranslate(); + const [deleteMeasure] = useMutation(deleteMeasureMutation); + const { toast } = useToast(); + const confirm = useConfirm(); + const dialogRef = useDialogRef(); - const onDelete = () => { + const handleDelete = () => { confirm( () => new Promise((resolve) => { - void deleteMeasure({ + deleteMeasure({ variables: { - input: { measureId: props.measure.id }, + input: { measureId: measure.id }, connections: [props.connectionId], }, - onCompleted: () => resolve(), + onCompleted(_, error) { + if (error) { + toast({ + title: __("Error"), + description: formatError( + __("Failed to delete measure"), + error as GraphQLError[], + ), + variant: "error", + }); + } else { + toast({ + title: __("Success"), + description: __("Measure deleted successfully"), + variant: "success", + }); + } + resolve(); + }, + onError(error) { + toast({ + title: __("Error"), + description: formatError( + __("Failed to delete measure"), + error as GraphQLError, + ), + variant: "error", + }); + resolve(); + }, }); }), { @@ -294,44 +421,44 @@ function MeasureRow(props: MeasureRowProps) { __( "This will permanently delete the measure \"%s\". This action cannot be undone.", ), - props.measure.name, + measure.name, ), }, ); }; - const dialogRef = useDialogRef(); - return ( <> - - - {props.measure.name} + + + {measure.name} + {measure.category} - + - {(props.measure.canUpdate || props.measure.canDelete) && ( + {props.hasAnyAction && ( - - {props.measure.canUpdate && ( - dialogRef.current?.open()} - > - {__("Edit")} - - )} - {props.measure.canDelete && ( - - {__("Delete")} - - )} - + {(measure.canUpdate || measure.canDelete) && ( + + {measure.canUpdate && ( + dialogRef.current?.open()} + > + {__("Edit")} + + )} + {measure.canDelete && ( + + {__("Delete")} + + )} + + )} )} diff --git a/apps/console/src/pages/organizations/measures/MeasuresPageLoader.tsx b/apps/console/src/pages/organizations/measures/MeasuresPageLoader.tsx new file mode 100644 index 000000000..7d6b35681 --- /dev/null +++ b/apps/console/src/pages/organizations/measures/MeasuresPageLoader.tsx @@ -0,0 +1,28 @@ +import { Suspense, useEffect } from "react"; +import { useQueryLoader } from "react-relay"; + +import type { MeasuresPageListQuery } from "#/__generated__/core/MeasuresPageListQuery.graphql"; +import { PageSkeleton } from "#/components/skeletons/PageSkeleton"; +import { useOrganizationId } from "#/hooks/useOrganizationId"; + +import MeasuresPage, { measuresPageQuery } from "./MeasuresPage"; + +export default function MeasuresPageLoader() { + const organizationId = useOrganizationId(); + const [queryRef, loadQuery] + = useQueryLoader(measuresPageQuery); + + useEffect(() => { + loadQuery({ organizationId }); + }, [loadQuery, organizationId]); + + if (!queryRef) { + return ; + } + + return ( + }> + + + ); +} diff --git a/apps/console/src/routes/measureRoutes.ts b/apps/console/src/routes/measureRoutes.ts index 2dfbf63b5..110bed6cb 100644 --- a/apps/console/src/routes/measureRoutes.ts +++ b/apps/console/src/routes/measureRoutes.ts @@ -8,24 +8,19 @@ import { Fragment } from "react"; import { loadQuery } from "react-relay"; import { redirect } from "react-router"; -import type { MeasureGraphListQuery } from "#/__generated__/core/MeasureGraphListQuery.graphql"; import type { MeasureGraphNodeQuery } from "#/__generated__/core/MeasureGraphNodeQuery.graphql"; import { LinkCardSkeleton } from "#/components/skeletons/LinkCardSkeleton"; import { PageSkeleton } from "#/components/skeletons/PageSkeleton"; import { coreEnvironment } from "#/environments"; -import { measureNodeQuery, measuresQuery } from "#/hooks/graph/MeasureGraph"; +import { measureNodeQuery } from "#/hooks/graph/MeasureGraph"; export const measureRoutes = [ { path: "measures", Fallback: PageSkeleton, - loader: loaderFromQueryLoader(({ organizationId }) => - loadQuery(coreEnvironment, measuresQuery, { - organizationId: organizationId, - }), - ), - Component: withQueryRef( - lazy(() => import("#/pages/organizations/measures/MeasuresPage")), + Component: lazy( + () => + import("#/pages/organizations/measures/MeasuresPageLoader"), ), children: [ { diff --git a/pkg/coredata/measure.go b/pkg/coredata/measure.go index 1b97e5533..bc7360787 100644 --- a/pkg/coredata/measure.go +++ b/pkg/coredata/measure.go @@ -326,6 +326,41 @@ WHERE return count, nil } +func (m *Measures) LoadDistinctCategoriesByOrganizationID( + ctx context.Context, + conn pg.Conn, + scope Scoper, + organizationID gid.GID, +) ([]string, error) { + q := ` +SELECT DISTINCT + category +FROM + measures +WHERE + %s + AND organization_id = @organization_id +ORDER BY + category ASC +` + q = fmt.Sprintf(q, scope.SQLFragment()) + + args := pgx.NamedArgs{"organization_id": organizationID} + maps.Copy(args, scope.SQLArguments()) + + rows, err := conn.Query(ctx, q, args) + if err != nil { + return nil, fmt.Errorf("cannot query measure categories: %w", err) + } + + categories, err := pgx.CollectRows(rows, pgx.RowTo[string]) + if err != nil { + return nil, fmt.Errorf("cannot collect measure categories: %w", err) + } + + return categories, nil +} + func (m *Measures) LoadByOrganizationID( ctx context.Context, conn pg.Conn, diff --git a/pkg/coredata/measure_filter.go b/pkg/coredata/measure_filter.go index 1133e0e94..628700848 100644 --- a/pkg/coredata/measure_filter.go +++ b/pkg/coredata/measure_filter.go @@ -20,22 +20,25 @@ import ( type ( MeasureFilter struct { - query *string - state *MeasureState + query *string + state *MeasureState + category *string } ) -func NewMeasureFilter(query *string, state *MeasureState) *MeasureFilter { +func NewMeasureFilter(query *string, state *MeasureState, category *string) *MeasureFilter { return &MeasureFilter{ - query: query, - state: state, + query: query, + state: state, + category: category, } } func (f *MeasureFilter) SQLArguments() pgx.NamedArgs { return pgx.NamedArgs{ - "query": f.query, - "state": f.state, + "query": f.query, + "state": f.state, + "category": f.category, } } @@ -60,6 +63,15 @@ AND ELSE state = @state::mitigation_state END - ) +) +AND +( + CASE + WHEN @category::text IS NULL OR @category::text = '' THEN + TRUE + ELSE + category = @category::text + END +) ` } diff --git a/pkg/probo/framework_service.go b/pkg/probo/framework_service.go index a3e26b11e..599c03124 100644 --- a/pkg/probo/framework_service.go +++ b/pkg/probo/framework_service.go @@ -208,7 +208,7 @@ func (s FrameworkService) Export( Direction: page.OrderDirectionAsc, }, ), - coredata.NewMeasureFilter(nil, nil), + coredata.NewMeasureFilter(nil, nil, nil), ) if err != nil { return fmt.Errorf("cannot load measures: %w", err) diff --git a/pkg/probo/measure_service.go b/pkg/probo/measure_service.go index 1203ce9a9..18678526a 100644 --- a/pkg/probo/measure_service.go +++ b/pkg/probo/measure_service.go @@ -237,6 +237,43 @@ func (s MeasureService) CountForOrganizationID( return count, nil } +func (s MeasureService) ListDistinctCategoriesForOrganizationID( + ctx context.Context, + organizationID gid.GID, +) ([]string, error) { + var categories []string + + err := s.svc.pg.WithConn( + ctx, + func(conn pg.Conn) error { + organization := &coredata.Organization{} + if err := organization.LoadByID(ctx, conn, s.svc.scope, organizationID); err != nil { + return fmt.Errorf("cannot load organization: %w", err) + } + + var measures coredata.Measures + var err error + categories, err = measures.LoadDistinctCategoriesByOrganizationID( + ctx, + conn, + s.svc.scope, + organization.ID, + ) + if err != nil { + return fmt.Errorf("cannot load measure categories: %w", err) + } + + return nil + }, + ) + + if err != nil { + return nil, err + } + + return categories, nil +} + func (s MeasureService) ListForOrganizationID( ctx context.Context, organizationID gid.GID, diff --git a/pkg/server/api/console/v1/schema.graphql b/pkg/server/api/console/v1/schema.graphql index 70c80e230..8da9abe55 100644 --- a/pkg/server/api/console/v1/schema.graphql +++ b/pkg/server/api/console/v1/schema.graphql @@ -1572,6 +1572,7 @@ input DocumentFilter { input MeasureFilter { query: String state: MeasureState + category: String } input RiskFilter { @@ -1857,6 +1858,8 @@ type Organization implements Node { filter: StateOfApplicabilityFilter = { snapshotId: null } ): StateOfApplicabilityConnection! @goField(forceResolver: true) + measureCategories: [String!]! @goField(forceResolver: true) + measures( first: Int after: CursorKey diff --git a/pkg/server/api/console/v1/v1_resolver.go b/pkg/server/api/console/v1/v1_resolver.go index c6d9fd428..bcadf6005 100644 --- a/pkg/server/api/console/v1/v1_resolver.go +++ b/pkg/server/api/console/v1/v1_resolver.go @@ -530,9 +530,9 @@ func (r *controlResolver) Measures(ctx context.Context, obj *types.Control, firs cursor := types.NewCursor(first, after, last, before, pageOrderBy) - var measureFilter = coredata.NewMeasureFilter(nil, nil) + var measureFilter = coredata.NewMeasureFilter(nil, nil, nil) if filter != nil { - measureFilter = coredata.NewMeasureFilter(filter.Query, filter.State) + measureFilter = coredata.NewMeasureFilter(filter.Query, filter.State, filter.Category) } page, err := prb.Measures.ListForControlID(ctx, obj.ID, cursor, measureFilter) @@ -6554,6 +6554,23 @@ func (r *organizationResolver) StatesOfApplicability(ctx context.Context, obj *t return types.NewStateOfApplicabilityConnection(page, r, obj.ID, stateOfApplicabilityFilter), nil } +// MeasureCategories is the resolver for the measureCategories field. +func (r *organizationResolver) MeasureCategories(ctx context.Context, obj *types.Organization) ([]string, error) { + if err := r.authorize(ctx, obj.ID, probo.ActionMeasureList); err != nil { + return nil, err + } + + prb := r.ProboService(ctx, obj.ID.TenantID()) + + categories, err := prb.Measures.ListDistinctCategoriesForOrganizationID(ctx, obj.ID) + if err != nil { + r.logger.ErrorCtx(ctx, "cannot list measure categories", log.Error(err)) + return nil, gqlutils.Internal(ctx) + } + + return categories, nil +} + // Measures is the resolver for the measures field. func (r *organizationResolver) Measures(ctx context.Context, obj *types.Organization, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.MeasureOrderBy, filter *types.MeasureFilter) (*types.MeasureConnection, error) { if err := r.authorize(ctx, obj.ID, probo.ActionMeasureList); err != nil { @@ -6575,9 +6592,9 @@ func (r *organizationResolver) Measures(ctx context.Context, obj *types.Organiza cursor := types.NewCursor(first, after, last, before, pageOrderBy) - var measureFilter = coredata.NewMeasureFilter(nil, nil) + var measureFilter = coredata.NewMeasureFilter(nil, nil, nil) if filter != nil { - measureFilter = coredata.NewMeasureFilter(filter.Query, filter.State) + measureFilter = coredata.NewMeasureFilter(filter.Query, filter.State, filter.Category) } page, err := prb.Measures.ListForOrganizationID(ctx, obj.ID, cursor, measureFilter) @@ -7790,9 +7807,9 @@ func (r *riskResolver) Measures(ctx context.Context, obj *types.Risk, first *int cursor := types.NewCursor(first, after, last, before, pageOrderBy) - var measureFilter = coredata.NewMeasureFilter(nil, nil) + var measureFilter = coredata.NewMeasureFilter(nil, nil, nil) if filter != nil { - measureFilter = coredata.NewMeasureFilter(filter.Query, filter.State) + measureFilter = coredata.NewMeasureFilter(filter.Query, filter.State, filter.Category) } page, err := prb.Measures.ListForRiskID(ctx, obj.ID, cursor, measureFilter) diff --git a/pkg/server/api/mcp/v1/schema.resolvers.go b/pkg/server/api/mcp/v1/schema.resolvers.go index 9b892a6f0..c951920b3 100644 --- a/pkg/server/api/mcp/v1/schema.resolvers.go +++ b/pkg/server/api/mcp/v1/schema.resolvers.go @@ -380,9 +380,9 @@ func (r *Resolver) ListMeasuresTool(ctx context.Context, req *mcp.CallToolReques cursor := types.NewCursor(input.Size, input.Cursor, pageOrderBy) - var measureFilter = coredata.NewMeasureFilter(nil, nil) + var measureFilter = coredata.NewMeasureFilter(nil, nil, nil) if input.Filter != nil { - measureFilter = coredata.NewMeasureFilter(input.Filter.Query, input.Filter.State) + measureFilter = coredata.NewMeasureFilter(input.Filter.Query, input.Filter.State, input.Filter.Category) } page, err := prb.Measures.ListForOrganizationID(ctx, input.OrganizationID, cursor, measureFilter) @@ -1624,7 +1624,7 @@ func (r *Resolver) ListControlMeasuresTool(ctx context.Context, req *mcp.CallToo cursor := types.NewCursor(input.Size, input.Cursor, pageOrderBy) - measurePage, err := prb.Measures.ListForControlID(ctx, input.ControlID, cursor, coredata.NewMeasureFilter(nil, nil)) + measurePage, err := prb.Measures.ListForControlID(ctx, input.ControlID, cursor, coredata.NewMeasureFilter(nil, nil, nil)) if err != nil { return nil, types.ListControlMeasuresOutput{}, fmt.Errorf("failed to list control measures: %w", err) } diff --git a/pkg/server/api/mcp/v1/specification.yaml b/pkg/server/api/mcp/v1/specification.yaml index bc8baa223..1157cd1d7 100644 --- a/pkg/server/api/mcp/v1/specification.yaml +++ b/pkg/server/api/mcp/v1/specification.yaml @@ -1434,6 +1434,9 @@ components: state: $ref: "#/components/schemas/MeasureState" description: Measure state filter + category: + type: string + description: Filter by measure category ListMeasuresOutput: type: object