Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 7 additions & 0 deletions d2graph/cyclediagram.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,7 @@
package d2graph

import "oss.terrastruct.com/d2/d2target"

func (obj *Object) IsCycleDiagram() bool {
return obj != nil && obj.Shape.Value == d2target.ShapeCycleDiagram
}
213 changes: 213 additions & 0 deletions d2layouts/d2cycle/layout.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,213 @@
package d2cycle

import (
"context"
"math"

"oss.terrastruct.com/d2/d2graph"
"oss.terrastruct.com/d2/lib/geo"
"oss.terrastruct.com/d2/lib/label"
"oss.terrastruct.com/util-go/go2"
)

const (
minRadius = 200.0
cyclePadding = 60.0
maxArcSweep = math.Pi / 2
intersectionPass = 32
)

func Layout(ctx context.Context, g *d2graph.Graph, layout d2graph.LayoutGraph) error {
if len(g.Root.ChildrenArray) == 0 {
return nil
}

for _, obj := range g.Objects {
positionLabelsIcons(obj)
}

radius := calculateRadius(g.Root.ChildrenArray)
center := positionObjects(g.Root.ChildrenArray, radius)
sizeRoot(g.Root)

for _, edge := range g.Edges {
routeCircularArc(edge, center, radius)
}

return nil
}

func calculateRadius(objects []*d2graph.Object) float64 {
if len(objects) < 2 {
return minRadius
}

maxHalfDiagonal := 0.0
for _, obj := range objects {
maxHalfDiagonal = math.Max(maxHalfDiagonal, math.Hypot(obj.Width/2, obj.Height/2))
}

chordRadius := (maxHalfDiagonal + cyclePadding) / math.Sin(math.Pi/float64(len(objects)))
return math.Max(minRadius, chordRadius)
}

func positionObjects(objects []*d2graph.Object, radius float64) *geo.Point {
maxHalfWidth, maxHalfHeight := 0.0, 0.0
for _, obj := range objects {
maxHalfWidth = math.Max(maxHalfWidth, obj.Width/2)
maxHalfHeight = math.Max(maxHalfHeight, obj.Height/2)
}

center := geo.NewPoint(radius+maxHalfWidth+cyclePadding, radius+maxHalfHeight+cyclePadding)
for i, obj := range objects {
angle := angleForIndex(i, len(objects))
nodeCenter := pointOnCircle(center, radius, angle)
obj.TopLeft = geo.NewPoint(nodeCenter.X-obj.Width/2, nodeCenter.Y-obj.Height/2)
}
return center
}

func sizeRoot(root *d2graph.Object) {
maxRight, maxBottom := 0.0, 0.0
for _, obj := range root.ChildrenArray {
maxRight = math.Max(maxRight, obj.TopLeft.X+obj.Width)
maxBottom = math.Max(maxBottom, obj.TopLeft.Y+obj.Height)
}
root.TopLeft = geo.NewPoint(0, 0)
root.Box = geo.NewBox(root.TopLeft, maxRight+cyclePadding, maxBottom+cyclePadding)
}

func routeCircularArc(edge *d2graph.Edge, center *geo.Point, radius float64) {
if edge.Src == nil || edge.Dst == nil || radius == 0 {
return
}

srcAngle := angleOf(center, edge.Src.Center())
dstAngle := angleOf(center, edge.Dst.Center())
if dstAngle <= srcAngle {
dstAngle += 2 * math.Pi
}

startAngle := exitAngle(edge.Src.Box, center, radius, srcAngle, dstAngle)
endAngle := enterAngle(edge.Dst.Box, center, radius, startAngle, dstAngle)
if endAngle <= startAngle {
edge.Route = []*geo.Point{edge.Src.Center(), edge.Dst.Center()}
edge.TraceToShape(edge.Route, 0, 1)
return
}

edge.Route = cubicArcRoute(center, radius, startAngle, endAngle)
edge.IsCurve = true
if edge.Label.Value != "" {
edge.LabelPosition = go2.Pointer(label.InsideMiddleCenter.String())
}
}

func angleForIndex(index, total int) float64 {
return -math.Pi/2 + 2*math.Pi*float64(index)/float64(total)
}

func angleOf(center, point *geo.Point) float64 {
return math.Atan2(point.Y-center.Y, point.X-center.X)
}

func pointOnCircle(center *geo.Point, radius, angle float64) *geo.Point {
return geo.NewPoint(center.X+radius*math.Cos(angle), center.Y+radius*math.Sin(angle))
}

func exitAngle(box *geo.Box, center *geo.Point, radius, startAngle, endAngle float64) float64 {
if box == nil || !box.Contains(pointOnCircle(center, radius, startAngle)) {
return startAngle
}

lo, hi := startAngle, endAngle
for i := 0; i < intersectionPass; i++ {
mid := (lo + hi) / 2
if box.Contains(pointOnCircle(center, radius, mid)) {
lo = mid
} else {
hi = mid
}
}
return hi
}

func enterAngle(box *geo.Box, center *geo.Point, radius, startAngle, endAngle float64) float64 {
if box == nil || !box.Contains(pointOnCircle(center, radius, endAngle)) {
return endAngle
}

lo, hi := startAngle, endAngle
for i := 0; i < intersectionPass; i++ {
mid := (lo + hi) / 2
if box.Contains(pointOnCircle(center, radius, mid)) {
hi = mid
} else {
lo = mid
}
}
return hi
}

