diff --git a/ci/release/changelogs/next.md b/ci/release/changelogs/next.md index 38cd1d2366..ffe01dbc1e 100644 --- a/ci/release/changelogs/next.md +++ b/ci/release/changelogs/next.md @@ -45,6 +45,9 @@ - Limit recursive glob matching and generated-field work so small diagrams fail with a clear error instead of consuming excessive CPU or memory. [#2925](https://github.com/d2lang/d2/pull/2925) + - Limit wildcard edge connections and selectors to bounded expansion work, + returning a clear error before excessive fanout consumes resources. + [#2926](https://github.com/d2lang/d2/pull/2926) - decoding and assets: - Cap decompressed URL-encoded D2 input at 16 MiB. [#2902](https://github.com/d2lang/d2/pull/2902) - Bound image references, locators, fetched and decoded bytes, cached data, and @@ -59,6 +62,9 @@ - Limit grids to 10,000 rows or columns and 1,000,000 total cells, returning a clear error for oversized dimensions instead of crashing or exhausting memory. [#2924](https://github.com/d2lang/d2/pull/2924) + - Limit Dagre and ELK inputs to 1,024 objects and 1,024 edges, and reject + overly interconnected graphs before entering non-cancellable layout. + [#2926](https://github.com/d2lang/d2/pull/2926) - links and paints: - Reject dangerous ordinary link schemes after decoding common obfuscation. [#2896](https://github.com/d2lang/d2/pull/2896) - Require gradient stop positions to be finite numbers or percentages. [#2905](https://github.com/d2lang/d2/pull/2905) diff --git a/d2compiler/compile.go b/d2compiler/compile.go index dde52d82d4..60fd8e1a5c 100644 --- a/d2compiler/compile.go +++ b/d2compiler/compile.go @@ -38,6 +38,13 @@ type CompileOptions struct { // materialization. Zero uses d2ir.DefaultMaxGlobExpansion. Explicit source // fields are not counted as materialization work. MaxGlobExpansion int64 + // MaxEdgeExpansion bounds distinct edge-segment and endpoint combinations + // considered by edge globs. Zero uses d2ir.DefaultMaxEdgeExpansion. Explicit + // edges do not consume this budget. + MaxEdgeExpansion int64 + // MaxEdgeExpansionWork bounds all endpoint-pair examinations performed by + // edge globs, including lazy replays. Zero uses the secure compiler default. + MaxEdgeExpansionWork int64 // FS is the file system used for resolving imports in the D2 text. Nil // disables imports. Callers that accept untrusted input should prefer a // filesystem constrained to the intended import root; lib/localfile provides @@ -62,6 +69,8 @@ func Compile(p string, r io.Reader, opts *CompileOptions) (*d2graph.Graph, *d2ta UTF16Pos: opts.UTF16Pos, MaxVariableExpansion: opts.MaxVariableExpansion, MaxGlobExpansion: opts.MaxGlobExpansion, + MaxEdgeExpansion: opts.MaxEdgeExpansion, + MaxEdgeExpansionWork: opts.MaxEdgeExpansionWork, FS: opts.FS, }) if err != nil { diff --git a/d2compiler/edge_expansion_test.go b/d2compiler/edge_expansion_test.go new file mode 100644 index 0000000000..cce6a5f7ad --- /dev/null +++ b/d2compiler/edge_expansion_test.go @@ -0,0 +1,63 @@ +package d2compiler_test + +import ( + "strings" + "testing" + + "github.com/d2lang/d2/d2compiler" +) + +func TestCompileEdgeExpansionLimit(t *testing.T) { + t.Parallel() + + _, _, err := d2compiler.Compile( + "edge-expansion.d2", + strings.NewReader("a\nb\nc\n* -> *\n"), + &d2compiler.CompileOptions{MaxEdgeExpansion: 8}, + ) + if err == nil || !strings.Contains(err.Error(), "edge glob expansion exceeds limit of 8 endpoint pairs") { + t.Fatalf("Compile error = %v, want edge expansion limit", err) + } +} + +func TestCompileEdgeSelectorExpansionLimit(t *testing.T) { + t.Parallel() + + _, _, err := d2compiler.Compile( + "edge-selector-expansion.d2", + strings.NewReader("a\nb\nc\n(* -> *)[*].style.opacity: 0\n"), + &d2compiler.CompileOptions{MaxEdgeExpansion: 8}, + ) + if err == nil || !strings.Contains(err.Error(), "edge glob expansion exceeds limit of 8 endpoint pairs") { + t.Fatalf("Compile error = %v, want edge expansion limit", err) + } +} + +func TestCompileEdgeExpansionWorkLimit(t *testing.T) { + t.Parallel() + + _, _, err := d2compiler.Compile( + "edge-expansion-work.d2", + strings.NewReader("(* -> *)[*].style.opacity: 0\na\nb\nc\n"), + &d2compiler.CompileOptions{ + MaxEdgeExpansion: 9, + MaxEdgeExpansionWork: 9, + }, + ) + if err == nil || !strings.Contains(err.Error(), "work limit of 9 endpoint-pair examinations") { + t.Fatalf("Compile error = %v, want edge expansion work limit", err) + } +} + +func TestCompileEdgeExpansionDoesNotChargeExplicitEdges(t *testing.T) { + t.Parallel() + + _, _, err := d2compiler.Compile( + "explicit-edges.d2", + strings.NewReader("a -> b\na -> b\na -> b\n"), + &d2compiler.CompileOptions{MaxEdgeExpansion: 1}, + ) + if err != nil { + t.Fatalf("Compile explicit edges: %v", err) + } +} diff --git a/d2ir/compile.go b/d2ir/compile.go index 391c73d3b8..9f718152ff 100644 --- a/d2ir/compile.go +++ b/d2ir/compile.go @@ -31,14 +31,17 @@ type globContext struct { } type compiler struct { - err *d2parser.ParseError - ctx context.Context - contextErr error - expansionErr error - globExpansionErr error - halted bool - variableExpansion *variableExpansionBudget - globExpansion *globExpansionBudget + err *d2parser.ParseError + ctx context.Context + contextErr error + expansionErr error + globExpansionErr error + halted bool + variableExpansion *variableExpansionBudget + globExpansion *globExpansionBudget + edgeExpansion *edgeExpansionBudget + edgeExpansionWork *edgeExpansionWorkBudget + edgeExpansionPairs map[edgeExpansionPair]struct{} fs fs.FS imports []string @@ -87,6 +90,13 @@ type CompileOptions struct { // materialization. Zero uses DefaultMaxGlobExpansion. Explicit source fields // are not counted as materialization work. MaxGlobExpansion int64 + // MaxEdgeExpansion bounds distinct edge-segment and endpoint combinations + // considered by edge globs. Zero uses DefaultMaxEdgeExpansion. Explicit edges + // do not consume this budget. + MaxEdgeExpansion int64 + // MaxEdgeExpansionWork bounds all endpoint-pair examinations performed by + // edge globs, including lazy replays. Zero uses DefaultMaxEdgeExpansionWork. + MaxEdgeExpansionWork int64 // FS resolves imports. Nil disables imports. The lib/localfile package // provides rooted and explicit unrestricted host-filesystem policies. FS fs.FS @@ -115,12 +125,22 @@ func Compile(ast *d2ast.Map, opts *CompileOptions) (*Map, []string, error) { if err != nil { return nil, nil, err } + edgeExpansion, err := newEdgeExpansionBudget(opts.MaxEdgeExpansion) + if err != nil { + return nil, nil, err + } + edgeExpansionWork, err := newEdgeExpansionWorkBudget(opts.MaxEdgeExpansionWork) + if err != nil { + return nil, nil, err + } c := &compiler{ err: &d2parser.ParseError{}, ctx: ctx, fs: opts.FS, variableExpansion: variableExpansion, globExpansion: globExpansion, + edgeExpansion: edgeExpansion, + edgeExpansionWork: edgeExpansionWork, seenImports: make(map[string]struct{}), parsedImports: make(map[string]*d2ast.Map), diff --git a/d2ir/d2ir.go b/d2ir/d2ir.go index de690cad34..2ed1eabdee 100644 --- a/d2ir/d2ir.go +++ b/d2ir/d2ir.go @@ -1563,7 +1563,7 @@ func (m *Map) getEdgesMode(eid *EdgeID, refctx *RefContext, c *compiler, indexed gctx = c.ensureGlobContext(refctx) } var ea []*Edge - m.getEdges(eid, refctx, gctx, indexed, &ea) + m.getEdges(eid, refctx, gctx, c, indexed, &ea) return ea } @@ -1590,7 +1590,10 @@ func (m *Map) getEdgesIndexed(eid *EdgeID) []*Edge { return edges } -func (m *Map) getEdges(eid *EdgeID, refctx *RefContext, gctx *globContext, indexed bool, ea *[]*Edge) error { +func (m *Map) getEdges(eid *EdgeID, refctx *RefContext, gctx *globContext, c *compiler, indexed bool, ea *[]*Edge) error { + if c != nil && c.stopped() { + return nil + } eid, m, common, err := eid.resolve(m) if err != nil { return err @@ -1608,11 +1611,14 @@ func (m *Map) getEdges(eid *EdgeID, refctx *RefContext, gctx *globContext, index } } } - fa, err := m.ensureFieldMode(commonKP, nil, false, nil, indexed) + fa, err := m.ensureFieldMode(commonKP, nil, false, c, indexed) if err != nil { return nil } for _, f := range fa { + if c != nil && c.stopped() { + return nil + } if _, ok := f.Composite.(*Array); ok { return d2parser.Errorf(refctx.Edge.Src, "cannot index into array") } @@ -1621,7 +1627,7 @@ func (m *Map) getEdges(eid *EdgeID, refctx *RefContext, gctx *globContext, index parent: f, } } - err = f.Map().getEdges(eid, refctx, gctx, indexed, ea) + err = f.Map().getEdges(eid, refctx, gctx, c, indexed, ea) if err != nil { return err } @@ -1629,17 +1635,23 @@ func (m *Map) getEdges(eid *EdgeID, refctx *RefContext, gctx *globContext, index return nil } - srcFA, err := refctx.ScopeMap.ensureFieldMode(refctx.Edge.Src, nil, false, nil, indexed) + srcFA, err := refctx.ScopeMap.ensureFieldMode(refctx.Edge.Src, nil, false, c, indexed) if err != nil { return err } - dstFA, err := refctx.ScopeMap.ensureFieldMode(refctx.Edge.Dst, nil, false, nil, indexed) + dstFA, err := refctx.ScopeMap.ensureFieldMode(refctx.Edge.Dst, nil, false, c, indexed) if err != nil { return err } for _, src := range srcFA { for _, dst := range dstFA { + if c != nil && c.stopped() { + return nil + } + if c != nil && (refctx.Edge.Src.HasGlob() || refctx.Edge.Dst.HasGlob()) && !c.reserveEdgeExpansion(refctx.Edge, gctx, src, dst, true) { + return nil + } eid2 := eid.Copy() eid2.SrcPath = RelIDA(m, src) eid2.DstPath = RelIDA(m, dst) @@ -1782,6 +1794,9 @@ func (m *Map) createEdge(eid *EdgeID, refctx *RefContext, gctx *globContext, c * if c != nil && c.stopped() { return nil } + if c != nil && (refctx.Edge.Src.HasGlob() || refctx.Edge.Dst.HasGlob()) && !c.reserveEdgeExpansion(refctx.Edge, gctx, src, dst, false) { + return nil + } if src == dst && (refctx.Edge.Src.HasGlob() || refctx.Edge.Dst.HasGlob()) { // Globs do not make self edges. continue diff --git a/d2ir/edge_expansion.go b/d2ir/edge_expansion.go new file mode 100644 index 0000000000..f0d5ddeed8 --- /dev/null +++ b/d2ir/edge_expansion.go @@ -0,0 +1,127 @@ +package d2ir + +import ( + "fmt" + + "github.com/d2lang/d2/d2ast" +) + +// DefaultMaxEdgeExpansion is the maximum number of distinct edge-segment and +// endpoint combinations that edge globs may consider during one compilation. +// Explicit edges do not consume this budget because their work is proportional +// to the input size. The limit is calibrated so a sparse star at the boundary +// remains bounded in the built-in Dagre and ELK layouts, while larger wildcard +// fanout stops before layout. +const DefaultMaxEdgeExpansion int64 = 1_024 + +// DefaultMaxEdgeExpansionWork is the maximum number of endpoint-pair +// examinations edge globs may perform, including repeated lazy replays. +const DefaultMaxEdgeExpansionWork int64 = 65_536 + +type edgeExpansionPair struct { + glob *globContext + segment *d2ast.Edge + src *Field + dst *Field + selector bool +} + +type edgeExpansionBudget struct { + limit int64 + used int64 +} + +type edgeExpansionLimitError struct { + limit int64 +} + +type edgeExpansionWorkBudget struct { + limit int64 + used int64 +} + +type edgeExpansionWorkLimitError struct { + limit int64 +} + +func (e *edgeExpansionWorkLimitError) Error() string { + return fmt.Sprintf("edge glob expansion exceeds work limit of %d endpoint-pair examinations", e.limit) +} + +func (e *edgeExpansionLimitError) Error() string { + return fmt.Sprintf("edge glob expansion exceeds limit of %d endpoint pairs", e.limit) +} + +func newEdgeExpansionBudget(limit int64) (*edgeExpansionBudget, error) { + if limit < 0 { + return nil, fmt.Errorf("MaxEdgeExpansion must not be negative") + } + if limit == 0 { + limit = DefaultMaxEdgeExpansion + } + return &edgeExpansionBudget{limit: limit}, nil +} + +func newEdgeExpansionWorkBudget(limit int64) (*edgeExpansionWorkBudget, error) { + if limit < 0 { + return nil, fmt.Errorf("MaxEdgeExpansionWork must not be negative") + } + if limit == 0 { + limit = DefaultMaxEdgeExpansionWork + } + return &edgeExpansionWorkBudget{limit: limit}, nil +} + +func (b *edgeExpansionBudget) reserve() error { + if b == nil || b.used >= b.limit { + limit := int64(0) + if b != nil { + limit = b.limit + } + return &edgeExpansionLimitError{limit: limit} + } + b.used++ + return nil +} + +func (b *edgeExpansionWorkBudget) reserve() error { + if b == nil || b.used >= b.limit { + limit := int64(0) + if b != nil { + limit = b.limit + } + return &edgeExpansionWorkLimitError{limit: limit} + } + b.used++ + return nil +} + +func (c *compiler) reserveEdgeExpansion(segment *d2ast.Edge, glob *globContext, src, dst *Field, selector bool) bool { + if c.stopped() { + return false + } + if err := c.edgeExpansionWork.reserve(); err != nil { + return c.rejectEdgeExpansion(segment, err) + } + pair := edgeExpansionPair{glob: glob, segment: segment, src: src, dst: dst, selector: selector} + if _, ok := c.edgeExpansionPairs[pair]; ok { + return true + } + if err := c.edgeExpansion.reserve(); err != nil { + return c.rejectEdgeExpansion(segment, err) + } + if c.edgeExpansionPairs == nil { + c.edgeExpansionPairs = make(map[edgeExpansionPair]struct{}) + } + c.edgeExpansionPairs[pair] = struct{}{} + return true +} + +func (c *compiler) rejectEdgeExpansion(n d2ast.Node, err error) bool { + c.expansionErr = err + if n != nil { + c.errorf(n, "%v", err) + } + c.halted = true + return false +} diff --git a/d2ir/edge_expansion_test.go b/d2ir/edge_expansion_test.go new file mode 100644 index 0000000000..d115f273a3 --- /dev/null +++ b/d2ir/edge_expansion_test.go @@ -0,0 +1,264 @@ +package d2ir + +import ( + "context" + "errors" + "fmt" + "strings" + "testing" + + "github.com/d2lang/d2/d2parser" +) + +func TestEdgeExpansionBudget(t *testing.T) { + t.Parallel() + + budget, err := newEdgeExpansionBudget(0) + if err != nil { + t.Fatal(err) + } + if budget.limit != DefaultMaxEdgeExpansion { + t.Fatalf("zero-value limit = %d, want %d", budget.limit, DefaultMaxEdgeExpansion) + } + for i := int64(0); i < DefaultMaxEdgeExpansion; i++ { + if err := budget.reserve(); err != nil { + t.Fatalf("reserve %d: %v", i, err) + } + } + if err := budget.reserve(); err == nil || !strings.Contains(err.Error(), fmt.Sprintf("limit of %d endpoint pairs", DefaultMaxEdgeExpansion)) { + t.Fatalf("overflow reserve error = %v, want edge expansion limit", err) + } + if budget.used != DefaultMaxEdgeExpansion { + t.Fatalf("failed reservation changed used budget to %d", budget.used) + } + if _, err := newEdgeExpansionBudget(-1); err == nil || !strings.Contains(err.Error(), "MaxEdgeExpansion must not be negative") { + t.Fatalf("negative limit error = %v", err) + } +} + +func TestEdgeExpansionWorkBudget(t *testing.T) { + t.Parallel() + + budget, err := newEdgeExpansionWorkBudget(0) + if err != nil { + t.Fatal(err) + } + if budget.limit != DefaultMaxEdgeExpansionWork { + t.Fatalf("zero-value limit = %d, want %d", budget.limit, DefaultMaxEdgeExpansionWork) + } + + budget, err = newEdgeExpansionWorkBudget(3) + if err != nil { + t.Fatal(err) + } + for i := 0; i < 3; i++ { + if err := budget.reserve(); err != nil { + t.Fatalf("reserve %d: %v", i, err) + } + } + if err := budget.reserve(); err == nil || !strings.Contains(err.Error(), "work limit of 3 endpoint-pair examinations") { + t.Fatalf("overflow reserve error = %v, want edge expansion work limit", err) + } + if budget.used != 3 { + t.Fatalf("failed reservation changed used budget to %d", budget.used) + } + if _, err := newEdgeExpansionWorkBudget(-1); err == nil || !strings.Contains(err.Error(), "MaxEdgeExpansionWork must not be negative") { + t.Fatalf("negative limit error = %v", err) + } +} + +func TestEdgeGlobExpansionBoundary(t *testing.T) { + t.Parallel() + + const source = "a\nb\nc\n* -> *\n" + if got, err := compileEdgeExpansion(t, source, 9, context.Background()); err != nil { + t.Fatalf("boundary compile: %v", err) + } else if got := got.EdgeCountRecursive(); got != 6 { + t.Fatalf("edge count = %d, want 6", got) + } + + if _, err := compileEdgeExpansion(t, source, 8, context.Background()); err == nil || !strings.Contains(err.Error(), "edge glob expansion exceeds limit of 8 endpoint pairs") { + t.Fatalf("over-limit compile error = %v, want edge expansion limit", err) + } +} + +func TestEdgeSelectorGlobExpansionBoundary(t *testing.T) { + t.Parallel() + + const source = "a\nb\nc\n(* -> *)[*].style.opacity: 0\n" + if got, err := compileEdgeExpansion(t, source, 9, context.Background()); err != nil { + t.Fatalf("boundary compile: %v", err) + } else if got := got.EdgeCountRecursive(); got != 0 { + t.Fatalf("edge count = %d, want 0", got) + } + + if _, err := compileEdgeExpansion(t, source, 8, context.Background()); err == nil || !strings.Contains(err.Error(), "edge glob expansion exceeds limit of 8 endpoint pairs") { + t.Fatalf("over-limit compile error = %v, want edge expansion limit", err) + } +} + +func TestEdgeSelectorGlobExpansionDoesNotRechargeUniqueFanout(t *testing.T) { + t.Parallel() + + const source = "(* -> *)[*].style.opacity: 0\na\nb\nc\n" + if _, err := compileEdgeExpansion(t, source, 9, context.Background()); err != nil { + t.Fatalf("lazy selector compile: %v", err) + } +} + +func TestEdgeSelectorGlobExpansionChargesLazyReplayWork(t *testing.T) { + t.Parallel() + + const source = "(* -> *)[*].style.opacity: 0\na\nb\nc\n" + if _, err := compileEdgeExpansionWithWork(t, source, 9, 9, context.Background()); err == nil || !strings.Contains(err.Error(), "work limit of 9 endpoint-pair examinations") { + t.Fatalf("lazy selector compile error = %v, want edge expansion work limit", err) + } +} + +func TestEdgeGlobExpansionChargesEachChainSegment(t *testing.T) { + t.Parallel() + + var source strings.Builder + for i := 0; i < 32; i++ { + fmt.Fprintf(&source, "n%d\n", i) + } + source.WriteString("* -> * -> *\n") + + if _, err := compileEdgeExpansion(t, source.String(), DefaultMaxEdgeExpansion, context.Background()); err == nil || !strings.Contains(err.Error(), fmt.Sprintf("edge glob expansion exceeds limit of %d endpoint pairs", DefaultMaxEdgeExpansion)) { + t.Fatalf("multi-segment chain compile error = %v, want edge expansion limit", err) + } +} + +func TestEdgeGlobExpansionDefaultRejectsCompleteGraphFanout(t *testing.T) { + t.Parallel() + + var source strings.Builder + for i := 0; i < 80; i++ { + fmt.Fprintf(&source, "n%d\n", i) + } + source.WriteString("* -> *\n") + + if _, err := compileEdgeExpansion(t, source.String(), 0, context.Background()); err == nil || !strings.Contains(err.Error(), fmt.Sprintf("edge glob expansion exceeds limit of %d endpoint pairs", DefaultMaxEdgeExpansion)) { + t.Fatalf("complete graph compile error = %v, want default edge expansion limit", err) + } +} + +func TestEdgeSelectorGlobExpansionDefaultRejectsFanout(t *testing.T) { + t.Parallel() + + var source strings.Builder + for i := 0; i < 128; i++ { + fmt.Fprintf(&source, "n%d\n", i) + } + source.WriteString("(* -> *)[*].style.opacity: 0\n") + + if _, err := compileEdgeExpansion(t, source.String(), 0, context.Background()); err == nil || !strings.Contains(err.Error(), fmt.Sprintf("edge glob expansion exceeds limit of %d endpoint pairs", DefaultMaxEdgeExpansion)) { + t.Fatalf("edge selector compile error = %v, want default edge expansion limit", err) + } +} + +func TestEdgeGlobExpansionDefaultSparseStarBoundary(t *testing.T) { + t.Parallel() + + compileStar := func(leaves int) (*Map, error) { + var source strings.Builder + source.WriteString("root\n") + for i := 0; i < leaves; i++ { + fmt.Fprintf(&source, "leaf%d\n", i) + } + source.WriteString("root -> *\n") + return compileEdgeExpansion(t, source.String(), 0, context.Background()) + } + + // The root plus DefaultMaxEdgeExpansion-1 leaves examines exactly the + // default number of endpoint pairs. + m, err := compileStar(int(DefaultMaxEdgeExpansion - 1)) + if err != nil { + t.Fatalf("default boundary compile: %v", err) + } + if got := m.EdgeCountRecursive(); got != int(DefaultMaxEdgeExpansion-1) { + t.Fatalf("edge count = %d, want %d", got, DefaultMaxEdgeExpansion-1) + } + + if _, err := compileStar(int(DefaultMaxEdgeExpansion)); err == nil || !strings.Contains(err.Error(), fmt.Sprintf("edge glob expansion exceeds limit of %d endpoint pairs", DefaultMaxEdgeExpansion)) { + t.Fatalf("over-limit sparse star error = %v, want default edge expansion limit", err) + } +} + +func TestEdgeGlobExpansionObservesCancellation(t *testing.T) { + t.Parallel() + + base := compileIndexTestSource(t, "a\nb\nc\n") + key, err := d2parser.ParseMapKey("* -> *") + if err != nil { + t.Fatal(err) + } + + // Locate a deterministic cancellation point after expansion has started but + // before all nine endpoint pairs have been examined. This avoids depending + // on the number of context checks performed by field lookup internals. + for cancelAt := 1; cancelAt < 128; cancelAt++ { + root := base.Copy(nil).(*Map) + ctx := &cancelAfterErrContext{Context: context.Background(), cancelAt: cancelAt} + c := &compiler{ + err: &d2parser.ParseError{}, + ctx: ctx, + variableExpansion: &variableExpansionBudget{limit: DefaultMaxVariableExpansion}, + edgeExpansion: &edgeExpansionBudget{limit: 100}, + globContextStack: [][]*globContext{{}}, + } + refctx := &RefContext{Key: key, Edge: key.Edges[0], ScopeMap: root} + _, _ = root.CreateEdge(NewEdgeIDs(key)[0], refctx, c) + if errors.Is(c.contextErr, context.Canceled) && c.edgeExpansion.used > 0 && c.edgeExpansion.used < 9 { + return + } + } + t.Fatal("could not observe cancellation during endpoint-pair expansion") +} + +func TestEdgeSelectorGlobExpansionObservesCancellation(t *testing.T) { + t.Parallel() + + base := compileIndexTestSource(t, "a\nb\nc\n") + key, err := d2parser.ParseMapKey("(* -> *)[*].style.opacity") + if err != nil { + t.Fatal(err) + } + + for cancelAt := 1; cancelAt < 128; cancelAt++ { + root := base.Copy(nil).(*Map) + ctx := &cancelAfterErrContext{Context: context.Background(), cancelAt: cancelAt} + c := &compiler{ + err: &d2parser.ParseError{}, + ctx: ctx, + variableExpansion: &variableExpansionBudget{limit: DefaultMaxVariableExpansion}, + edgeExpansion: &edgeExpansionBudget{limit: 100}, + globContextStack: [][]*globContext{{}}, + } + refctx := &RefContext{Key: key, Edge: key.Edges[0], ScopeMap: root} + _ = root.getEdgesForCompile(NewEdgeIDs(key)[0], refctx, c) + if errors.Is(c.contextErr, context.Canceled) && c.edgeExpansion.used > 0 && c.edgeExpansion.used < 9 { + return + } + } + t.Fatal("could not observe cancellation during selector endpoint-pair expansion") +} + +func compileEdgeExpansion(t *testing.T, source string, limit int64, ctx context.Context) (*Map, error) { + t.Helper() + return compileEdgeExpansionWithWork(t, source, limit, 0, ctx) +} + +func compileEdgeExpansionWithWork(t *testing.T, source string, limit, workLimit int64, ctx context.Context) (*Map, error) { + t.Helper() + ast, err := d2parser.Parse("edge-expansion.d2", strings.NewReader(source), nil) + if err != nil { + t.Fatal(err) + } + m, _, err := Compile(ast, &CompileOptions{ + Context: ctx, + MaxEdgeExpansion: limit, + MaxEdgeExpansionWork: workLimit, + }) + return m, err +} diff --git a/d2ir/expansion.go b/d2ir/expansion.go index eeecc89777..65f58e043a 100644 --- a/d2ir/expansion.go +++ b/d2ir/expansion.go @@ -66,6 +66,12 @@ func (c *compiler) stopped() bool { if c.globExpansion == nil { c.globExpansion = &globExpansionBudget{limit: DefaultMaxGlobExpansion} } + if c.edgeExpansion == nil { + c.edgeExpansion = &edgeExpansionBudget{limit: DefaultMaxEdgeExpansion} + } + if c.edgeExpansionWork == nil { + c.edgeExpansionWork = &edgeExpansionWorkBudget{limit: DefaultMaxEdgeExpansionWork} + } if err := c.ctx.Err(); err != nil { c.contextErr = err c.halted = true diff --git a/d2layouts/d2dagrelayout/layout.go b/d2layouts/d2dagrelayout/layout.go index 4208e1806d..7eb59b2e26 100644 --- a/d2layouts/d2dagrelayout/layout.go +++ b/d2layouts/d2dagrelayout/layout.go @@ -12,6 +12,7 @@ import ( "github.com/d2lang/util-go/go2" "github.com/d2lang/d2/d2graph" + "github.com/d2lang/d2/d2layouts/internal/layoutguard" "github.com/d2lang/d2/d2target" "github.com/d2lang/d2/lib/geo" "github.com/d2lang/d2/lib/label" @@ -50,10 +51,16 @@ func DefaultLayout(ctx context.Context, g *d2graph.Graph) (err error) { } func Layout(ctx context.Context, g *d2graph.Graph, opts *ConfigurableOpts) (err error) { + if ctx == nil { + ctx = context.Background() + } if opts == nil { opts = &DefaultOpts } defer xdefer.Errorf(&err, "failed to dagre layout") + if err := layoutguard.CheckGraph(ctx, "dagre", g); err != nil { + return err + } rootAttrs := dagreOpts{ ConfigurableOpts: ConfigurableOpts{ @@ -173,9 +180,17 @@ func Layout(ctx context.Context, g *d2graph.Graph, opts *ConfigurableOpts) (err ) } + if err := ctx.Err(); err != nil { + return err + } + // dagro does not accept a context. Check immediately around the call; the + // topology preflight above bounds inputs before entering this section. if err := dagro.Layout(dagreGraph); err != nil { return err } + if err := ctx.Err(); err != nil { + return err + } for _, id := range dagreGraph.Nodes() { dn, ok := dagreGraph.Node(id).(dagro.Attrs) diff --git a/d2layouts/d2dagrelayout/layout_test.go b/d2layouts/d2dagrelayout/layout_test.go index aebf813f91..7ea7b5ef4d 100644 --- a/d2layouts/d2dagrelayout/layout_test.go +++ b/d2layouts/d2dagrelayout/layout_test.go @@ -2,6 +2,7 @@ package d2dagrelayout import ( "context" + "errors" "math" "strings" "testing" @@ -14,6 +15,34 @@ import ( "github.com/d2lang/util-go/go2" ) +func TestLayoutRejectsDenseGraphBeforeDagre(t *testing.T) { + t.Parallel() + g := d2graph.NewGraph() + for i := 0; i < 20; i++ { + g.Objects = append(g.Objects, &d2graph.Object{Graph: g, Parent: g.Root}) + } + for src := range g.Objects { + for dst := range g.Objects { + if src != dst { + g.Edges = append(g.Edges, &d2graph.Edge{Src: g.Objects[src], Dst: g.Objects[dst]}) + } + } + } + err := DefaultLayout(context.Background(), g) + if err == nil || !strings.Contains(err.Error(), "layout graph is too interconnected") { + t.Fatalf("DefaultLayout error = %v, want interaction-work limit", err) + } +} + +func TestLayoutChecksCancellationBeforeDagre(t *testing.T) { + t.Parallel() + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if err := DefaultLayout(ctx, &d2graph.Graph{}); !errors.Is(err, context.Canceled) { + t.Fatalf("DefaultLayout error = %v, want context.Canceled", err) + } +} + func TestDeduplicateRoutePoints(t *testing.T) { t.Parallel() points := []*geo.Point{ diff --git a/d2layouts/d2elklayout/layout.go b/d2layouts/d2elklayout/layout.go index 464dac05dd..e0dd619fa2 100644 --- a/d2layouts/d2elklayout/layout.go +++ b/d2layouts/d2elklayout/layout.go @@ -19,6 +19,7 @@ import ( "github.com/d2lang/util-go/go2" "github.com/d2lang/d2/d2graph" + "github.com/d2lang/d2/d2layouts/internal/layoutguard" "github.com/d2lang/d2/d2target" "github.com/d2lang/d2/lib/geo" "github.com/d2lang/d2/lib/label" @@ -242,11 +243,14 @@ func DefaultLayout(ctx context.Context, g *d2graph.Graph) (err error) { } func Layout(ctx context.Context, g *d2graph.Graph, opts *ConfigurableOpts) (err error) { + if ctx == nil { + ctx = context.Background() + } if opts == nil { opts = &DefaultOpts } defer xdefer.Errorf(&err, "failed to ELK layout") - if err := ctx.Err(); err != nil { + if err := layoutguard.CheckGraph(ctx, "ELK", g); err != nil { return err } graphStats := collectELKGraphStats(g) @@ -478,10 +482,18 @@ func Layout(ctx context.Context, g *d2graph.Graph, opts *ConfigurableOpts) (err return err } + if err := ctx.Err(); err != nil { + return err + } + // ELK does not accept a context. Check immediately around the call; the + // topology preflight above bounds inputs before entering this section. jsonOut, err := elk.LayoutJSON(raw) if err != nil { return fmt.Errorf("native ELK layout failed: %w", err) } + if err := ctx.Err(); err != nil { + return err + } err = json.Unmarshal(jsonOut, &elkGraph) if err != nil { diff --git a/d2layouts/d2elklayout/security_test.go b/d2layouts/d2elklayout/security_test.go new file mode 100644 index 0000000000..00de9b7df6 --- /dev/null +++ b/d2layouts/d2elklayout/security_test.go @@ -0,0 +1,38 @@ +package d2elklayout + +import ( + "context" + "errors" + "strings" + "testing" + + "github.com/d2lang/d2/d2graph" +) + +func TestLayoutRejectsDenseGraphBeforeELK(t *testing.T) { + t.Parallel() + g := d2graph.NewGraph() + for i := 0; i < 20; i++ { + g.Objects = append(g.Objects, &d2graph.Object{Graph: g, Parent: g.Root}) + } + for src := range g.Objects { + for dst := range g.Objects { + if src != dst { + g.Edges = append(g.Edges, &d2graph.Edge{Src: g.Objects[src], Dst: g.Objects[dst]}) + } + } + } + err := DefaultLayout(context.Background(), g) + if err == nil || !strings.Contains(err.Error(), "layout graph is too interconnected") { + t.Fatalf("DefaultLayout error = %v, want interaction-work limit", err) + } +} + +func TestLayoutChecksCancellationBeforeELK(t *testing.T) { + t.Parallel() + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if err := DefaultLayout(ctx, &d2graph.Graph{}); !errors.Is(err, context.Canceled) { + t.Fatalf("DefaultLayout error = %v, want context.Canceled", err) + } +} diff --git a/d2layouts/internal/layoutguard/graph.go b/d2layouts/internal/layoutguard/graph.go new file mode 100644 index 0000000000..43d3c1e62e --- /dev/null +++ b/d2layouts/internal/layoutguard/graph.go @@ -0,0 +1,86 @@ +// Package layoutguard protects layout engines that cannot be interrupted once +// their native layout call begins. +package layoutguard + +import ( + "context" + "fmt" + + "github.com/d2lang/d2/d2graph" +) + +const ( + // Dagre and ELK do not accept a context during their native layout calls. + // Keep their raw inputs at or below the compiler's calibrated wildcard + // boundary so direct graph callers cannot bypass compiler admission. + maxLayoutObjects = 1_024 + maxLayoutEdges = 1_024 + // Small graphs are safe to pass through without topology accounting and may + // intentionally contain many parallel or self-referential edges. + unconditionallyAllowedEdges = 256 + // Edge interaction work is the sum, for every edge, of the smaller incident + // degree of its endpoints. It remains low for stars, trees, and other sparse + // diagrams, but grows quickly for dense subgraphs and cannot be diluted by + // adding unrelated objects. + maxEdgeInteractionWork int64 = 8_192 +) + +// CheckGraph rejects malformed, oversized, or dense graphs before passing +// control to a layout engine that cannot observe context cancellation during +// its native layout call. +func CheckGraph(ctx context.Context, engine string, g *d2graph.Graph) error { + if ctx == nil { + ctx = context.Background() + } + if err := ctx.Err(); err != nil { + return err + } + if g == nil { + return fmt.Errorf("%s layout requires a graph", engine) + } + if len(g.Objects) > maxLayoutObjects { + return fmt.Errorf("%s layout object count %d exceeds safe limit of %d", engine, len(g.Objects), maxLayoutObjects) + } + if len(g.Edges) > maxLayoutEdges { + return fmt.Errorf("%s layout edge count %d exceeds safe limit of %d", engine, len(g.Edges), maxLayoutEdges) + } + + incident := make(map[*d2graph.Object]int, len(g.Objects)) + for i, edge := range g.Edges { + if i%64 == 0 { + if err := ctx.Err(); err != nil { + return err + } + } + if edge == nil || edge.Src == nil || edge.Dst == nil { + return fmt.Errorf("%s layout graph contains an edge without endpoints", engine) + } + incident[edge.Src]++ + incident[edge.Dst]++ + } + if len(g.Edges) <= unconditionallyAllowedEdges { + return nil + } + + var work int64 + for i, edge := range g.Edges { + if i%64 == 0 { + if err := ctx.Err(); err != nil { + return err + } + } + units := incident[edge.Src] + if incident[edge.Dst] < units { + units = incident[edge.Dst] + } + if int64(units) > maxEdgeInteractionWork-work { + return fmt.Errorf( + "%s layout graph is too interconnected: edge interaction work exceeds safe limit of %d", + engine, + maxEdgeInteractionWork, + ) + } + work += int64(units) + } + return nil +} diff --git a/d2layouts/internal/layoutguard/graph_test.go b/d2layouts/internal/layoutguard/graph_test.go new file mode 100644 index 0000000000..ac9937bb7d --- /dev/null +++ b/d2layouts/internal/layoutguard/graph_test.go @@ -0,0 +1,134 @@ +package layoutguard + +import ( + "context" + "errors" + "strings" + "testing" + + "github.com/d2lang/d2/d2graph" +) + +func TestCheckGraphPreservesSparseGraphs(t *testing.T) { + t.Parallel() + g := testGraph(maxLayoutObjects) + for i := 1; i < len(g.Objects); i++ { + g.Edges = append(g.Edges, &d2graph.Edge{Src: g.Objects[0], Dst: g.Objects[i]}) + } + if err := CheckGraph(context.Background(), "test", g); err != nil { + t.Fatal(err) + } +} + +func TestCheckGraphRawObjectBoundary(t *testing.T) { + t.Parallel() + if err := CheckGraph(context.Background(), "test", testGraph(maxLayoutObjects)); err != nil { + t.Fatalf("object boundary: %v", err) + } + if err := CheckGraph(context.Background(), "test", testGraph(maxLayoutObjects+1)); err == nil || !strings.Contains(err.Error(), "object count 1025 exceeds safe limit of 1024") { + t.Fatalf("over-limit object error = %v", err) + } +} + +func TestCheckGraphRawEdgeBoundary(t *testing.T) { + t.Parallel() + g := testGraph(512) + for i := 0; i < maxLayoutEdges; i++ { + g.Edges = append(g.Edges, &d2graph.Edge{ + Src: g.Objects[i%len(g.Objects)], + Dst: g.Objects[(i+1)%len(g.Objects)], + }) + } + if err := CheckGraph(context.Background(), "test", g); err != nil { + t.Fatalf("edge boundary: %v", err) + } + g.Edges = append(g.Edges, &d2graph.Edge{Src: g.Objects[0], Dst: g.Objects[1]}) + if err := CheckGraph(context.Background(), "test", g); err == nil || !strings.Contains(err.Error(), "edge count 1025 exceeds safe limit of 1024") { + t.Fatalf("over-limit edge error = %v", err) + } +} + +func TestCheckGraphRejectsDenseSubgraphs(t *testing.T) { + t.Parallel() + g := completeTestGraph(20) + // Unrelated objects must not dilute the dense component's work estimate. + for i := 0; i < 1_000; i++ { + g.Objects = append(g.Objects, &d2graph.Object{}) + } + err := CheckGraph(context.Background(), "test", g) + if err == nil || !strings.Contains(err.Error(), "edge interaction work exceeds safe limit of 8192") { + t.Fatalf("CheckGraph error = %v, want interaction-work limit", err) + } +} + +func TestCheckGraphDenseBoundary(t *testing.T) { + t.Parallel() + if err := CheckGraph(context.Background(), "test", completeTestGraph(16)); err != nil { + t.Fatalf("16-object complete graph: %v", err) + } + if err := CheckGraph(context.Background(), "test", completeTestGraph(17)); err == nil { + t.Fatal("17-object complete graph was not rejected") + } +} + +func TestCheckGraphPreservesSmallParallelGraphs(t *testing.T) { + t.Parallel() + g := testGraph(2) + for i := 0; i < unconditionallyAllowedEdges; i++ { + g.Edges = append(g.Edges, &d2graph.Edge{Src: g.Objects[0], Dst: g.Objects[1]}) + } + if err := CheckGraph(context.Background(), "test", g); err != nil { + t.Fatal(err) + } +} + +func TestCheckGraphRejectsMalformedSmallGraph(t *testing.T) { + t.Parallel() + + tests := map[string]*d2graph.Edge{ + "nil edge": nil, + "nil source": {Dst: &d2graph.Object{}}, + "nil target": {Src: &d2graph.Object{}}, + "nil endpoints": {}, + } + for name, edge := range tests { + t.Run(name, func(t *testing.T) { + t.Parallel() + g := d2graph.NewGraph() + g.Edges = append(g.Edges, edge) + if err := CheckGraph(context.Background(), "test", g); err == nil || !strings.Contains(err.Error(), "edge without endpoints") { + t.Fatalf("CheckGraph error = %v, want missing-endpoint error", err) + } + }) + } +} + +func TestCheckGraphObservesCancellation(t *testing.T) { + t.Parallel() + ctx, cancel := context.WithCancel(context.Background()) + cancel() + if err := CheckGraph(ctx, "test", &d2graph.Graph{}); !errors.Is(err, context.Canceled) { + t.Fatalf("CheckGraph error = %v, want context.Canceled", err) + } +} + +func completeTestGraph(objectCount int) *d2graph.Graph { + g := testGraph(objectCount) + for src := range g.Objects { + for dst := range g.Objects { + if src == dst { + continue + } + g.Edges = append(g.Edges, &d2graph.Edge{Src: g.Objects[src], Dst: g.Objects[dst]}) + } + } + return g +} + +func testGraph(objectCount int) *d2graph.Graph { + g := d2graph.NewGraph() + for i := 0; i < objectCount; i++ { + g.Objects = append(g.Objects, &d2graph.Object{Graph: g, Parent: g.Root}) + } + return g +} diff --git a/d2lib/d2.go b/d2lib/d2.go index 1153f0c8f9..28c87ad9ca 100644 --- a/d2lib/d2.go +++ b/d2lib/d2.go @@ -30,6 +30,13 @@ type CompileOptions struct { // materialization. Zero uses the secure compiler default. Explicit source // fields are not counted as materialization work. MaxGlobExpansion int64 + // MaxEdgeExpansion bounds distinct edge-segment and endpoint combinations + // considered by edge globs. Zero uses the secure compiler default. Explicit + // edges do not consume this budget. + MaxEdgeExpansion int64 + // MaxEdgeExpansionWork bounds all endpoint-pair examinations performed by + // edge globs, including lazy replays. Zero uses the secure compiler default. + MaxEdgeExpansionWork int64 // FS is the file system used for resolving imports in the D2 text. Nil // disables imports. Callers that accept untrusted input should prefer a // filesystem constrained to the intended import root; lib/localfile provides @@ -90,6 +97,8 @@ func compileInput(ctx context.Context, input string, compileOpts *CompileOptions UTF16Pos: compileOpts.UTF16Pos, MaxVariableExpansion: compileOpts.MaxVariableExpansion, MaxGlobExpansion: compileOpts.MaxGlobExpansion, + MaxEdgeExpansion: compileOpts.MaxEdgeExpansion, + MaxEdgeExpansionWork: compileOpts.MaxEdgeExpansionWork, FS: compileOpts.FS, }) if err != nil { diff --git a/d2lib/d2_test.go b/d2lib/d2_test.go index 285a047cfe..4e3b061aca 100644 --- a/d2lib/d2_test.go +++ b/d2lib/d2_test.go @@ -46,6 +46,25 @@ func TestCompilePropagatesGlobExpansionLimit(t *testing.T) { } } +func TestCompilePropagatesEdgeExpansionLimit(t *testing.T) { + _, _, err := Compile(context.Background(), "a\nb\nc\n* -> *\n", &CompileOptions{ + MaxEdgeExpansion: 8, + }, nil) + if err == nil || !strings.Contains(err.Error(), "edge glob expansion exceeds limit of 8 endpoint pairs") { + t.Fatalf("Compile() error = %v, want edge expansion limit", err) + } +} + +func TestCompilePropagatesEdgeExpansionWorkLimit(t *testing.T) { + _, _, err := Compile(context.Background(), "(* -> *)[*].style.opacity: 0\na\nb\nc\n", &CompileOptions{ + MaxEdgeExpansion: 9, + MaxEdgeExpansionWork: 9, + }, nil) + if err == nil || !strings.Contains(err.Error(), "work limit of 9 endpoint-pair examinations") { + t.Fatalf("Compile() error = %v, want edge expansion work limit", err) + } +} + func TestNilFSDeniesImports(t *testing.T) { directory := t.TempDir() if err := os.WriteFile(filepath.Join(directory, "secret.d2"), []byte("disclosed"), 0o600); err != nil {