Render mermaid diagram per risk assessment scope
Each scope card now shows a flowchart of its nodes, processes, and threats, with a distinct shape per type: stadium for entities, hexagon for boundaries, rectangle for assets, cylinder for data, and a red hexagon for threats attached via dashed edges to their process target. The Mermaid source is built on the backend and exposed as a new `mermaid` field on RiskAssessmentScope; the frontend just renders it via @probo/ui's MermaidDiagram and shows a copy button + legend. Signed-off-by: Sacha Al Himdani <sacha@getprobo.com>
This commit is contained in:
@@ -106,6 +106,45 @@ WHERE
|
||||
return nil
|
||||
}
|
||||
|
||||
func (ns *RiskAssessmentNodes) LoadAllByRiskAssessmentScopeID(
|
||||
ctx context.Context,
|
||||
conn pg.Querier,
|
||||
scope Scoper,
|
||||
riskAssessmentScopeID gid.GID,
|
||||
) error {
|
||||
q := `
|
||||
SELECT
|
||||
id,
|
||||
organization_id,
|
||||
risk_assessment_scope_id,
|
||||
node_type,
|
||||
name,
|
||||
created_at,
|
||||
updated_at
|
||||
FROM
|
||||
risk_assessment_nodes
|
||||
WHERE
|
||||
%s
|
||||
AND risk_assessment_scope_id = @risk_assessment_scope_id
|
||||
ORDER BY
|
||||
created_at ASC, id ASC
|
||||
`
|
||||
q = fmt.Sprintf(q, scope.SQLFragment())
|
||||
args := pgx.NamedArgs{"risk_assessment_scope_id": riskAssessmentScopeID}
|
||||
maps.Copy(args, scope.SQLArguments())
|
||||
|
||||
rows, err := conn.Query(ctx, q, args)
|
||||
if err != nil {
|
||||
return fmt.Errorf("cannot query risk assessment nodes: %w", err)
|
||||
}
|
||||
results, err := pgx.CollectRows(rows, pgx.RowToAddrOfStructByName[RiskAssessmentNode])
|
||||
if err != nil {
|
||||
return fmt.Errorf("cannot collect risk assessment nodes: %w", err)
|
||||
}
|
||||
*ns = results
|
||||
return nil
|
||||
}
|
||||
|
||||
func (ns *RiskAssessmentNodes) CountByRiskAssessmentScopeID(
|
||||
ctx context.Context,
|
||||
conn pg.Querier,
|
||||
|
||||
@@ -108,6 +108,46 @@ WHERE
|
||||
return nil
|
||||
}
|
||||
|
||||
func (ps *RiskAssessmentProcesses) LoadAllByRiskAssessmentScopeID(
|
||||
ctx context.Context,
|
||||
conn pg.Querier,
|
||||
scope Scoper,
|
||||
riskAssessmentScopeID gid.GID,
|
||||
) error {
|
||||
q := `
|
||||
SELECT
|
||||
id,
|
||||
organization_id,
|
||||
risk_assessment_scope_id,
|
||||
source_node_id,
|
||||
target_node_id,
|
||||
name,
|
||||
created_at,
|
||||
updated_at
|
||||
FROM
|
||||
risk_assessment_processes
|
||||
WHERE
|
||||
%s
|
||||
AND risk_assessment_scope_id = @risk_assessment_scope_id
|
||||
ORDER BY
|
||||
created_at ASC, id ASC
|
||||
`
|
||||
q = fmt.Sprintf(q, scope.SQLFragment())
|
||||
args := pgx.NamedArgs{"risk_assessment_scope_id": riskAssessmentScopeID}
|
||||
maps.Copy(args, scope.SQLArguments())
|
||||
|
||||
rows, err := conn.Query(ctx, q, args)
|
||||
if err != nil {
|
||||
return fmt.Errorf("cannot query risk assessment processes: %w", err)
|
||||
}
|
||||
results, err := pgx.CollectRows(rows, pgx.RowToAddrOfStructByName[RiskAssessmentProcess])
|
||||
if err != nil {
|
||||
return fmt.Errorf("cannot collect risk assessment processes: %w", err)
|
||||
}
|
||||
*ps = results
|
||||
return nil
|
||||
}
|
||||
|
||||
func (ps *RiskAssessmentProcesses) CountByRiskAssessmentScopeID(
|
||||
ctx context.Context,
|
||||
conn pg.Querier,
|
||||
|
||||
@@ -108,6 +108,46 @@ WHERE
|
||||
return nil
|
||||
}
|
||||
|
||||
func (ts *RiskAssessmentThreats) LoadAllByRiskAssessmentScopeID(
|
||||
ctx context.Context,
|
||||
conn pg.Querier,
|
||||
scope Scoper,
|
||||
riskAssessmentScopeID gid.GID,
|
||||
) error {
|
||||
q := `
|
||||
SELECT
|
||||
id,
|
||||
organization_id,
|
||||
risk_assessment_scope_id,
|
||||
process_id,
|
||||
name,
|
||||
category,
|
||||
created_at,
|
||||
updated_at
|
||||
FROM
|
||||
risk_assessment_threats
|
||||
WHERE
|
||||
%s
|
||||
AND risk_assessment_scope_id = @risk_assessment_scope_id
|
||||
ORDER BY
|
||||
created_at ASC, id ASC
|
||||
`
|
||||
q = fmt.Sprintf(q, scope.SQLFragment())
|
||||
args := pgx.NamedArgs{"risk_assessment_scope_id": riskAssessmentScopeID}
|
||||
maps.Copy(args, scope.SQLArguments())
|
||||
|
||||
rows, err := conn.Query(ctx, q, args)
|
||||
if err != nil {
|
||||
return fmt.Errorf("cannot query risk threats: %w", err)
|
||||
}
|
||||
results, err := pgx.CollectRows(rows, pgx.RowToAddrOfStructByName[RiskAssessmentThreat])
|
||||
if err != nil {
|
||||
return fmt.Errorf("cannot collect risk threats: %w", err)
|
||||
}
|
||||
*ts = results
|
||||
return nil
|
||||
}
|
||||
|
||||
func (ts *RiskAssessmentThreats) CountByRiskAssessmentScopeID(
|
||||
ctx context.Context,
|
||||
conn pg.Querier,
|
||||
|
||||
157
pkg/riskmanagement/mermaid.go
Normal file
157
pkg/riskmanagement/mermaid.go
Normal file
@@ -0,0 +1,157 @@
|
||||
// Copyright (c) 2026 Probo Inc <hello@getprobo.com>.
|
||||
//
|
||||
// Permission to use, copy, modify, and/or distribute this software for any
|
||||
// purpose with or without fee is hereby granted, provided that the above
|
||||
// copyright notice and this permission notice appear in all copies.
|
||||
//
|
||||
// THE SOFTWARE IS PROVIDED "AS IS" AND THE AUTHOR DISCLAIMS ALL WARRANTIES WITH
|
||||
// REGARD TO THIS SOFTWARE INCLUDING ALL IMPLIED WARRANTIES OF MERCHANTABILITY
|
||||
// AND FITNESS. IN NO EVENT SHALL THE AUTHOR BE LIABLE FOR ANY SPECIAL, DIRECT,
|
||||
// INDIRECT, OR CONSEQUENTIAL DAMAGES OR ANY DAMAGES WHATSOEVER RESULTING FROM
|
||||
// LOSS OF USE, DATA OR PROFITS, WHETHER IN AN ACTION OF CONTRACT, NEGLIGENCE OR
|
||||
// OTHER TORTIOUS ACTION, ARISING OUT OF OR IN CONNECTION WITH THE USE OR
|
||||
// PERFORMANCE OF THIS SOFTWARE.
|
||||
|
||||
package riskmanagement
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"go.gearno.de/kit/pg"
|
||||
"go.probo.inc/probo/pkg/coredata"
|
||||
"go.probo.inc/probo/pkg/gid"
|
||||
)
|
||||
|
||||
func (s *Service) BuildScopeMermaidChart(ctx context.Context, scope coredata.Scoper, scopeID gid.GID) (string, error) {
|
||||
var (
|
||||
nodes coredata.RiskAssessmentNodes
|
||||
processes coredata.RiskAssessmentProcesses
|
||||
threats coredata.RiskAssessmentThreats
|
||||
)
|
||||
|
||||
err := s.pg.WithConn(ctx, func(ctx context.Context, conn pg.Querier) error {
|
||||
if err := nodes.LoadAllByRiskAssessmentScopeID(ctx, conn, scope, scopeID); err != nil {
|
||||
return fmt.Errorf("cannot load nodes: %w", err)
|
||||
}
|
||||
if err := processes.LoadAllByRiskAssessmentScopeID(ctx, conn, scope, scopeID); err != nil {
|
||||
return fmt.Errorf("cannot load processes: %w", err)
|
||||
}
|
||||
if err := threats.LoadAllByRiskAssessmentScopeID(ctx, conn, scope, scopeID); err != nil {
|
||||
return fmt.Errorf("cannot load threats: %w", err)
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
|
||||
return buildScopeMermaidChart(nodes, processes, threats), nil
|
||||
}
|
||||
|
||||
func buildScopeMermaidChart(
|
||||
nodes coredata.RiskAssessmentNodes,
|
||||
processes coredata.RiskAssessmentProcesses,
|
||||
threats coredata.RiskAssessmentThreats,
|
||||
) string {
|
||||
if len(nodes) == 0 {
|
||||
return ""
|
||||
}
|
||||
|
||||
nodeAlias := make(map[gid.GID]string, len(nodes))
|
||||
for i, n := range nodes {
|
||||
nodeAlias[n.ID] = fmt.Sprintf("n%d", i)
|
||||
}
|
||||
|
||||
var b strings.Builder
|
||||
b.WriteString("flowchart LR\n")
|
||||
|
||||
for _, n := range nodes {
|
||||
id := nodeAlias[n.ID]
|
||||
fmt.Fprintf(&b, " %s\n", mermaidNodeShape(n.NodeType, id, n.Name))
|
||||
fmt.Fprintf(&b, " class %s %s\n", id, mermaidNodeClass(n.NodeType))
|
||||
}
|
||||
|
||||
for _, p := range processes {
|
||||
src, srcOK := nodeAlias[p.SourceNodeID]
|
||||
dst, dstOK := nodeAlias[p.TargetNodeID]
|
||||
if !srcOK || !dstOK {
|
||||
continue
|
||||
}
|
||||
fmt.Fprintf(&b, " %s -- \"%s\" --> %s\n", src, escapeMermaidLabel(p.Name), dst)
|
||||
}
|
||||
|
||||
processTarget := make(map[gid.GID]gid.GID, len(processes))
|
||||
for _, p := range processes {
|
||||
processTarget[p.ID] = p.TargetNodeID
|
||||
}
|
||||
|
||||
for i, t := range threats {
|
||||
target, ok := processTarget[t.ProcessID]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
targetAlias, ok := nodeAlias[target]
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
tid := fmt.Sprintf("t%d", i)
|
||||
label := escapeMermaidLabel(fmt.Sprintf("%s (%s)", t.Name, t.Category))
|
||||
fmt.Fprintf(&b, " %s{{\"%s\"}}\n", tid, label)
|
||||
fmt.Fprintf(&b, " class %s nodeThreat\n", tid)
|
||||
fmt.Fprintf(&b, " %s -.-> %s\n", tid, targetAlias)
|
||||
}
|
||||
|
||||
b.WriteString(" classDef nodeEntity fill:#dbeafe,stroke:#1d4ed8,color:#1e3a8a\n")
|
||||
b.WriteString(" classDef nodeBoundary fill:#fef3c7,stroke:#b45309,color:#78350f\n")
|
||||
b.WriteString(" classDef nodeAsset fill:#e5e7eb,stroke:#374151,color:#111827\n")
|
||||
b.WriteString(" classDef nodeData fill:#dcfce7,stroke:#15803d,color:#14532d\n")
|
||||
b.WriteString(" classDef nodeThreat fill:#fee2e2,stroke:#b91c1c,color:#7f1d1d\n")
|
||||
|
||||
return strings.TrimRight(b.String(), "\n")
|
||||
}
|
||||
|
||||
func mermaidNodeShape(t coredata.RiskAssessmentNodeType, id, name string) string {
|
||||
label := `"` + escapeMermaidLabel(name) + `"`
|
||||
switch t {
|
||||
case coredata.RiskAssessmentNodeTypeEntity:
|
||||
return fmt.Sprintf("%s([%s])", id, label)
|
||||
case coredata.RiskAssessmentNodeTypeBoundary:
|
||||
return fmt.Sprintf("%s{{%s}}", id, label)
|
||||
case coredata.RiskAssessmentNodeTypeData:
|
||||
return fmt.Sprintf("%s[(%s)]", id, label)
|
||||
case coredata.RiskAssessmentNodeTypeAsset:
|
||||
fallthrough
|
||||
default:
|
||||
return fmt.Sprintf("%s[%s]", id, label)
|
||||
}
|
||||
}
|
||||
|
||||
func mermaidNodeClass(t coredata.RiskAssessmentNodeType) string {
|
||||
switch t {
|
||||
case coredata.RiskAssessmentNodeTypeEntity:
|
||||
return "nodeEntity"
|
||||
case coredata.RiskAssessmentNodeTypeBoundary:
|
||||
return "nodeBoundary"
|
||||
case coredata.RiskAssessmentNodeTypeData:
|
||||
return "nodeData"
|
||||
case coredata.RiskAssessmentNodeTypeAsset:
|
||||
fallthrough
|
||||
default:
|
||||
return "nodeAsset"
|
||||
}
|
||||
}
|
||||
|
||||
var mermaidLabelReplacer = strings.NewReplacer(
|
||||
"&", "&",
|
||||
`"`, "#quot;",
|
||||
"<", "<",
|
||||
">", ">",
|
||||
"\r\n", " ",
|
||||
"\n", " ",
|
||||
)
|
||||
|
||||
func escapeMermaidLabel(s string) string {
|
||||
return mermaidLabelReplacer.Replace(s)
|
||||
}
|
||||
@@ -228,6 +228,8 @@ type RiskAssessmentScope implements Node {
|
||||
orderBy: RiskAssessmentScenarioOrder
|
||||
): RiskAssessmentScenarioConnection @goField(forceResolver: true)
|
||||
|
||||
mermaidChart: String! @goField(forceResolver: true)
|
||||
|
||||
createdAt: Datetime!
|
||||
updatedAt: Datetime!
|
||||
}
|
||||
|
||||
@@ -854,6 +854,20 @@ func (r *riskAssessmentScopeResolver) Scenarios(ctx context.Context, obj *types.
|
||||
return types.NewRiskAssessmentScenarioConnection(p, r, obj.ID), nil
|
||||
}
|
||||
|
||||
// MermaidChart is the resolver for the mermaidChart field.
|
||||
func (r *riskAssessmentScopeResolver) MermaidChart(ctx context.Context, obj *types.RiskAssessmentScope) (string, error) {
|
||||
if err := r.authorize(ctx, obj.ID, probo.ActionRiskAssessmentScopeGet); err != nil {
|
||||
return "", err
|
||||
}
|
||||
scope := coredata.NewScopeFromObjectID(obj.ID)
|
||||
chart, err := r.riskManagement.BuildScopeMermaidChart(ctx, scope, obj.ID)
|
||||
if err != nil {
|
||||
r.logger.ErrorCtx(ctx, "cannot build risk assessment scope mermaid chart", log.Error(err))
|
||||
return "", gqlutils.Internal(ctx)
|
||||
}
|
||||
return chart, nil
|
||||
}
|
||||
|
||||
// TotalCount is the resolver for the totalCount field.
|
||||
func (r *riskAssessmentScopeConnectionResolver) TotalCount(ctx context.Context, obj *types.RiskAssessmentScopeConnection) (*int, error) {
|
||||
if err := r.authorize(ctx, obj.ParentID, probo.ActionRiskAssessmentScopeList); err != nil {
|
||||
|
||||
Reference in New Issue
Block a user