func cubicArcRoute(center *geo.Point, radius, startAngle, endAngle float64) []*geo.Point {
route := []*geo.Point{pointOnCircle(center, radius, startAngle)}
remaining := endAngle - startAngle
segments := int(math.Ceil(remaining / maxArcSweep))
sweep := remaining / float64(segments)

for i := 0; i < segments; i++ {
a0 := startAngle + sweep*float64(i)
a1 := a0 + sweep
k := 4.0 / 3.0 * math.Tan((a1-a0)/4.0)

p0 := pointOnCircle(center, radius, a0)
p3 := pointOnCircle(center, radius, a1)
cp1 := geo.NewPoint(
p0.X+(-math.Sin(a0))*radius*k,
p0.Y+(math.Cos(a0))*radius*k,
)
cp2 := geo.NewPoint(
p3.X-(-math.Sin(a1))*radius*k,
p3.Y-(math.Cos(a1))*radius*k,
)
route = append(route, cp1, cp2, p3)
}
return route
}

func positionLabelsIcons(obj *d2graph.Object) {
if obj.Icon != nil && obj.IconPosition == nil {
if len(obj.ChildrenArray) > 0 {
obj.IconPosition = go2.Pointer(label.OutsideTopLeft.String())
if obj.LabelPosition == nil {
obj.LabelPosition = go2.Pointer(label.OutsideTopRight.String())
return
}
} else if obj.SQLTable != nil || obj.Class != nil || obj.Language != "" {
obj.IconPosition = go2.Pointer(label.OutsideTopLeft.String())
} else {
obj.IconPosition = go2.Pointer(label.InsideMiddleCenter.String())
}
}

if obj.HasLabel() && obj.LabelPosition == nil {
if len(obj.ChildrenArray) > 0 {
obj.LabelPosition = go2.Pointer(label.OutsideTopCenter.String())
} else if obj.HasOutsideBottomLabel() {
obj.LabelPosition = go2.Pointer(label.OutsideBottomCenter.String())
} else if obj.Icon != nil {
obj.LabelPosition = go2.Pointer(label.InsideTopCenter.String())
} else {
obj.LabelPosition = go2.Pointer(label.InsideMiddleCenter.String())
}

if float64(obj.LabelDimensions.Width) > obj.Width ||
float64(obj.LabelDimensions.Height) > obj.Height {
if len(obj.ChildrenArray) > 0 {
obj.LabelPosition = go2.Pointer(label.OutsideTopCenter.String())
} else {
obj.LabelPosition = go2.Pointer(label.OutsideBottomCenter.String())
}
}
}
}
10 changes: 10 additions & 0 deletions d2layouts/d2layouts.go
Original file line number Diff line number Diff line change
Expand Up @@ -9,6 +9,7 @@ import (
"strings"

"oss.terrastruct.com/d2/d2graph"
"oss.terrastruct.com/d2/d2layouts/d2cycle"
"oss.terrastruct.com/d2/d2layouts/d2grid"
"oss.terrastruct.com/d2/d2layouts/d2near"
"oss.terrastruct.com/d2/d2layouts/d2sequence"
Expand All @@ -24,6 +25,7 @@ type DiagramType string
const (
DefaultGraphType DiagramType = ""
ConstantNearGraph DiagramType = "constant-near"
CycleDiagram DiagramType = "cycle-diagram"
GridDiagram DiagramType = "grid-diagram"
SequenceDiagram DiagramType = "sequence-diagram"
)
Expand Down Expand Up @@ -248,6 +250,12 @@ func LayoutNested(ctx context.Context, g *d2graph.Graph, graphInfo GraphInfo, co
var err error
if len(g.Objects) > 0 {
switch graphInfo.DiagramType {
case CycleDiagram:
log.Debug(ctx, "layout cycle", slog.Any("rootlevel", g.RootLevel), slog.Any("shapes", g.PrintString()))
if err = d2cycle.Layout(ctx, g, coreLayout); err != nil {
return err
}

case GridDiagram:
log.Debug(ctx, "layout grid", slog.Any("rootlevel", g.RootLevel), slog.Any("shapes", g.PrintString()))
if err = d2grid.Layout(ctx, g); err != nil {
Expand Down Expand Up @@ -362,6 +370,8 @@ func NestedGraphInfo(obj *d2graph.Object) (gi GraphInfo) {
}
if obj.IsSequenceDiagram() {
gi.DiagramType = SequenceDiagram
} else if obj.IsCycleDiagram() {
gi.DiagramType = CycleDiagram
} else if obj.IsGridDiagram() {
gi.DiagramType = GridDiagram
}
Expand Down
3 changes: 3 additions & 0 deletions d2target/d2target.go
Original file line number Diff line number Diff line change
Expand Up @@ -1072,6 +1072,7 @@ const (
ShapeSQLTable = "sql_table"
ShapeImage = "image"
ShapeSequenceDiagram = "sequence_diagram"
ShapeCycleDiagram = "cycle"
ShapeHierarchy = "hierarchy"
)

Expand Down Expand Up @@ -1100,6 +1101,7 @@ var Shapes = []string{
ShapeSQLTable,
ShapeImage,
ShapeSequenceDiagram,
ShapeCycleDiagram,
ShapeHierarchy,
}

Expand Down Expand Up @@ -1170,6 +1172,7 @@ var DSL_SHAPE_TO_SHAPE_TYPE = map[string]string{
ShapeSQLTable: shape.TABLE_TYPE,
ShapeImage: shape.IMAGE_TYPE,
ShapeSequenceDiagram: shape.SQUARE_TYPE,
ShapeCycleDiagram: shape.SQUARE_TYPE,
ShapeHierarchy: shape.SQUARE_TYPE,
}

Expand Down
Loading