diff --git a/pkg/coredata/people.go b/pkg/coredata/people.go index 97690e9b2..f0c0c4ede 100644 --- a/pkg/coredata/people.go +++ b/pkg/coredata/people.go @@ -293,6 +293,38 @@ DELETE FROM peoples WHERE %s AND id = @people_id return err } +func (p *Peoples) CountByOrganizationID( + ctx context.Context, + conn pg.Conn, + scope Scoper, + organizationID gid.GID, +) (int, error) { + q := ` +SELECT + COUNT(id) +FROM + peoples +WHERE + %s + AND organization_id = @organization_id +` + + q = fmt.Sprintf(q, scope.SQLFragment()) + + args := pgx.StrictNamedArgs{"organization_id": organizationID} + 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 count people: %w", err) + } + + return count, nil +} + func (p *Peoples) LoadByOrganizationID( ctx context.Context, conn pg.Conn, diff --git a/pkg/probo/people_service.go b/pkg/probo/people_service.go index 94bae22f9..c8cc6ca26 100644 --- a/pkg/probo/people_service.go +++ b/pkg/probo/people_service.go @@ -95,6 +95,32 @@ func (s PeopleService) GetByUserID( return people, nil } +func (s PeopleService) CountForOrganizationID( + ctx context.Context, + organizationID gid.GID, +) (int, error) { + var count int + + err := s.svc.pg.WithConn( + ctx, + func(conn pg.Conn) (err error) { + peoples := coredata.Peoples{} + count, err = peoples.CountByOrganizationID(ctx, conn, s.svc.scope, organizationID) + if err != nil { + return fmt.Errorf("cannot count peoples: %w", err) + } + + return nil + }, + ) + + if err != nil { + return 0, err + } + + return count, nil +} + func (s PeopleService) 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 7d67a59c8..1fda560fd 100644 --- a/pkg/server/api/console/v1/schema.graphql +++ b/pkg/server/api/console/v1/schema.graphql @@ -1070,7 +1070,11 @@ type UserEdge { node: User! } -type PeopleConnection { +type PeopleConnection + @goModel( + model: "github.com/getprobo/probo/pkg/server/api/console/v1/types.PeopleConnection" + ) { + totalCount: Int! @goField(forceResolver: true) edges: [PeopleEdge!]! pageInfo: PageInfo! } diff --git a/pkg/server/api/console/v1/schema/schema.go b/pkg/server/api/console/v1/schema/schema.go index 2a3c3f4e0..edbb3d308 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 { MeasureConnection() MeasureConnectionResolver Mutation() MutationResolver Organization() OrganizationResolver + PeopleConnection() PeopleConnectionResolver Query() QueryResolver Risk() RiskResolver RiskConnection() RiskConnectionResolver @@ -605,8 +606,9 @@ type ComplexityRoot struct { } PeopleConnection 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 } PeopleEdge struct { @@ -1026,6 +1028,9 @@ type OrganizationResolver interface { Assets(ctx context.Context, obj *types.Organization, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.AssetOrder) (*types.AssetConnection, error) Data(ctx context.Context, obj *types.Organization, first *int, after *page.CursorKey, last *int, before *page.CursorKey, orderBy *types.DatumOrder) (*types.DatumConnection, error) } +type PeopleConnectionResolver interface { + TotalCount(ctx context.Context, obj *types.PeopleConnection) (int, error) +} type QueryResolver interface { Node(ctx context.Context, id gid.GID) (types.Node, error) Viewer(ctx context.Context) (*types.Viewer, error) @@ -3588,6 +3593,13 @@ func (e *executableSchema) Complexity(ctx context.Context, typeName, field strin return e.complexity.PeopleConnection.PageInfo(childComplexity), true + case "PeopleConnection.totalCount": + if e.complexity.PeopleConnection.TotalCount == nil { + break + } + + return e.complexity.PeopleConnection.TotalCount(childComplexity), true + case "PeopleEdge.cursor": if e.complexity.PeopleEdge.Cursor == nil { break @@ -5863,7 +5875,11 @@ type UserEdge { node: User! } -type PeopleConnection { +type PeopleConnection + @goModel( + model: "github.com/getprobo/probo/pkg/server/api/console/v1/types.PeopleConnection" + ) { + totalCount: Int! @goField(forceResolver: true) edges: [PeopleEdge!]! pageInfo: PageInfo! } @@ -25844,6 +25860,8 @@ func (ec *executionContext) fieldContext_Organization_peoples(ctx context.Contex IsResolver: true, Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { switch field.Name { + case "totalCount": + return ec.fieldContext_PeopleConnection_totalCount(ctx, field) case "edges": return ec.fieldContext_PeopleConnection_edges(ctx, field) case "pageInfo": @@ -27157,6 +27175,50 @@ func (ec *executionContext) fieldContext_People_updatedAt(_ context.Context, fie return fc, nil } +func (ec *executionContext) _PeopleConnection_totalCount(ctx context.Context, field graphql.CollectedField, obj *types.PeopleConnection) (ret graphql.Marshaler) { + fc, err := ec.fieldContext_PeopleConnection_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.PeopleConnection().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_PeopleConnection_totalCount(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { + fc = &graphql.FieldContext{ + Object: "PeopleConnection", + 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) _PeopleConnection_edges(ctx context.Context, field graphql.CollectedField, obj *types.PeopleConnection) (ret graphql.Marshaler) { fc, err := ec.fieldContext_PeopleConnection_edges(ctx, field) if err != nil { @@ -27233,9 +27295,9 @@ func (ec *executionContext) _PeopleConnection_pageInfo(ctx context.Context, fiel } 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_PeopleConnection_pageInfo(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { @@ -46166,15 +46228,51 @@ func (ec *executionContext) _PeopleConnection(ctx context.Context, sel ast.Selec switch field.Name { case "__typename": out.Values[i] = graphql.MarshalString("PeopleConnection") + 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._PeopleConnection_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._PeopleConnection_edges(ctx, field, obj) if out.Values[i] == graphql.Null { - out.Invalids++ + atomic.AddUint32(&out.Invalids, 1) } case "pageInfo": out.Values[i] = ec._PeopleConnection_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/people.go b/pkg/server/api/console/v1/types/people.go index c0f39abcc..c41dba7d5 100644 --- a/pkg/server/api/console/v1/types/people.go +++ b/pkg/server/api/console/v1/types/people.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 ( PeopleOrderBy OrderBy[coredata.PeopleOrderField] + + PeopleConnection struct { + TotalCount int + Edges []*PeopleEdge + PageInfo PageInfo + + Resolver any + ParentID gid.GID + } ) -func NewPeopleConnection(p *page.Page[*coredata.People, coredata.PeopleOrderField]) *PeopleConnection { +func NewPeopleConnection( + p *page.Page[*coredata.People, coredata.PeopleOrderField], + parentType any, + parentID gid.GID, +) *PeopleConnection { var edges = make([]*PeopleEdge, len(p.Data)) for i := range edges { @@ -32,7 +46,10 @@ func NewPeopleConnection(p *page.Page[*coredata.People, coredata.PeopleOrderFiel return &PeopleConnection{ 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 96a8d982a..00f3486f6 100644 --- a/pkg/server/api/console/v1/types/types.go +++ b/pkg/server/api/console/v1/types/types.go @@ -792,11 +792,6 @@ type People struct { func (People) IsNode() {} func (this People) GetID() gid.GID { return this.ID } -type PeopleConnection struct { - Edges []*PeopleEdge `json:"edges"` - PageInfo *PageInfo `json:"pageInfo"` -} - type PeopleEdge struct { Cursor page.CursorKey `json:"cursor"` Node *People `json:"node"` diff --git a/pkg/server/api/console/v1/v1_resolver.go b/pkg/server/api/console/v1/v1_resolver.go index 44b8a5031..fee9bc805 100644 --- a/pkg/server/api/console/v1/v1_resolver.go +++ b/pkg/server/api/console/v1/v1_resolver.go @@ -2200,7 +2200,7 @@ func (r *organizationResolver) Peoples(ctx context.Context, obj *types.Organizat panic(fmt.Errorf("cannot list organization peoples: %w", err)) } - return types.NewPeopleConnection(page), nil + return types.NewPeopleConnection(page, r, obj.ID), nil } // Documents is the resolver for the documents field. @@ -2368,6 +2368,22 @@ func (r *organizationResolver) Data(ctx context.Context, obj *types.Organization return types.NewDataConnection(page), nil } +// TotalCount is the resolver for the totalCount field. +func (r *peopleConnectionResolver) TotalCount(ctx context.Context, obj *types.PeopleConnection) (int, error) { + svc := GetTenantService(ctx, r.proboSvc, obj.ParentID.TenantID()) + + switch obj.Resolver.(type) { + case *organizationResolver: + count, err := svc.Peoples.CountForOrganizationID(ctx, obj.ParentID) + if err != nil { + return 0, fmt.Errorf("cannot count peoples: %w", err) + } + return count, nil + } + + panic(fmt.Errorf("unsupported resolver: %T", obj.Resolver)) +} + // Node is the resolver for the node field. func (r *queryResolver) Node(ctx context.Context, id gid.GID) (types.Node, error) { svc := GetTenantService(ctx, r.proboSvc, id.TenantID()) @@ -3027,6 +3043,11 @@ func (r *Resolver) Mutation() schema.MutationResolver { return &mutationResolver // Organization returns schema.OrganizationResolver implementation. func (r *Resolver) Organization() schema.OrganizationResolver { return &organizationResolver{r} } +// PeopleConnection returns schema.PeopleConnectionResolver implementation. +func (r *Resolver) PeopleConnection() schema.PeopleConnectionResolver { + return &peopleConnectionResolver{r} +} + // Query returns schema.QueryResolver implementation. func (r *Resolver) Query() schema.QueryResolver { return &queryResolver{r} } @@ -3082,6 +3103,7 @@ type measureResolver struct{ *Resolver } type measureConnectionResolver struct{ *Resolver } type mutationResolver struct{ *Resolver } type organizationResolver struct{ *Resolver } +type peopleConnectionResolver struct{ *Resolver } type queryResolver struct{ *Resolver } type riskResolver struct{ *Resolver } type riskConnectionResolver struct{ *Resolver }