From b43e61e3b9c71ea8d48fa03278e6fc9beb502eb5 Mon Sep 17 00:00:00 2001 From: Bryan Frimin Date: Fri, 23 May 2025 10:33:57 -0700 Subject: [PATCH] Improve data isolation for risks Signed-off-by: Bryan Frimin --- pkg/coredata/risk.go | 5 +- pkg/coredata/risk_mesure.go | 6 +- pkg/coredata/risk_policy.go | 6 +- pkg/probo/risk_service.go | 83 ++++++--- pkg/server/api/console/v1/schema.graphql | 6 +- pkg/server/api/console/v1/schema/schema.go | 196 +++++++++++++++++---- pkg/server/api/console/v1/types/types.go | 6 +- pkg/server/api/console/v1/v1_resolver.go | 10 +- 8 files changed, 248 insertions(+), 70 deletions(-) diff --git a/pkg/coredata/risk.go b/pkg/coredata/risk.go index a585117fd..7e11c25a0 100644 --- a/pkg/coredata/risk.go +++ b/pkg/coredata/risk.go @@ -324,15 +324,14 @@ func (r *Risk) Delete( ctx context.Context, conn pg.Conn, scope Scoper, + riskID gid.GID, ) error { q := ` DELETE FROM risks WHERE %s AND id = @id ` q = fmt.Sprintf(q, scope.SQLFragment()) - args := pgx.StrictNamedArgs{ - "id": r.ID, - } + args := pgx.StrictNamedArgs{"id": riskID} maps.Copy(args, scope.SQLArguments()) _, err := conn.Exec(ctx, q, args) diff --git a/pkg/coredata/risk_mesure.go b/pkg/coredata/risk_mesure.go index 516c51893..3f3dbe433 100644 --- a/pkg/coredata/risk_mesure.go +++ b/pkg/coredata/risk_mesure.go @@ -71,6 +71,8 @@ func (rm RiskMeasure) Delete( ctx context.Context, conn pg.Conn, scope Scoper, + riskID gid.GID, + measureID gid.GID, ) error { q := ` DELETE @@ -85,8 +87,8 @@ WHERE q = fmt.Sprintf(q, scope.SQLFragment()) args := pgx.StrictNamedArgs{ - "risk_id": rm.RiskID, - "measure_id": rm.MeasureID, + "risk_id": riskID, + "measure_id": measureID, } maps.Copy(args, scope.SQLArguments()) diff --git a/pkg/coredata/risk_policy.go b/pkg/coredata/risk_policy.go index e6d690a97..a9275e1b8 100644 --- a/pkg/coredata/risk_policy.go +++ b/pkg/coredata/risk_policy.go @@ -71,6 +71,8 @@ func (rp RiskPolicy) Delete( ctx context.Context, conn pg.Conn, scope Scoper, + riskID gid.GID, + policyID gid.GID, ) error { q := ` DELETE @@ -85,8 +87,8 @@ WHERE q = fmt.Sprintf(q, scope.SQLFragment()) args := pgx.StrictNamedArgs{ - "risk_id": rp.RiskID, - "policy_id": rp.PolicyID, + "risk_id": riskID, + "policy_id": policyID, } maps.Copy(args, scope.SQLArguments()) diff --git a/pkg/probo/risk_service.go b/pkg/probo/risk_service.go index 34e9190eb..93d5320b4 100644 --- a/pkg/probo/risk_service.go +++ b/pkg/probo/risk_service.go @@ -84,40 +84,68 @@ func (s RiskService) CreatePolicyMapping( ctx context.Context, riskID gid.GID, policyID gid.GID, -) error { - riskPolicy := &coredata.RiskPolicy{ - RiskID: riskID, - PolicyID: policyID, - TenantID: s.svc.scope.GetTenantID(), - CreatedAt: time.Now(), - } +) (*coredata.Risk, *coredata.Policy, error) { + risk := &coredata.Risk{} + policy := &coredata.Policy{} - return s.svc.pg.WithConn( + err := s.svc.pg.WithConn( ctx, func(conn pg.Conn) error { + if err := risk.LoadByID(ctx, conn, s.svc.scope, riskID); err != nil { + return fmt.Errorf("cannot load risk: %w", err) + } + + if err := policy.LoadByID(ctx, conn, s.svc.scope, policyID); err != nil { + return fmt.Errorf("cannot load policy: %w", err) + } + + riskPolicy := &coredata.RiskPolicy{ + RiskID: risk.ID, + PolicyID: policy.ID, + TenantID: s.svc.scope.GetTenantID(), + CreatedAt: time.Now(), + } + return riskPolicy.Insert(ctx, conn, s.svc.scope) }, ) + + if err != nil { + return nil, nil, fmt.Errorf("cannot create risk policy mapping: %w", err) + } + + return risk, policy, nil } func (s RiskService) DeletePolicyMapping( ctx context.Context, riskID gid.GID, policyID gid.GID, -) error { - riskPolicy := &coredata.RiskPolicy{ - RiskID: riskID, - PolicyID: policyID, - TenantID: s.svc.scope.GetTenantID(), - CreatedAt: time.Now(), - } +) (*coredata.Risk, *coredata.Policy, error) { + riskPolicy := &coredata.RiskPolicy{} + risk := &coredata.Risk{} + policy := &coredata.Policy{} - return s.svc.pg.WithConn( + err := s.svc.pg.WithConn( ctx, func(conn pg.Conn) error { - return riskPolicy.Delete(ctx, conn, s.svc.scope) + if err := risk.LoadByID(ctx, conn, s.svc.scope, riskID); err != nil { + return fmt.Errorf("cannot load risk: %w", err) + } + + if err := policy.LoadByID(ctx, conn, s.svc.scope, policyID); err != nil { + return fmt.Errorf("cannot load policy: %w", err) + } + + return riskPolicy.Delete(ctx, conn, s.svc.scope, risk.ID, policy.ID) }, ) + + if err != nil { + return nil, nil, fmt.Errorf("cannot delete risk policy mapping: %w", err) + } + + return risk, policy, nil } func (s RiskService) CreateMeasureMapping( @@ -183,7 +211,7 @@ func (s RiskService) DeleteMeasureMapping( CreatedAt: time.Now(), } - return riskMeasure.Delete(ctx, conn, s.svc.scope) + return riskMeasure.Delete(ctx, conn, s.svc.scope, risk.ID, measure.ID) }, ) @@ -199,10 +227,11 @@ func (s RiskService) Create( req CreateRiskRequest, ) (*coredata.Risk, error) { now := time.Now() - riskID := gid.New(s.svc.scope.GetTenantID(), coredata.RiskEntityType) + people := coredata.People{} + organization := coredata.Organization{} risk := &coredata.Risk{ - ID: riskID, + ID: gid.New(s.svc.scope.GetTenantID(), coredata.RiskEntityType), OrganizationID: req.OrganizationID, Name: req.Name, Description: req.Description, @@ -232,8 +261,11 @@ func (s RiskService) Create( err := s.svc.pg.WithConn( ctx, func(conn pg.Conn) error { + if err := organization.LoadByID(ctx, conn, s.svc.scope, req.OrganizationID); err != nil { + return fmt.Errorf("cannot load organization: %w", err) + } + if req.OwnerID != nil { - people := coredata.People{} if err := people.LoadByID(ctx, conn, s.svc.scope, *req.OwnerID); err != nil { return fmt.Errorf("cannot load owner: %w", err) } @@ -312,6 +344,11 @@ func (s RiskService) Update( } if req.OwnerID != nil { + people := coredata.People{} + if err := people.LoadByID(ctx, conn, s.svc.scope, *req.OwnerID); err != nil { + return fmt.Errorf("cannot load owner: %w", err) + } + risk.OwnerID = req.OwnerID } @@ -343,12 +380,12 @@ func (s RiskService) Delete( ctx context.Context, riskID gid.GID, ) error { - risk := &coredata.Risk{ID: riskID} + risk := &coredata.Risk{} return s.svc.pg.WithConn( ctx, func(conn pg.Conn) error { - return risk.Delete(ctx, conn, s.svc.scope) + return risk.Delete(ctx, conn, s.svc.scope, riskID) }, ) } diff --git a/pkg/server/api/console/v1/schema.graphql b/pkg/server/api/console/v1/schema.graphql index 022617558..b99082c81 100644 --- a/pkg/server/api/console/v1/schema.graphql +++ b/pkg/server/api/console/v1/schema.graphql @@ -1506,11 +1506,13 @@ type DeleteRiskMeasureMappingPayload { } type CreateRiskPolicyMappingPayload { - success: Boolean! + riskEdge: RiskEdge! + policyEdge: PolicyEdge! } type DeleteRiskPolicyMappingPayload { - success: Boolean! + deletedRiskId: ID! + deletedPolicyId: ID! } type RequestEvidencePayload { diff --git a/pkg/server/api/console/v1/schema/schema.go b/pkg/server/api/console/v1/schema/schema.go index 1b9bdfb04..b2f1b3eb5 100644 --- a/pkg/server/api/console/v1/schema/schema.go +++ b/pkg/server/api/console/v1/schema/schema.go @@ -160,7 +160,8 @@ type ComplexityRoot struct { } CreateRiskPolicyMappingPayload struct { - Success func(childComplexity int) int + PolicyEdge func(childComplexity int) int + RiskEdge func(childComplexity int) int } CreateTaskPayload struct { @@ -217,7 +218,8 @@ type ComplexityRoot struct { } DeleteRiskPolicyMappingPayload struct { - Success func(childComplexity int) int + DeletedPolicyID func(childComplexity int) int + DeletedRiskID func(childComplexity int) int } DeleteTaskPayload struct { @@ -1179,12 +1181,19 @@ func (e *executableSchema) Complexity(ctx context.Context, typeName, field strin return e.complexity.CreateRiskPayload.RiskEdge(childComplexity), true - case "CreateRiskPolicyMappingPayload.success": - if e.complexity.CreateRiskPolicyMappingPayload.Success == nil { + case "CreateRiskPolicyMappingPayload.policyEdge": + if e.complexity.CreateRiskPolicyMappingPayload.PolicyEdge == nil { break } - return e.complexity.CreateRiskPolicyMappingPayload.Success(childComplexity), true + return e.complexity.CreateRiskPolicyMappingPayload.PolicyEdge(childComplexity), true + + case "CreateRiskPolicyMappingPayload.riskEdge": + if e.complexity.CreateRiskPolicyMappingPayload.RiskEdge == nil { + break + } + + return e.complexity.CreateRiskPolicyMappingPayload.RiskEdge(childComplexity), true case "CreateTaskPayload.taskEdge": if e.complexity.CreateTaskPayload.TaskEdge == nil { @@ -1284,12 +1293,19 @@ func (e *executableSchema) Complexity(ctx context.Context, typeName, field strin return e.complexity.DeleteRiskPayload.DeletedRiskID(childComplexity), true - case "DeleteRiskPolicyMappingPayload.success": - if e.complexity.DeleteRiskPolicyMappingPayload.Success == nil { + case "DeleteRiskPolicyMappingPayload.deletedPolicyId": + if e.complexity.DeleteRiskPolicyMappingPayload.DeletedPolicyID == nil { break } - return e.complexity.DeleteRiskPolicyMappingPayload.Success(childComplexity), true + return e.complexity.DeleteRiskPolicyMappingPayload.DeletedPolicyID(childComplexity), true + + case "DeleteRiskPolicyMappingPayload.deletedRiskId": + if e.complexity.DeleteRiskPolicyMappingPayload.DeletedRiskID == nil { + break + } + + return e.complexity.DeleteRiskPolicyMappingPayload.DeletedRiskID(childComplexity), true case "DeleteTaskPayload.deletedTaskId": if e.complexity.DeleteTaskPayload.DeletedTaskID == nil { @@ -5493,11 +5509,13 @@ type DeleteRiskMeasureMappingPayload { } type CreateRiskPolicyMappingPayload { - success: Boolean! + riskEdge: RiskEdge! + policyEdge: PolicyEdge! } type DeleteRiskPolicyMappingPayload { - success: Boolean! + deletedRiskId: ID! + deletedPolicyId: ID! } type RequestEvidencePayload { @@ -11444,8 +11462,8 @@ func (ec *executionContext) fieldContext_CreateRiskPayload_riskEdge(_ context.Co return fc, nil } -func (ec *executionContext) _CreateRiskPolicyMappingPayload_success(ctx context.Context, field graphql.CollectedField, obj *types.CreateRiskPolicyMappingPayload) (ret graphql.Marshaler) { - fc, err := ec.fieldContext_CreateRiskPolicyMappingPayload_success(ctx, field) +func (ec *executionContext) _CreateRiskPolicyMappingPayload_riskEdge(ctx context.Context, field graphql.CollectedField, obj *types.CreateRiskPolicyMappingPayload) (ret graphql.Marshaler) { + fc, err := ec.fieldContext_CreateRiskPolicyMappingPayload_riskEdge(ctx, field) if err != nil { return graphql.Null } @@ -11458,7 +11476,7 @@ func (ec *executionContext) _CreateRiskPolicyMappingPayload_success(ctx context. }() resTmp, err := ec.ResolverMiddleware(ctx, func(rctx context.Context) (any, error) { ctx = rctx // use context from middleware stack in children - return obj.Success, nil + return obj.RiskEdge, nil }) if err != nil { ec.Error(ctx, err) @@ -11470,19 +11488,75 @@ func (ec *executionContext) _CreateRiskPolicyMappingPayload_success(ctx context. } return graphql.Null } - res := resTmp.(bool) + res := resTmp.(*types.RiskEdge) fc.Result = res - return ec.marshalNBoolean2bool(ctx, field.Selections, res) + return ec.marshalNRiskEdge2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐRiskEdge(ctx, field.Selections, res) } -func (ec *executionContext) fieldContext_CreateRiskPolicyMappingPayload_success(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { +func (ec *executionContext) fieldContext_CreateRiskPolicyMappingPayload_riskEdge(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { fc = &graphql.FieldContext{ Object: "CreateRiskPolicyMappingPayload", Field: field, IsMethod: false, IsResolver: false, Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { - return nil, errors.New("field of type Boolean does not have child fields") + switch field.Name { + case "cursor": + return ec.fieldContext_RiskEdge_cursor(ctx, field) + case "node": + return ec.fieldContext_RiskEdge_node(ctx, field) + } + return nil, fmt.Errorf("no field named %q was found under type RiskEdge", field.Name) + }, + } + return fc, nil +} + +func (ec *executionContext) _CreateRiskPolicyMappingPayload_policyEdge(ctx context.Context, field graphql.CollectedField, obj *types.CreateRiskPolicyMappingPayload) (ret graphql.Marshaler) { + fc, err := ec.fieldContext_CreateRiskPolicyMappingPayload_policyEdge(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 obj.PolicyEdge, nil + }) + 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.(*types.PolicyEdge) + fc.Result = res + return ec.marshalNPolicyEdge2ᚖgithubᚗcomᚋgetproboᚋproboᚋpkgᚋserverᚋapiᚋconsoleᚋv1ᚋtypesᚐPolicyEdge(ctx, field.Selections, res) +} + +func (ec *executionContext) fieldContext_CreateRiskPolicyMappingPayload_policyEdge(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { + fc = &graphql.FieldContext{ + Object: "CreateRiskPolicyMappingPayload", + Field: field, + IsMethod: false, + IsResolver: false, + Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { + switch field.Name { + case "cursor": + return ec.fieldContext_PolicyEdge_cursor(ctx, field) + case "node": + return ec.fieldContext_PolicyEdge_node(ctx, field) + } + return nil, fmt.Errorf("no field named %q was found under type PolicyEdge", field.Name) }, } return fc, nil @@ -12122,8 +12196,8 @@ func (ec *executionContext) fieldContext_DeleteRiskPayload_deletedRiskId(_ conte return fc, nil } -func (ec *executionContext) _DeleteRiskPolicyMappingPayload_success(ctx context.Context, field graphql.CollectedField, obj *types.DeleteRiskPolicyMappingPayload) (ret graphql.Marshaler) { - fc, err := ec.fieldContext_DeleteRiskPolicyMappingPayload_success(ctx, field) +func (ec *executionContext) _DeleteRiskPolicyMappingPayload_deletedRiskId(ctx context.Context, field graphql.CollectedField, obj *types.DeleteRiskPolicyMappingPayload) (ret graphql.Marshaler) { + fc, err := ec.fieldContext_DeleteRiskPolicyMappingPayload_deletedRiskId(ctx, field) if err != nil { return graphql.Null } @@ -12136,7 +12210,7 @@ func (ec *executionContext) _DeleteRiskPolicyMappingPayload_success(ctx context. }() resTmp, err := ec.ResolverMiddleware(ctx, func(rctx context.Context) (any, error) { ctx = rctx // use context from middleware stack in children - return obj.Success, nil + return obj.DeletedRiskID, nil }) if err != nil { ec.Error(ctx, err) @@ -12148,19 +12222,63 @@ func (ec *executionContext) _DeleteRiskPolicyMappingPayload_success(ctx context. } return graphql.Null } - res := resTmp.(bool) + res := resTmp.(gid.GID) fc.Result = res - return ec.marshalNBoolean2bool(ctx, field.Selections, res) + return ec.marshalNID2githubᚗcomᚋgetproboᚋproboᚋpkgᚋgidᚐGID(ctx, field.Selections, res) } -func (ec *executionContext) fieldContext_DeleteRiskPolicyMappingPayload_success(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { +func (ec *executionContext) fieldContext_DeleteRiskPolicyMappingPayload_deletedRiskId(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { fc = &graphql.FieldContext{ Object: "DeleteRiskPolicyMappingPayload", Field: field, IsMethod: false, IsResolver: false, Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { - return nil, errors.New("field of type Boolean does not have child fields") + return nil, errors.New("field of type ID does not have child fields") + }, + } + return fc, nil +} + +func (ec *executionContext) _DeleteRiskPolicyMappingPayload_deletedPolicyId(ctx context.Context, field graphql.CollectedField, obj *types.DeleteRiskPolicyMappingPayload) (ret graphql.Marshaler) { + fc, err := ec.fieldContext_DeleteRiskPolicyMappingPayload_deletedPolicyId(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 obj.DeletedPolicyID, nil + }) + 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.(gid.GID) + fc.Result = res + return ec.marshalNID2githubᚗcomᚋgetproboᚋproboᚋpkgᚋgidᚐGID(ctx, field.Selections, res) +} + +func (ec *executionContext) fieldContext_DeleteRiskPolicyMappingPayload_deletedPolicyId(_ context.Context, field graphql.CollectedField) (fc *graphql.FieldContext, err error) { + fc = &graphql.FieldContext{ + Object: "DeleteRiskPolicyMappingPayload", + Field: field, + IsMethod: false, + IsResolver: false, + Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { + return nil, errors.New("field of type ID does not have child fields") }, } return fc, nil @@ -16747,8 +16865,10 @@ func (ec *executionContext) fieldContext_Mutation_createRiskPolicyMapping(ctx co IsResolver: true, Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { switch field.Name { - case "success": - return ec.fieldContext_CreateRiskPolicyMappingPayload_success(ctx, field) + case "riskEdge": + return ec.fieldContext_CreateRiskPolicyMappingPayload_riskEdge(ctx, field) + case "policyEdge": + return ec.fieldContext_CreateRiskPolicyMappingPayload_policyEdge(ctx, field) } return nil, fmt.Errorf("no field named %q was found under type CreateRiskPolicyMappingPayload", field.Name) }, @@ -16806,8 +16926,10 @@ func (ec *executionContext) fieldContext_Mutation_deleteRiskPolicyMapping(ctx co IsResolver: true, Child: func(ctx context.Context, field graphql.CollectedField) (*graphql.FieldContext, error) { switch field.Name { - case "success": - return ec.fieldContext_DeleteRiskPolicyMappingPayload_success(ctx, field) + case "deletedRiskId": + return ec.fieldContext_DeleteRiskPolicyMappingPayload_deletedRiskId(ctx, field) + case "deletedPolicyId": + return ec.fieldContext_DeleteRiskPolicyMappingPayload_deletedPolicyId(ctx, field) } return nil, fmt.Errorf("no field named %q was found under type DeleteRiskPolicyMappingPayload", field.Name) }, @@ -34198,8 +34320,13 @@ func (ec *executionContext) _CreateRiskPolicyMappingPayload(ctx context.Context, switch field.Name { case "__typename": out.Values[i] = graphql.MarshalString("CreateRiskPolicyMappingPayload") - case "success": - out.Values[i] = ec._CreateRiskPolicyMappingPayload_success(ctx, field, obj) + case "riskEdge": + out.Values[i] = ec._CreateRiskPolicyMappingPayload_riskEdge(ctx, field, obj) + if out.Values[i] == graphql.Null { + out.Invalids++ + } + case "policyEdge": + out.Values[i] = ec._CreateRiskPolicyMappingPayload_policyEdge(ctx, field, obj) if out.Values[i] == graphql.Null { out.Invalids++ } @@ -34749,8 +34876,13 @@ func (ec *executionContext) _DeleteRiskPolicyMappingPayload(ctx context.Context, switch field.Name { case "__typename": out.Values[i] = graphql.MarshalString("DeleteRiskPolicyMappingPayload") - case "success": - out.Values[i] = ec._DeleteRiskPolicyMappingPayload_success(ctx, field, obj) + case "deletedRiskId": + out.Values[i] = ec._DeleteRiskPolicyMappingPayload_deletedRiskId(ctx, field, obj) + if out.Values[i] == graphql.Null { + out.Invalids++ + } + case "deletedPolicyId": + out.Values[i] = ec._DeleteRiskPolicyMappingPayload_deletedPolicyId(ctx, field, obj) if out.Values[i] == graphql.Null { out.Invalids++ } diff --git a/pkg/server/api/console/v1/types/types.go b/pkg/server/api/console/v1/types/types.go index d5c5fcb4f..3db3a975c 100644 --- a/pkg/server/api/console/v1/types/types.go +++ b/pkg/server/api/console/v1/types/types.go @@ -214,7 +214,8 @@ type CreateRiskPolicyMappingInput struct { } type CreateRiskPolicyMappingPayload struct { - Success bool `json:"success"` + RiskEdge *RiskEdge `json:"riskEdge"` + PolicyEdge *PolicyEdge `json:"policyEdge"` } type CreateTaskInput struct { @@ -357,7 +358,8 @@ type DeleteRiskPolicyMappingInput struct { } type DeleteRiskPolicyMappingPayload struct { - Success bool `json:"success"` + DeletedRiskID gid.GID `json:"deletedRiskId"` + DeletedPolicyID gid.GID `json:"deletedPolicyId"` } type DeleteTaskInput struct { diff --git a/pkg/server/api/console/v1/v1_resolver.go b/pkg/server/api/console/v1/v1_resolver.go index edb7c6129..dc0dab124 100644 --- a/pkg/server/api/console/v1/v1_resolver.go +++ b/pkg/server/api/console/v1/v1_resolver.go @@ -949,13 +949,14 @@ func (r *mutationResolver) DeleteRiskMeasureMapping(ctx context.Context, input t func (r *mutationResolver) CreateRiskPolicyMapping(ctx context.Context, input types.CreateRiskPolicyMappingInput) (*types.CreateRiskPolicyMappingPayload, error) { svc := GetTenantService(ctx, r.proboSvc, input.RiskID.TenantID()) - err := svc.Risks.CreatePolicyMapping(ctx, input.RiskID, input.PolicyID) + risk, policy, err := svc.Risks.CreatePolicyMapping(ctx, input.RiskID, input.PolicyID) if err != nil { panic(fmt.Errorf("cannot create risk policy mapping: %w", err)) } return &types.CreateRiskPolicyMappingPayload{ - Success: true, + RiskEdge: types.NewRiskEdge(risk, coredata.RiskOrderFieldCreatedAt), + PolicyEdge: types.NewPolicyEdge(policy, coredata.PolicyOrderFieldTitle), }, nil } @@ -963,13 +964,14 @@ func (r *mutationResolver) CreateRiskPolicyMapping(ctx context.Context, input ty func (r *mutationResolver) DeleteRiskPolicyMapping(ctx context.Context, input types.DeleteRiskPolicyMappingInput) (*types.DeleteRiskPolicyMappingPayload, error) { svc := GetTenantService(ctx, r.proboSvc, input.RiskID.TenantID()) - err := svc.Risks.DeletePolicyMapping(ctx, input.RiskID, input.PolicyID) + risk, policy, err := svc.Risks.DeletePolicyMapping(ctx, input.RiskID, input.PolicyID) if err != nil { panic(fmt.Errorf("cannot delete risk policy mapping: %w", err)) } return &types.DeleteRiskPolicyMappingPayload{ - Success: true, + DeletedRiskID: risk.ID, + DeletedPolicyID: policy.ID, }, nil }