diff --git a/pkg/coredata/measure.go b/pkg/coredata/measure.go index 6d3d029b9..bd477b889 100644 --- a/pkg/coredata/measure.go +++ b/pkg/coredata/measure.go @@ -52,6 +52,47 @@ func (m Measure) CursorKey(orderBy MeasureOrderField) page.CursorKey { panic(fmt.Sprintf("unsupported order by: %s", orderBy)) } +func (m *Measures) CountByRiskID( + ctx context.Context, + conn pg.Conn, + scope Scoper, + riskID gid.GID, + filter *MeasureFilter, +) (int, error) { + q := ` +WITH msrs AS ( + SELECT + m.id + FROM + measures m + INNER JOIN + risks_measures rm ON m.id = rm.measure_id + WHERE + rm.risk_id = @risk_id +) +SELECT + COUNT(id) +FROM + msrs +WHERE %s + AND %s +` + q = fmt.Sprintf(q, scope.SQLFragment(), filter.SQLFragment()) + + args := pgx.NamedArgs{"risk_id": riskID} + maps.Copy(args, scope.SQLArguments()) + maps.Copy(args, filter.SQLArguments()) + + row := conn.QueryRow(ctx, q, args) + + var count int + if err := row.Scan(&count); err != nil { + return 0, fmt.Errorf("cannot scan count: %w", err) + } + + return count, nil +} + func (m *Measures) LoadByRiskID( ctx context.Context, conn pg.Conn, @@ -118,6 +159,47 @@ WHERE %s return nil } +func (m *Measures) CountByControlID( + ctx context.Context, + conn pg.Conn, + scope Scoper, + controlID gid.GID, + filter *MeasureFilter, +) (int, error) { + q := ` +WITH mtgtns AS ( + SELECT + m.id + FROM + measures m + INNER JOIN + controls_measures cm ON m.id = cm.measure_id + WHERE + cm.control_id = @control_id + ) + SELECT + COUNT(id) + FROM + mtgtns + WHERE %s + AND %s + ` + q = fmt.Sprintf(q, scope.SQLFragment(), filter.SQLFragment()) + + args := pgx.NamedArgs{"control_id": controlID} + maps.Copy(args, scope.SQLArguments()) + maps.Copy(args, filter.SQLArguments()) + + row := conn.QueryRow(ctx, q, args) + + var count int + if err := row.Scan(&count); err != nil { + return 0, fmt.Errorf("cannot scan count: %w", err) + } + + return count, nil +} + func (m *Measures) LoadByControlID( ctx context.Context, conn pg.Conn, @@ -184,6 +266,39 @@ WHERE %s return nil } +func (m *Measures) CountByOrganizationID( + ctx context.Context, + conn pg.Conn, + scope Scoper, + organizationID gid.GID, + filter *MeasureFilter, +) (int, error) { + q := ` +SELECT + COUNT(id) +FROM + measures +WHERE + %s + AND organization_id = @organization_id + AND %s +` + q = fmt.Sprintf(q, scope.SQLFragment(), filter.SQLFragment()) + + args := pgx.NamedArgs{"organization_id": organizationID} + maps.Copy(args, scope.SQLArguments()) + maps.Copy(args, filter.SQLArguments()) + + row := conn.QueryRow(ctx, q, args) + + var count int + if err := row.Scan(&count); err != nil { + return 0, fmt.Errorf("cannot scan count: %w", err) + } + + return count, nil +} + func (m *Measures) LoadByOrganizationID( ctx context.Context, conn pg.Conn, diff --git a/pkg/probo/measure_service.go b/pkg/probo/measure_service.go index c0af5d08a..a943a2f53 100644 --- a/pkg/probo/measure_service.go +++ b/pkg/probo/measure_service.go @@ -69,6 +69,32 @@ type ( } ) +func (s MeasureService) CountForRiskID( + ctx context.Context, + riskID gid.GID, + filter *coredata.MeasureFilter, +) (int, error) { + var count int + + err := s.svc.pg.WithConn( + ctx, + func(conn pg.Conn) (err error) { + measures := &coredata.Measures{} + count, err = measures.CountByRiskID(ctx, conn, s.svc.scope, riskID, filter) + if err != nil { + return fmt.Errorf("cannot count measures: %w", err) + } + + return nil + }, + ) + + if err != nil { + return 0, err + } + + return count, nil +} func (s MeasureService) ListForRiskID( ctx context.Context, riskID gid.GID, @@ -101,6 +127,33 @@ func (s MeasureService) ListForRiskID( return page.NewPage(measures, cursor), nil } +func (s MeasureService) CountForControlID( + ctx context.Context, + controlID gid.GID, + filter *coredata.MeasureFilter, +) (int, error) { + var count int + + err := s.svc.pg.WithConn( + ctx, + func(conn pg.Conn) (err error) { + measures := &coredata.Measures{} + count, err = measures.CountByControlID(ctx, conn, s.svc.scope, controlID, filter) + if err != nil { + return fmt.Errorf("cannot count measures: %w", err) + } + + return nil + }, + ) + + if err != nil { + return 0, err + } + + return count, nil +} + func (s MeasureService) ListForControlID( ctx context.Context, controlID gid.GID, @@ -133,6 +186,72 @@ func (s MeasureService) ListForControlID( return page.NewPage(measures, cursor), nil } +func (s MeasureService) CountForOrganizationID( + ctx context.Context, + organizationID gid.GID, + filter *coredata.MeasureFilter, +) (int, error) { + var count int + + err := s.svc.pg.WithConn( + ctx, + func(conn pg.Conn) (err error) { + measures := &coredata.Measures{} + count, err = measures.CountByOrganizationID(ctx, conn, s.svc.scope, organizationID, filter) + if err != nil { + return fmt.Errorf("cannot count measures: %w", err) + } + + return nil + }, + ) + + if err != nil { + return 0, err + } + + return count, nil +} + +func (s MeasureService) ListForOrganizationID( + ctx context.Context, + organizationID gid.GID, + cursor *page.Cursor[coredata.MeasureOrderField], + filter *coredata.MeasureFilter, +) (*page.Page[*coredata.Measure, coredata.MeasureOrderField], error) { + var measures coredata.Measures + organization := &coredata.Organization{} + + err := s.svc.pg.WithConn( + ctx, + func(conn pg.Conn) error { + if err := organization.LoadByID(ctx, conn, s.svc.scope, organizationID); err != nil { + return fmt.Errorf("cannot load organization: %w", err) + } + + err := measures.LoadByOrganizationID( + ctx, + conn, + s.svc.scope, + organization.ID, + cursor, + filter, + ) + if err != nil { + return fmt.Errorf("cannot load measures: %w", err) + } + + return nil + }, + ) + + if err != nil { + return nil, err + } + + return page.NewPage(measures, cursor), nil +} + func (s MeasureService) Get( ctx context.Context, measureID gid.GID, @@ -319,45 +438,6 @@ func (s MeasureService) Update( return measure, nil } -func (s MeasureService) ListForOrganizationID( - ctx context.Context, - organizationID gid.GID, - cursor *page.Cursor[coredata.MeasureOrderField], - filter *coredata.MeasureFilter, -) (*page.Page[*coredata.Measure, coredata.MeasureOrderField], error) { - var measures coredata.Measures - organization := &coredata.Organization{} - - err := s.svc.pg.WithConn( - ctx, - func(conn pg.Conn) error { - if err := organization.LoadByID(ctx, conn, s.svc.scope, organizationID); err != nil { - return fmt.Errorf("cannot load organization: %w", err) - } - - err := measures.LoadByOrganizationID( - ctx, - conn, - s.svc.scope, - organization.ID, - cursor, - filter, - ) - if err != nil { - return fmt.Errorf("cannot load measures: %w", err) - } - - return nil - }, - ) - - if err != nil { - return nil, err - } - - return page.NewPage(measures, cursor), nil -} - func (s MeasureService) Create( ctx context.Context, req CreateMeasureRequest, diff --git a/pkg/server/api/console/v1/schema.graphql b/pkg/server/api/console/v1/schema.graphql index 0a60591e7..9d7f79371 100644 --- a/pkg/server/api/console/v1/schema.graphql +++ b/pkg/server/api/console/v1/schema.graphql @@ -1118,7 +1118,11 @@ type ControlEdge { node: Control! } -type MeasureConnection { +type MeasureConnection + @goModel( + model: "github.com/getprobo/probo/pkg/server/api/console/v1/types.MeasureConnection" + ) { + totalCount: Int! @goField(forceResolver: true) edges: [MeasureEdge!]! pageInfo: PageInfo! } diff --git a/pkg/server/api/console/v1/schema/schema.go b/pkg/server/api/console/v1/schema/schema.go index 96f993cb1..4eabef5a5 100644 --- a/pkg/server/api/console/v1/schema/schema.go +++ b/pkg/server/api/console/v1/schema/schema.go @@ -53,6 +53,7 @@ type ResolverRoot interface { Framework() FrameworkResolver FrameworkConnection() FrameworkConnectionResolver Measure() MeasureResolver + MeasureConnection() MeasureConnectionResolver Mutation() MutationResolver Organization() OrganizationResolver Query() QueryResolver @@ -470,8 +471,9 @@ type ComplexityRoot struct { } MeasureConnection struct { - Edges func(childComplexity int) int - PageInfo func(childComplexity int) int + Edges func(childComplexity int) int + PageInfo func(childComplexity int) int + TotalCount func(childComplexity int) int } MeasureEdge struct { @@ -925,6 +927,9 @@ type MeasureResolver interface { Risks(ctx context.Context, obj *types.Measure, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.RiskOrderBy, filter *types.RiskFilter) (*types.RiskConnection, error) Controls(ctx context.Context, obj *types.Measure, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.ControlOrderBy, filter *types.ControlFilter) (*types.ControlConnection, error) } +type MeasureConnectionResolver interface { + TotalCount(ctx context.Context, obj *types.MeasureConnection) (int, error) +} type MutationResolver interface { CreateOrganization(ctx context.Context, input types.CreateOrganizationInput) (*types.CreateOrganizationPayload, error) UpdateOrganization(ctx context.Context, input types.UpdateOrganizationInput) (*types.UpdateOrganizationPayload, error) @@ -2448,6 +2453,13 @@ func (e *executableSchema) Complexity(ctx context.Context, typeName, field strin return e.complexity.MeasureConnection.PageInfo(childComplexity), true + case "MeasureConnection.totalCount": + if e.complexity.MeasureConnection.TotalCount == nil { + break + } + + return e.complexity.MeasureConnection.TotalCount(childComplexity), true + case "MeasureEdge.cursor": if e.complexity.MeasureEdge.Cursor == nil { break @@ -5839,7 +5851,11 @@ type ControlEdge { node: Control! } -type MeasureConnection { +type MeasureConnection + @goModel( + model: "github.com/getprobo/probo/pkg/server/api/console/v1/types.MeasureConnection" + ) { + totalCount: Int! @goField(forceResolver: true) edges: [MeasureEdge!]! pageInfo: PageInfo! } @@ -13453,6 +13469,8 @@ func (ec *executionContext) fieldContext_Control_measures(ctx context.Context, f IsResolver: true, Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { switch field.Name { + case "totalCount": + return ec.fieldContext_MeasureConnection_totalCount(ctx, field) case "edges": return ec.fieldContext_MeasureConnection_edges(ctx, field) case "pageInfo": @@ -21170,6 +21188,50 @@ func (ec *executionContext) fieldContext_Measure_updatedAt(_ context.Context, fi return fc, nil } +func (ec *executionContext) _MeasureConnection_totalCount(ctx context.Context, field graphql.CollectedField, obj *types.MeasureConnection) (ret graphql.Marshaler) { + fc, err := ec.fieldContext_MeasureConnection_totalCount(ctx, field) + if err != nil { + return graphql.Null + } + ctx = graphql.WithFieldContext(ctx, fc) + defer func() { + if r := recover(); r != nil { + ec.Error(ctx, ec.Recover(ctx, r)) + ret = graphql.Null + } + }() + resTmp, err := ec.ResolverMiddleware(ctx, func(rctx context.Context) (any, error) { + ctx = rctx // use context from middleware stack in children + return ec.resolvers.MeasureConnection().TotalCount(rctx, obj) + }) + if err != nil { + ec.Error(ctx, err) + return graphql.Null + } + if resTmp == nil { + if !graphql.HasFieldError(ctx, fc) { + ec.Errorf(ctx, "must not be null") + } + return graphql.Null + } + res := resTmp.(int) + fc.Result = res + return ec.marshalNInt2int(ctx, field.Selections, res) +} + +func (ec *executionContext) fieldContext_MeasureConnection_totalCount(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { + fc = &graphql.FieldContext{ + Object: "MeasureConnection", + Field: field, + IsMethod: true, + IsResolver: true, + Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { + return nil, errors.New("field of type Int does not have child fields") + }, + } + return fc, nil +} + func (ec *executionContext) _MeasureConnection_edges(ctx context.Context, field graphql.CollectedField, obj *types.MeasureConnection) (ret graphql.Marshaler) { fc, err := ec.fieldContext_MeasureConnection_edges(ctx, field) if err != nil { @@ -21246,9 +21308,9 @@ func (ec *executionContext) _MeasureConnection_pageInfo(ctx context.Context, fie } return graphql.Null } - res := resTmp.(*types.PageInfo) + res := resTmp.(types.PageInfo) fc.Result = res - return ec.marshalNPageInfo2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐPageInfo(ctx, field.Selections, res) + return ec.marshalNPageInfo2githubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐPageInfo(ctx, field.Selections, res) } func (ec *executionContext) fieldContext_MeasureConnection_pageInfo(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { @@ -25722,6 +25784,8 @@ func (ec *executionContext) fieldContext_Organization_measures(ctx context.Conte IsResolver: true, Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { switch field.Name { + case "totalCount": + return ec.fieldContext_MeasureConnection_totalCount(ctx, field) case "edges": return ec.fieldContext_MeasureConnection_edges(ctx, field) case "pageInfo": @@ -28351,6 +28415,8 @@ func (ec *executionContext) fieldContext_Risk_measures(ctx context.Context, fiel IsResolver: true, Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { switch field.Name { + case "totalCount": + return ec.fieldContext_MeasureConnection_totalCount(ctx, field) case "edges": return ec.fieldContext_MeasureConnection_edges(ctx, field) case "pageInfo": @@ -44367,15 +44433,51 @@ func (ec *executionContext) _MeasureConnection(ctx context.Context, sel ast.Sele switch field.Name { case "__typename": out.Values[i] = graphql.MarshalString("MeasureConnection") + case "totalCount": + field := field + + innerFunc := func(ctx context.Context, fs *graphql.FieldSet) (res graphql.Marshaler) { + defer func() { + if r := recover(); r != nil { + ec.Error(ctx, ec.Recover(ctx, r)) + } + }() + res = ec._MeasureConnection_totalCount(ctx, field, obj) + if res == graphql.Null { + atomic.AddUint32(&fs.Invalids, 1) + } + return res + } + + if field.Deferrable != nil { + dfs, ok := deferred[field.Deferrable.Label] + di := 0 + if ok { + dfs.AddField(field) + di = len(dfs.Values) - 1 + } else { + dfs = graphql.NewFieldSet([]graphql.CollectedField{field}) + deferred[field.Deferrable.Label] = dfs + } + dfs.Concurrently(di, func(ctx context.Context) graphql.Marshaler { + return innerFunc(ctx, dfs) + }) + + // don't run the out.Concurrently() call below + out.Values[i] = graphql.Null + continue + } + + out.Concurrently(i, func(ctx context.Context) graphql.Marshaler { return innerFunc(ctx, out) }) case "edges": out.Values[i] = ec._MeasureConnection_edges(ctx, field, obj) if out.Values[i] == graphql.Null { - out.Invalids++ + atomic.AddUint32(&out.Invalids, 1) } case "pageInfo": out.Values[i] = ec._MeasureConnection_pageInfo(ctx, field, obj) if out.Values[i] == graphql.Null { - out.Invalids++ + atomic.AddUint32(&out.Invalids, 1) } default: panic("unknown field " + strconv.Quote(field.Name)) diff --git a/pkg/server/api/console/v1/types/mesure.go b/pkg/server/api/console/v1/types/mesure.go index 19454a688..121fdba00 100644 --- a/pkg/server/api/console/v1/types/mesure.go +++ b/pkg/server/api/console/v1/types/mesure.go @@ -16,14 +16,30 @@ package types import ( "github.com/getprobo/probo/pkg/coredata" + "github.com/getprobo/probo/pkg/gid" "github.com/getprobo/probo/pkg/page" ) type ( MeasureOrderBy OrderBy[coredata.MeasureOrderField] + + MeasureConnection struct { + TotalCount int + Edges []*MeasureEdge + PageInfo PageInfo + + Resolver any + ParentID gid.GID + Filters *coredata.MeasureFilter + } ) -func NewMeasureConnection(p *page.Page[*coredata.Measure, coredata.MeasureOrderField]) *MeasureConnection { +func NewMeasureConnection( + p *page.Page[*coredata.Measure, coredata.MeasureOrderField], + parentType any, + parentID gid.GID, + filters *coredata.MeasureFilter, +) *MeasureConnection { var edges = make([]*MeasureEdge, len(p.Data)) for i := range edges { @@ -32,7 +48,11 @@ func NewMeasureConnection(p *page.Page[*coredata.Measure, coredata.MeasureOrderF return &MeasureConnection{ Edges: edges, - PageInfo: NewPageInfo(p), + PageInfo: *NewPageInfo(p), + + Resolver: parentType, + ParentID: parentID, + Filters: filters, } } diff --git a/pkg/server/api/console/v1/types/types.go b/pkg/server/api/console/v1/types/types.go index 7fc9f91c8..198aa716c 100644 --- a/pkg/server/api/console/v1/types/types.go +++ b/pkg/server/api/console/v1/types/types.go @@ -729,11 +729,6 @@ type Measure struct { func (Measure) IsNode() {} func (this Measure) GetID() gid.GID { return this.ID } -type MeasureConnection struct { - Edges []*MeasureEdge `json:"edges"` - PageInfo *PageInfo `json:"pageInfo"` -} - type MeasureEdge struct { Cursor page.CursorKey `json:"cursor"` Node *Measure `json:"node"` diff --git a/pkg/server/api/console/v1/v1_resolver.go b/pkg/server/api/console/v1/v1_resolver.go index 6cf84893e..56337e032 100644 --- a/pkg/server/api/console/v1/v1_resolver.go +++ b/pkg/server/api/console/v1/v1_resolver.go @@ -134,7 +134,7 @@ func (r *controlResolver) Measures(ctx context.Context, obj *types.Control, firs return nil, fmt.Errorf("cannot list measures: %w", err) } - return types.NewMeasureConnection(page), nil + return types.NewMeasureConnection(page, r, obj.ID, measureFilter), nil } // Documents is the resolver for the documents field. @@ -710,6 +710,34 @@ func (r *measureResolver) Controls(ctx context.Context, obj *types.Measure, firs return types.NewControlConnection(page, r, obj.ID, controlFilter), nil } +// TotalCount is the resolver for the totalCount field. +func (r *measureConnectionResolver) TotalCount(ctx context.Context, obj *types.MeasureConnection) (int, error) { + svc := GetTenantService(ctx, r.proboSvc, obj.ParentID.TenantID()) + + switch obj.Resolver.(type) { + case *organizationResolver: + count, err := svc.Measures.CountForOrganizationID(ctx, obj.ParentID, obj.Filters) + if err != nil { + return 0, fmt.Errorf("cannot count measures: %w", err) + } + return count, nil + case *controlResolver: + count, err := svc.Measures.CountForControlID(ctx, obj.ParentID, obj.Filters) + if err != nil { + return 0, fmt.Errorf("cannot count measures: %w", err) + } + return count, nil + case *riskResolver: + count, err := svc.Measures.CountForRiskID(ctx, obj.ParentID, obj.Filters) + if err != nil { + return 0, fmt.Errorf("cannot count measures: %w", err) + } + return count, nil + } + + panic(fmt.Errorf("unsupported resolver: %T", obj.Resolver)) +} + // CreateOrganization is the resolver for the createOrganization field. func (r *mutationResolver) CreateOrganization(ctx context.Context, input types.CreateOrganizationInput) (*types.CreateOrganizationPayload, error) { svc := r.proboSvc.WithTenant(gid.NewTenantID()) @@ -2182,7 +2210,7 @@ func (r *organizationResolver) Measures(ctx context.Context, obj *types.Organiza panic(fmt.Errorf("cannot list organization measures: %w", err)) } - return types.NewMeasureConnection(page), nil + return types.NewMeasureConnection(page, r, obj.ID, measureFilter), nil } // Risks is the resolver for the risks field. @@ -2475,7 +2503,7 @@ func (r *riskResolver) Measures(ctx context.Context, obj *types.Risk, first *int panic(fmt.Errorf("cannot list risk measures: %w", err)) } - return types.NewMeasureConnection(page), nil + return types.NewMeasureConnection(page, r, obj.ID, measureFilter), nil } // Documents is the resolver for the documents field. @@ -2856,6 +2884,11 @@ func (r *Resolver) FrameworkConnection() schema.FrameworkConnectionResolver { // Measure returns schema.MeasureResolver implementation. func (r *Resolver) Measure() schema.MeasureResolver { return &measureResolver{r} } +// MeasureConnection returns schema.MeasureConnectionResolver implementation. +func (r *Resolver) MeasureConnection() schema.MeasureConnectionResolver { + return &measureConnectionResolver{r} +} + // Mutation returns schema.MutationResolver implementation. func (r *Resolver) Mutation() schema.MutationResolver { return &mutationResolver{r} } @@ -2901,6 +2934,7 @@ type evidenceResolver struct{ *Resolver } type frameworkResolver struct{ *Resolver } type frameworkConnectionResolver struct{ *Resolver } type measureResolver struct{ *Resolver } +type measureConnectionResolver struct{ *Resolver } type mutationResolver struct{ *Resolver } type organizationResolver struct{ *Resolver } type queryResolver struct{ *Resolver }