From f7ee529b3b10c90c91ac48216e300c95707992a1 Mon Sep 17 00:00:00 2001 From: Bryan Frimin Date: Mon, 9 Jun 2025 21:42:30 -0700 Subject: [PATCH] Add evidences totalCount support Signed-off-by: Bryan Frimin --- pkg/coredata/evidence.go | 106 ++++++++++++++---- pkg/probo/evidence_service.go | 52 +++++++++ pkg/server/api/console/v1/schema.graphql | 6 +- pkg/server/api/console/v1/schema/schema.go | 114 ++++++++++++++++++-- pkg/server/api/console/v1/types/evidence.go | 21 +++- pkg/server/api/console/v1/types/types.go | 5 - pkg/server/api/console/v1/v1_resolver.go | 32 +++++- 7 files changed, 298 insertions(+), 38 deletions(-) diff --git a/pkg/coredata/evidence.go b/pkg/coredata/evidence.go index 794f3eb0c..792112702 100644 --- a/pkg/coredata/evidence.go +++ b/pkg/coredata/evidence.go @@ -239,6 +239,38 @@ LIMIT 1; return nil } +func (e *Evidences) CountByMeasureID( + ctx context.Context, + conn pg.Conn, + scope Scoper, + measureID gid.GID, +) (int, error) { + q := ` +SELECT + COUNT(id) +FROM + evidences +WHERE + %s + AND measure_id = @measure_id + ` + + q = fmt.Sprintf(q, scope.SQLFragment()) + + args := pgx.StrictNamedArgs{"measure_id": measureID} + maps.Copy(args, scope.SQLArguments()) + + row := conn.QueryRow(ctx, q, args) + + var count int + err := row.Scan(&count) + if err != nil { + return 0, fmt.Errorf("cannot collect evidence: %w", err) + } + + return count, nil +} + func (e *Evidences) LoadByMeasureID( ctx context.Context, conn pg.Conn, @@ -247,27 +279,27 @@ func (e *Evidences) LoadByMeasureID( cursor *page.Cursor[EvidenceOrderField], ) error { q := ` - SELECT - id, - measure_id, - task_id, - reference_id, - state, - type, - object_key, - mime_type, - size, - filename, - url, - description, - created_at, - updated_at - FROM - evidences - WHERE - %s - AND measure_id = @measure_id - AND %s +SELECT + id, + measure_id, + task_id, + reference_id, + state, + type, + object_key, + mime_type, + size, + filename, + url, + description, + created_at, + updated_at +FROM + evidences +WHERE + %s + AND measure_id = @measure_id + AND %s ` q = fmt.Sprintf(q, scope.SQLFragment(), cursor.SQLFragment()) @@ -291,6 +323,38 @@ func (e *Evidences) LoadByMeasureID( return nil } +func (e *Evidences) CountByTaskID( + ctx context.Context, + conn pg.Conn, + scope Scoper, + taskID gid.GID, +) (int, error) { + q := ` +SELECT + COUNT(id) +FROM + evidences +WHERE + %s + AND task_id = @task_id +` + + q = fmt.Sprintf(q, scope.SQLFragment()) + + args := pgx.StrictNamedArgs{"task_id": taskID} + maps.Copy(args, scope.SQLArguments()) + + row := conn.QueryRow(ctx, q, args) + + var count int + err := row.Scan(&count) + if err != nil { + return 0, fmt.Errorf("cannot collect evidence: %w", err) + } + + return count, nil +} + func (e *Evidences) LoadByTaskID( ctx context.Context, conn pg.Conn, diff --git a/pkg/probo/evidence_service.go b/pkg/probo/evidence_service.go index 36b715490..9bc121dae 100644 --- a/pkg/probo/evidence_service.go +++ b/pkg/probo/evidence_service.go @@ -441,6 +441,32 @@ func (s EvidenceService) GenerateFileURL( return &presignedReq.URL, nil } +func (s EvidenceService) CountForMeasureID( + ctx context.Context, + measureID gid.GID, +) (int, error) { + var count int + + err := s.svc.pg.WithConn( + ctx, + func(conn pg.Conn) (err error) { + evidences := coredata.Evidences{} + count, err = evidences.CountByMeasureID(ctx, conn, s.svc.scope, measureID) + if err != nil { + return fmt.Errorf("cannot count evidences: %w", err) + } + + return nil + }, + ) + + if err != nil { + return 0, err + } + + return count, nil +} + func (s EvidenceService) ListForMeasureID( ctx context.Context, measureID gid.GID, @@ -468,6 +494,32 @@ func (s EvidenceService) ListForMeasureID( return page.NewPage(evidences, cursor), nil } +func (s EvidenceService) CountForTaskID( + ctx context.Context, + taskID gid.GID, +) (int, error) { + var count int + + err := s.svc.pg.WithConn( + ctx, + func(conn pg.Conn) (err error) { + evidences := coredata.Evidences{} + count, err = evidences.CountByTaskID(ctx, conn, s.svc.scope, taskID) + if err != nil { + return fmt.Errorf("cannot count evidences: %w", err) + } + + return nil + }, + ) + + if err != nil { + return 0, err + } + + return count, nil +} + func (s EvidenceService) ListForTaskID( ctx context.Context, taskID gid.GID, diff --git a/pkg/server/api/console/v1/schema.graphql b/pkg/server/api/console/v1/schema.graphql index eb9eaef63..a83981b98 100644 --- a/pkg/server/api/console/v1/schema.graphql +++ b/pkg/server/api/console/v1/schema.graphql @@ -1146,7 +1146,11 @@ type TaskEdge { node: Task! } -type EvidenceConnection { +type EvidenceConnection + @goModel( + model: "github.com/getprobo/probo/pkg/server/api/console/v1/types.EvidenceConnection" + ) { + totalCount: Int! @goField(forceResolver: true) edges: [EvidenceEdge!]! pageInfo: PageInfo! } diff --git a/pkg/server/api/console/v1/schema/schema.go b/pkg/server/api/console/v1/schema/schema.go index 6d8c8f983..84319d707 100644 --- a/pkg/server/api/console/v1/schema/schema.go +++ b/pkg/server/api/console/v1/schema/schema.go @@ -51,6 +51,7 @@ type ResolverRoot interface { DocumentVersion() DocumentVersionResolver DocumentVersionSignature() DocumentVersionSignatureResolver Evidence() EvidenceResolver + EvidenceConnection() EvidenceConnectionResolver Framework() FrameworkResolver FrameworkConnection() FrameworkConnectionResolver Measure() MeasureResolver @@ -406,8 +407,9 @@ type ComplexityRoot struct { } EvidenceConnection 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 } EvidenceEdge struct { @@ -923,6 +925,9 @@ type EvidenceResolver interface { Task(ctx context.Context, obj *types.Evidence) (*types.Task, error) Measure(ctx context.Context, obj *types.Evidence) (*types.Measure, error) } +type EvidenceConnectionResolver interface { + TotalCount(ctx context.Context, obj *types.EvidenceConnection) (int, error) +} type FrameworkResolver interface { Organization(ctx context.Context, obj *types.Framework) (*types.Organization, error) Controls(ctx context.Context, obj *types.Framework, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.ControlOrderBy, filter *types.ControlFilter) (*types.ControlConnection, error) @@ -2219,6 +2224,13 @@ func (e *executableSchema) Complexity(ctx context.Context, typeName, field strin return e.complexity.EvidenceConnection.PageInfo(childComplexity), true + case "EvidenceConnection.totalCount": + if e.complexity.EvidenceConnection.TotalCount == nil { + break + } + + return e.complexity.EvidenceConnection.TotalCount(childComplexity), true + case "EvidenceEdge.cursor": if e.complexity.EvidenceEdge.Cursor == nil { break @@ -5915,7 +5927,11 @@ type TaskEdge { node: Task! } -type EvidenceConnection { +type EvidenceConnection + @goModel( + model: "github.com/getprobo/probo/pkg/server/api/console/v1/types.EvidenceConnection" + ) { + totalCount: Int! @goField(forceResolver: true) edges: [EvidenceEdge!]! pageInfo: PageInfo! } @@ -19611,6 +19627,50 @@ func (ec *executionContext) fieldContext_Evidence_updatedAt(_ context.Context, f return fc, nil } +func (ec *executionContext) _EvidenceConnection_totalCount(ctx context.Context, field graphql.CollectedField, obj *types.EvidenceConnection) (ret graphql.Marshaler) { + fc, err := ec.fieldContext_EvidenceConnection_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.EvidenceConnection().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_EvidenceConnection_totalCount(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { + fc = &graphql.FieldContext{ + Object: "EvidenceConnection", + 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) _EvidenceConnection_edges(ctx context.Context, field graphql.CollectedField, obj *types.EvidenceConnection) (ret graphql.Marshaler) { fc, err := ec.fieldContext_EvidenceConnection_edges(ctx, field) if err != nil { @@ -19687,9 +19747,9 @@ func (ec *executionContext) _EvidenceConnection_pageInfo(ctx context.Context, fi } 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_EvidenceConnection_pageInfo(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { @@ -20987,6 +21047,8 @@ func (ec *executionContext) fieldContext_Measure_evidences(ctx context.Context, IsResolver: true, Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { switch field.Name { + case "totalCount": + return ec.fieldContext_EvidenceConnection_totalCount(ctx, field) case "edges": return ec.fieldContext_EvidenceConnection_edges(ctx, field) case "pageInfo": @@ -29670,6 +29732,8 @@ func (ec *executionContext) fieldContext_Task_evidences(ctx context.Context, fie IsResolver: true, Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { switch field.Name { + case "totalCount": + return ec.fieldContext_EvidenceConnection_totalCount(ctx, field) case "edges": return ec.fieldContext_EvidenceConnection_edges(ctx, field) case "pageInfo": @@ -43873,15 +43937,51 @@ func (ec *executionContext) _EvidenceConnection(ctx context.Context, sel ast.Sel switch field.Name { case "__typename": out.Values[i] = graphql.MarshalString("EvidenceConnection") + 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._EvidenceConnection_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._EvidenceConnection_edges(ctx, field, obj) if out.Values[i] == graphql.Null { - out.Invalids++ + atomic.AddUint32(&out.Invalids, 1) } case "pageInfo": out.Values[i] = ec._EvidenceConnection_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/evidence.go b/pkg/server/api/console/v1/types/evidence.go index 2d79a4f19..a85d94592 100644 --- a/pkg/server/api/console/v1/types/evidence.go +++ b/pkg/server/api/console/v1/types/evidence.go @@ -16,14 +16,28 @@ package types import ( "github.com/getprobo/probo/pkg/coredata" + "github.com/getprobo/probo/pkg/gid" "github.com/getprobo/probo/pkg/page" ) type ( EvidenceOrderBy OrderBy[coredata.EvidenceOrderField] + + EvidenceConnection struct { + TotalCount int + Edges []*EvidenceEdge + PageInfo PageInfo + + Resolver any + ParentID gid.GID + } ) -func NewEvidenceConnection(p *page.Page[*coredata.Evidence, coredata.EvidenceOrderField]) *EvidenceConnection { +func NewEvidenceConnection( + p *page.Page[*coredata.Evidence, coredata.EvidenceOrderField], + parentType any, + parentID gid.GID, +) *EvidenceConnection { var edges = make([]*EvidenceEdge, len(p.Data)) for i := range edges { @@ -32,7 +46,10 @@ func NewEvidenceConnection(p *page.Page[*coredata.Evidence, coredata.EvidenceOrd return &EvidenceConnection{ Edges: edges, - PageInfo: NewPageInfo(p), + PageInfo: *NewPageInfo(p), + + Resolver: parentType, + ParentID: parentID, } } diff --git a/pkg/server/api/console/v1/types/types.go b/pkg/server/api/console/v1/types/types.go index 205f8c3e1..9bf04fe31 100644 --- a/pkg/server/api/console/v1/types/types.go +++ b/pkg/server/api/console/v1/types/types.go @@ -624,11 +624,6 @@ type Evidence struct { func (Evidence) IsNode() {} func (this Evidence) GetID() gid.GID { return this.ID } -type EvidenceConnection struct { - Edges []*EvidenceEdge `json:"edges"` - PageInfo *PageInfo `json:"pageInfo"` -} - type EvidenceEdge struct { Cursor page.CursorKey `json:"cursor"` Node *Evidence `json:"node"` diff --git a/pkg/server/api/console/v1/v1_resolver.go b/pkg/server/api/console/v1/v1_resolver.go index 3d2783996..735e4f935 100644 --- a/pkg/server/api/console/v1/v1_resolver.go +++ b/pkg/server/api/console/v1/v1_resolver.go @@ -565,6 +565,28 @@ func (r *evidenceResolver) Measure(ctx context.Context, obj *types.Evidence) (*t return types.NewMeasure(measure), nil } +// TotalCount is the resolver for the totalCount field. +func (r *evidenceConnectionResolver) TotalCount(ctx context.Context, obj *types.EvidenceConnection) (int, error) { + svc := GetTenantService(ctx, r.proboSvc, obj.ParentID.TenantID()) + + switch obj.Resolver.(type) { + case *measureResolver: + count, err := svc.Evidences.CountForMeasureID(ctx, obj.ParentID) + if err != nil { + return 0, fmt.Errorf("cannot count tasks: %w", err) + } + return count, nil + case *taskResolver: + count, err := svc.Evidences.CountForTaskID(ctx, obj.ParentID) + if err != nil { + return 0, fmt.Errorf("cannot count tasks: %w", err) + } + return count, nil + } + + panic(fmt.Errorf("unsupported resolver: %T", obj.Resolver)) +} + // Organization is the resolver for the organization field. func (r *frameworkResolver) Organization(ctx context.Context, obj *types.Framework) (*types.Organization, error) { svc := GetTenantService(ctx, r.proboSvc, obj.ID.TenantID()) @@ -650,7 +672,7 @@ func (r *measureResolver) Evidences(ctx context.Context, obj *types.Measure, fir return nil, fmt.Errorf("cannot list measure evidences: %w", err) } - return types.NewEvidenceConnection(page), nil + return types.NewEvidenceConnection(page, r, obj.ID), nil } // Tasks is the resolver for the tasks field. @@ -2692,7 +2714,7 @@ func (r *taskResolver) Evidences(ctx context.Context, obj *types.Task, first *in panic(fmt.Errorf("failed to list task evidences: %w", err)) } - return types.NewEvidenceConnection(page), nil + return types.NewEvidenceConnection(page, r, obj.ID), nil } // TotalCount is the resolver for the totalCount field. @@ -2950,6 +2972,11 @@ func (r *Resolver) DocumentVersionSignature() schema.DocumentVersionSignatureRes // Evidence returns schema.EvidenceResolver implementation. func (r *Resolver) Evidence() schema.EvidenceResolver { return &evidenceResolver{r} } +// EvidenceConnection returns schema.EvidenceConnectionResolver implementation. +func (r *Resolver) EvidenceConnection() schema.EvidenceConnectionResolver { + return &evidenceConnectionResolver{r} +} + // Framework returns schema.FrameworkResolver implementation. func (r *Resolver) Framework() schema.FrameworkResolver { return &frameworkResolver{r} } @@ -3015,6 +3042,7 @@ type documentConnectionResolver struct{ *Resolver } type documentVersionResolver struct{ *Resolver } type documentVersionSignatureResolver struct{ *Resolver } type evidenceResolver struct{ *Resolver } +type evidenceConnectionResolver struct{ *Resolver } type frameworkResolver struct{ *Resolver } type frameworkConnectionResolver struct{ *Resolver } type measureResolver struct{ *Resolver }