diff --git a/pkg/coredata/risk.go b/pkg/coredata/risk.go index 535f2f88e..d58d01ecd 100644 --- a/pkg/coredata/risk.go +++ b/pkg/coredata/risk.go @@ -68,6 +68,47 @@ func (r *Risk) CursorKey(orderBy RiskOrderField) page.CursorKey { panic(fmt.Sprintf("unsupported order by: %s", orderBy)) } +func (r *Risks) CountByMeasureID( + ctx context.Context, + conn pg.Conn, + scope Scoper, + measureID gid.GID, + filter *RiskFilter, +) (int, error) { + q := ` +WITH rsks AS ( + SELECT + r.id + FROM + risks r + INNER JOIN + risks_measures rm ON r.id = rm.risk_id + WHERE + rm.measure_id = @measure_id +) +SELECT + COUNT(id) +FROM + rsks +WHERE %s + AND %s +` + q = fmt.Sprintf(q, scope.SQLFragment(), filter.SQLFragment()) + + args := pgx.NamedArgs{"measure_id": measureID} + 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 (r *Risks) LoadByMeasureID( ctx context.Context, conn pg.Conn, @@ -148,6 +189,37 @@ WHERE %s return nil } +func (r *Risks) CountByOrganizationID( + ctx context.Context, + conn pg.Conn, + scope Scoper, + organizationID gid.GID, + filter *RiskFilter, +) (int, error) { + q := ` +SELECT + COUNT(id) +FROM risks +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 (r *Risks) LoadByOrganizationID( ctx context.Context, conn pg.Conn, diff --git a/pkg/probo/risk_service.go b/pkg/probo/risk_service.go index afd351645..b3d4f1e11 100644 --- a/pkg/probo/risk_service.go +++ b/pkg/probo/risk_service.go @@ -59,6 +59,33 @@ type ( } ) +func (s RiskService) CountForMeasureID( + ctx context.Context, + measureID gid.GID, + filter *coredata.RiskFilter, +) (int, error) { + var count int + + err := s.svc.pg.WithConn( + ctx, + func(conn pg.Conn) (err error) { + risks := &coredata.Risks{} + count, err = risks.CountByMeasureID(ctx, conn, s.svc.scope, measureID, filter) + if err != nil { + return fmt.Errorf("cannot count risks: %w", err) + } + + return nil + }, + ) + + if err != nil { + return 0, fmt.Errorf("cannot count risks: %w", err) + } + + return count, nil +} + func (s RiskService) ListForMeasureID( ctx context.Context, measureID gid.GID, @@ -81,6 +108,62 @@ func (s RiskService) ListForMeasureID( return page.NewPage(risks, cursor), nil } +func (s RiskService) CountForOrganizationID( + ctx context.Context, + organizationID gid.GID, + filter *coredata.RiskFilter, +) (int, error) { + var count int + + err := s.svc.pg.WithConn( + ctx, + func(conn pg.Conn) (err error) { + risks := &coredata.Risks{} + count, err = risks.CountByOrganizationID(ctx, conn, s.svc.scope, organizationID, filter) + if err != nil { + return fmt.Errorf("cannot count risks: %w", err) + } + + return nil + }, + ) + + if err != nil { + return 0, fmt.Errorf("cannot count risks: %w", err) + } + + return count, nil +} + +func (s RiskService) ListForOrganizationID( + ctx context.Context, + organizationID gid.GID, + cursor *page.Cursor[coredata.RiskOrderField], + filter *coredata.RiskFilter, +) (*page.Page[*coredata.Risk, coredata.RiskOrderField], error) { + var risks coredata.Risks + + err := s.svc.pg.WithConn( + ctx, + func(conn pg.Conn) error { + return risks.LoadByOrganizationID( + ctx, + conn, + s.svc.scope, + organizationID, + cursor, + filter, + ) + }, + ) + + if err != nil { + return nil, fmt.Errorf("cannot list risks: %w", err) + } + + return page.NewPage(risks, cursor), nil +} + func (s RiskService) CreateDocumentMapping( ctx context.Context, riskID gid.GID, @@ -390,32 +473,3 @@ func (s RiskService) Delete( }, ) } - -func (s RiskService) ListForOrganizationID( - ctx context.Context, - organizationID gid.GID, - cursor *page.Cursor[coredata.RiskOrderField], - filter *coredata.RiskFilter, -) (*page.Page[*coredata.Risk, coredata.RiskOrderField], error) { - var risks coredata.Risks - - err := s.svc.pg.WithConn( - ctx, - func(conn pg.Conn) error { - return risks.LoadByOrganizationID( - ctx, - conn, - s.svc.scope, - organizationID, - cursor, - filter, - ) - }, - ) - - if err != nil { - return nil, fmt.Errorf("cannot list risks: %w", err) - } - - return page.NewPage(risks, cursor), nil -} diff --git a/pkg/server/api/console/v1/schema.graphql b/pkg/server/api/console/v1/schema.graphql index 9d7f79371..a34aa9856 100644 --- a/pkg/server/api/console/v1/schema.graphql +++ b/pkg/server/api/console/v1/schema.graphql @@ -1162,7 +1162,11 @@ type DocumentEdge { node: Document! } -type RiskConnection { +type RiskConnection + @goModel( + model: "github.com/getprobo/probo/pkg/server/api/console/v1/types.RiskConnection" + ) { + totalCount: Int! @goField(forceResolver: true) edges: [RiskEdge!]! pageInfo: PageInfo! } diff --git a/pkg/server/api/console/v1/schema/schema.go b/pkg/server/api/console/v1/schema/schema.go index 4eabef5a5..287c86969 100644 --- a/pkg/server/api/console/v1/schema/schema.go +++ b/pkg/server/api/console/v1/schema/schema.go @@ -58,6 +58,7 @@ type ResolverRoot interface { Organization() OrganizationResolver Query() QueryResolver Risk() RiskResolver + RiskConnection() RiskConnectionResolver Task() TaskResolver User() UserResolver Vendor() VendorResolver @@ -652,8 +653,9 @@ type ComplexityRoot struct { } RiskConnection 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 } RiskEdge struct { @@ -1021,6 +1023,9 @@ type RiskResolver interface { Documents(ctx context.Context, obj *types.Risk, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.DocumentOrderBy, filter *types.DocumentFilter) (*types.DocumentConnection, error) Controls(ctx context.Context, obj *types.Risk, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.ControlOrderBy, filter *types.ControlFilter) (*types.ControlConnection, error) } +type RiskConnectionResolver interface { + TotalCount(ctx context.Context, obj *types.RiskConnection) (int, error) +} type TaskResolver interface { AssignedTo(ctx context.Context, obj *types.Task) (*types.People, error) Organization(ctx context.Context, obj *types.Task) (*types.Organization, error) @@ -3779,6 +3784,13 @@ func (e *executableSchema) Complexity(ctx context.Context, typeName, field strin return e.complexity.RiskConnection.PageInfo(childComplexity), true + case "RiskConnection.totalCount": + if e.complexity.RiskConnection.TotalCount == nil { + break + } + + return e.complexity.RiskConnection.TotalCount(childComplexity), true + case "RiskEdge.cursor": if e.complexity.RiskEdge.Cursor == nil { break @@ -5895,7 +5907,11 @@ type DocumentEdge { node: Document! } -type RiskConnection { +type RiskConnection + @goModel( + model: "github.com/getprobo/probo/pkg/server/api/console/v1/types.RiskConnection" + ) { + totalCount: Int! @goField(forceResolver: true) edges: [RiskEdge!]! pageInfo: PageInfo! } @@ -21015,6 +21031,8 @@ func (ec *executionContext) fieldContext_Measure_risks(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_RiskConnection_totalCount(ctx, field) case "edges": return ec.fieldContext_RiskConnection_edges(ctx, field) case "pageInfo": @@ -25847,6 +25865,8 @@ func (ec *executionContext) fieldContext_Organization_risks(ctx context.Context, IsResolver: true, Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { switch field.Name { + case "totalCount": + return ec.fieldContext_RiskConnection_totalCount(ctx, field) case "edges": return ec.fieldContext_RiskConnection_edges(ctx, field) case "pageInfo": @@ -28651,6 +28671,50 @@ func (ec *executionContext) fieldContext_Risk_updatedAt(_ context.Context, field return fc, nil } +func (ec *executionContext) _RiskConnection_totalCount(ctx context.Context, field graphql.CollectedField, obj *types.RiskConnection) (ret graphql.Marshaler) { + fc, err := ec.fieldContext_RiskConnection_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.RiskConnection().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_RiskConnection_totalCount(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { + fc = &graphql.FieldContext{ + Object: "RiskConnection", + 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) _RiskConnection_edges(ctx context.Context, field graphql.CollectedField, obj *types.RiskConnection) (ret graphql.Marshaler) { fc, err := ec.fieldContext_RiskConnection_edges(ctx, field) if err != nil { @@ -28727,9 +28791,9 @@ func (ec *executionContext) _RiskConnection_pageInfo(ctx context.Context, field } 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_RiskConnection_pageInfo(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { @@ -46394,15 +46458,51 @@ func (ec *executionContext) _RiskConnection(ctx context.Context, sel ast.Selecti switch field.Name { case "__typename": out.Values[i] = graphql.MarshalString("RiskConnection") + 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._RiskConnection_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._RiskConnection_edges(ctx, field, obj) if out.Values[i] == graphql.Null { - out.Invalids++ + atomic.AddUint32(&out.Invalids, 1) } case "pageInfo": out.Values[i] = ec._RiskConnection_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/risk.go b/pkg/server/api/console/v1/types/risk.go index fd403068b..784bb8fba 100644 --- a/pkg/server/api/console/v1/types/risk.go +++ b/pkg/server/api/console/v1/types/risk.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 ( RiskOrderBy OrderBy[coredata.RiskOrderField] + + RiskConnection struct { + TotalCount int + Edges []*RiskEdge + PageInfo PageInfo + + Resolver any + ParentID gid.GID + Filters *coredata.RiskFilter + } ) -func NewRiskConnection(p *page.Page[*coredata.Risk, coredata.RiskOrderField]) *RiskConnection { +func NewRiskConnection( + p *page.Page[*coredata.Risk, coredata.RiskOrderField], + parentType any, + parentID gid.GID, + filters *coredata.RiskFilter, +) *RiskConnection { var edges = make([]*RiskEdge, len(p.Data)) for i := range edges { @@ -32,7 +48,11 @@ func NewRiskConnection(p *page.Page[*coredata.Risk, coredata.RiskOrderField]) *R return &RiskConnection{ 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 198aa716c..4d5f4980d 100644 --- a/pkg/server/api/console/v1/types/types.go +++ b/pkg/server/api/console/v1/types/types.go @@ -879,11 +879,6 @@ type Risk struct { func (Risk) IsNode() {} func (this Risk) GetID() gid.GID { return this.ID } -type RiskConnection struct { - Edges []*RiskEdge `json:"edges"` - PageInfo *PageInfo `json:"pageInfo"` -} - type RiskEdge struct { Cursor page.CursorKey `json:"cursor"` Node *Risk `json:"node"` diff --git a/pkg/server/api/console/v1/v1_resolver.go b/pkg/server/api/console/v1/v1_resolver.go index 56337e032..97cbe1044 100644 --- a/pkg/server/api/console/v1/v1_resolver.go +++ b/pkg/server/api/console/v1/v1_resolver.go @@ -677,7 +677,7 @@ func (r *measureResolver) Risks(ctx context.Context, obj *types.Measure, first * return nil, fmt.Errorf("cannot list measure risks: %w", err) } - return types.NewRiskConnection(page), nil + return types.NewRiskConnection(page, r, obj.ID, riskFilter), nil } // Controls is the resolver for the controls field. @@ -2240,7 +2240,7 @@ func (r *organizationResolver) Risks(ctx context.Context, obj *types.Organizatio panic(fmt.Errorf("cannot list organization risks: %w", err)) } - return types.NewRiskConnection(page), nil + return types.NewRiskConnection(page, r, obj.ID, riskFilter), nil } // Tasks is the resolver for the tasks field. @@ -2566,6 +2566,28 @@ func (r *riskResolver) Controls(ctx context.Context, obj *types.Risk, first *int return types.NewControlConnection(page, r, obj.ID, controlFilter), nil } +// TotalCount is the resolver for the totalCount field. +func (r *riskConnectionResolver) TotalCount(ctx context.Context, obj *types.RiskConnection) (int, error) { + svc := GetTenantService(ctx, r.proboSvc, obj.ParentID.TenantID()) + + switch obj.Resolver.(type) { + case *measureResolver: + count, err := svc.Risks.CountForMeasureID(ctx, obj.ParentID, obj.Filters) + if err != nil { + return 0, fmt.Errorf("cannot count risks: %w", err) + } + return count, nil + case *organizationResolver: + count, err := svc.Risks.CountForOrganizationID(ctx, obj.ParentID, obj.Filters) + if err != nil { + return 0, fmt.Errorf("cannot count risks: %w", err) + } + return count, nil + } + + panic(fmt.Errorf("unsupported resolver: %T", obj.Resolver)) +} + // AssignedTo is the resolver for the assignedTo field. func (r *taskResolver) AssignedTo(ctx context.Context, obj *types.Task) (*types.People, error) { svc := GetTenantService(ctx, r.proboSvc, obj.ID.TenantID()) @@ -2901,6 +2923,9 @@ func (r *Resolver) Query() schema.QueryResolver { return &queryResolver{r} } // Risk returns schema.RiskResolver implementation. func (r *Resolver) Risk() schema.RiskResolver { return &riskResolver{r} } +// RiskConnection returns schema.RiskConnectionResolver implementation. +func (r *Resolver) RiskConnection() schema.RiskConnectionResolver { return &riskConnectionResolver{r} } + // Task returns schema.TaskResolver implementation. func (r *Resolver) Task() schema.TaskResolver { return &taskResolver{r} } @@ -2939,6 +2964,7 @@ type mutationResolver struct{ *Resolver } type organizationResolver struct{ *Resolver } type queryResolver struct{ *Resolver } type riskResolver struct{ *Resolver } +type riskConnectionResolver struct{ *Resolver } type taskResolver struct{ *Resolver } type userResolver struct{ *Resolver } type vendorResolver struct{ *Resolver }