diff --git a/README.md b/README.md
index b9fefc6f..92de8c8f 100644
--- a/README.md
+++ b/README.md
@@ -81,7 +81,7 @@ const BasicFlow = () => ;
- `edgeTypes`: object with [edge types](#edge-types--custom-edges)
- `style`: css properties
- `className`: additional class name
-- `connectionLineType`: connection line type = `straight` or `bezier`
+- `connectionLineType`: connection line type = `default` (bezier), `straight`, `step`, `smoothstep`
- `connectionLineStyle`: connection style as svg attributes
- `deleteKeyCode`: default: `8` (delete)
- `selectionKeyCode`: default: `16` (shift)
diff --git a/example/src/CustomNode/index.js b/example/src/CustomNode/index.js
index f65827c6..7c910cd7 100644
--- a/example/src/CustomNode/index.js
+++ b/example/src/CustomNode/index.js
@@ -4,9 +4,9 @@ import ReactFlow, { isEdge, removeElements, addEdge, MiniMap, Controls } from 'r
import ColorSelectorNode from './ColorSelectorNode';
-const onNodeDragStop = node => console.log('drag stop', node);
-const onElementClick = element => console.log('click', element);
-const onLoad = reactFlowInstance => console.log('graph loaded:', reactFlowInstance);
+const onNodeDragStop = (node) => console.log('drag stop', node);
+const onElementClick = (element) => console.log('click', element);
+const onLoad = (reactFlowInstance) => console.log('graph loaded:', reactFlowInstance);
const initBgColor = '#f0e742';
@@ -16,39 +16,47 @@ const CustomNodeFlow = () => {
useEffect(() => {
const onChange = (evt) => {
- setElements(els => els.map(e => {
- if (isEdge(e) || e.id !== '2') {
- return e;
- }
-
- const color = evt.target.value;
-
- setBgColor(color);
-
- return {
- ...e,
- data: {
- ...e.data,
- color
+ setElements((els) =>
+ els.map((e) => {
+ if (isEdge(e) || e.id !== '2') {
+ return e;
}
- };
- }));
+
+ const color = evt.target.value;
+
+ setBgColor(color);
+
+ return {
+ ...e,
+ data: {
+ ...e.data,
+ color,
+ },
+ };
+ })
+ );
};
setElements([
{ id: '1', type: 'input', data: { label: 'An input node' }, position: { x: 0, y: 50 }, sourcePosition: 'right' },
- { id: '2', type: 'selectorNode', data: { onChange: onChange, color: initBgColor }, style: { border: '1px solid #777', padding: 10 }, position: { x: 250, y: 50 } },
+ {
+ id: '2',
+ type: 'selectorNode',
+ data: { onChange: onChange, color: initBgColor },
+ style: { border: '1px solid #777', padding: 10 },
+ position: { x: 250, y: 50 },
+ },
{ id: '3', type: 'output', data: { label: 'Output A' }, position: { x: 550, y: 25 }, targetPosition: 'left' },
{ id: '4', type: 'output', data: { label: 'Output B' }, position: { x: 550, y: 100 }, targetPosition: 'left' },
{ id: 'e1-2', source: '1', target: '2', animated: true, style: { stroke: '#fff' } },
{ id: 'e2a-3', source: '2__a', target: '3', animated: true, style: { stroke: '#fff' } },
{ id: 'e2b-4', source: '2__b', target: '4', animated: true, style: { stroke: '#fff' } },
- ])
+ ]);
}, []);
- const onElementsRemove = (elementsToRemove) => setElements(els => removeElements(elementsToRemove, els));
- const onConnect = (params) => setElements(els => addEdge(params, els));
+ const onElementsRemove = (elementsToRemove) => setElements((els) => removeElements(elementsToRemove, els));
+ const onConnect = (params) => setElements((els) => addEdge(params, els));
return (
{
selectorNode: ColorSelectorNode,
}}
connectionLineStyle={{ stroke: '#ddd', strokeWidth: 2 }}
- connectionLineType="bezier"
snapToGrid={true}
snapGrid={[16, 16]}
>
{
+ nodeColor={(n) => {
if (n.type === 'input') return 'blue';
if (n.type === 'selectorNode') return bgColor;
if (n.type === 'output') return 'green';
@@ -77,6 +84,6 @@ const CustomNodeFlow = () => {
);
-}
+};
-export default CustomNodeFlow;
\ No newline at end of file
+export default CustomNodeFlow;
diff --git a/example/src/Overview/index.js b/example/src/Overview/index.js
index 6858a6eb..dffbb66f 100644
--- a/example/src/Overview/index.js
+++ b/example/src/Overview/index.js
@@ -127,7 +127,6 @@ const OverviewFlow = () => {
style={{ width: '100%', height: '100%' }}
onLoad={onLoad}
connectionLineStyle={{ stroke: '#ddd', strokeWidth: 2 }}
- connectionLineType="bezier"
snapToGrid={true}
snapGrid={[16, 16]}
>
diff --git a/src/components/ConnectionLine/index.tsx b/src/components/ConnectionLine/index.tsx
index e21ce534..d7b7033d 100644
--- a/src/components/ConnectionLine/index.tsx
+++ b/src/components/ConnectionLine/index.tsx
@@ -1,6 +1,9 @@
import React, { useEffect, useState, CSSProperties } from 'react';
import cx from 'classnames';
+import { getBezierPath } from '../Edges/BezierEdge';
+import { getStepPath } from '../Edges/StepEdge';
+import { getSmoothStepPath } from '../Edges/SmoothStepEdge';
import { ElementId, Node, Transform, HandleElement, Position, ConnectionLineType, HandleType } from '../../types';
interface ConnectionLineProps {
@@ -57,17 +60,46 @@ export default ({
const targetY = (connectionPositionY - transform[1]) * (1 / transform[2]);
let dAttr: string = '';
+ const xOffset = Math.abs(targetX - sourceX) / 2;
+ const centerX = targetX < sourceX ? targetX + xOffset : targetX - xOffset;
+
+ const yOffset = Math.abs(targetY - sourceY) / 2;
+ const centerY = targetY < sourceY ? targetY + yOffset : targetY - yOffset;
+ const isRightOrLeft = sourceHandle?.position === Position.Left || sourceHandle?.position === Position.Right;
+ const targetPosition = isRightOrLeft ? Position.Left : Position.Top;
if (connectionLineType === ConnectionLineType.Bezier) {
- if (sourceHandle?.position === Position.Left || sourceHandle?.position === Position.Right) {
- const xOffset = Math.abs(targetX - sourceX) / 2;
- const centerX = targetX < sourceX ? targetX + xOffset : targetX - xOffset;
- dAttr = `M${sourceX},${sourceY} C${centerX},${sourceY} ${centerX},${targetY} ${targetX},${targetY}`;
- } else {
- const yOffset = Math.abs(targetY - sourceY) / 2;
- const centerY = targetY < sourceY ? targetY + yOffset : targetY - yOffset;
- dAttr = `M${sourceX},${sourceY} C${sourceX},${centerY} ${targetX},${centerY} ${targetX},${targetY}`;
- }
+ dAttr = getBezierPath({
+ centerX,
+ centerY,
+ sourceX,
+ sourceY,
+ sourcePosition: sourceHandle?.position,
+ targetX,
+ targetY,
+ targetPosition,
+ });
+ } else if (connectionLineType === ConnectionLineType.Step) {
+ dAttr = getStepPath({
+ centerY,
+ sourceX,
+ sourceY,
+ targetX,
+ targetY,
+ });
+ } else if (connectionLineType === ConnectionLineType.SmoothStep) {
+ dAttr = getSmoothStepPath({
+ xOffset,
+ yOffset,
+ centerX,
+ centerY,
+ sourceX,
+ sourceY,
+ sourcePosition: sourceHandle?.position,
+ targetX,
+ targetY,
+ targetPosition,
+ });
} else {
dAttr = `M${sourceX},${sourceY} ${targetX},${targetY}`;
}
diff --git a/src/components/Edges/BezierEdge.tsx b/src/components/Edges/BezierEdge.tsx
index a6bdc8a5..816ed3db 100644
--- a/src/components/Edges/BezierEdge.tsx
+++ b/src/components/Edges/BezierEdge.tsx
@@ -3,6 +3,42 @@ import React, { memo } from 'react';
import EdgeText from './EdgeText';
import { EdgeBezierProps, Position } from '../../types';
+interface GetBezierPathParams {
+ centerX: number;
+ centerY: number;
+ sourceX: number;
+ sourceY: number;
+ sourcePosition?: Position;
+ targetX: number;
+ targetY: number;
+ targetPosition?: Position;
+}
+
+export function getBezierPath({
+ centerX,
+ centerY,
+ sourceX,
+ sourceY,
+ sourcePosition = Position.Bottom,
+ targetX,
+ targetY,
+ targetPosition = Position.Top,
+}: GetBezierPathParams): string {
+ let path = `M${sourceX},${sourceY} C${sourceX},${centerY} ${targetX},${centerY} ${targetX},${targetY}`;
+
+ const leftAndRight = [Position.Left, Position.Right];
+
+ if (leftAndRight.includes(sourcePosition) && leftAndRight.includes(targetPosition)) {
+ path = `M${sourceX},${sourceY} C${centerX},${sourceY} ${centerX},${targetY} ${targetX},${targetY}`;
+ } else if (leftAndRight.includes(targetPosition)) {
+ path = `M${sourceX},${sourceY} C${sourceX},${targetY} ${sourceX},${targetY} ${targetX},${targetY}`;
+ } else if (leftAndRight.includes(sourcePosition)) {
+ path = `M${sourceX},${sourceY} C${targetX},${sourceY} ${targetX},${sourceY} ${targetX},${targetY}`;
+ }
+
+ return path;
+}
+
export default memo(
({
sourceX,
@@ -23,17 +59,16 @@ export default memo(
const xOffset = Math.abs(targetX - sourceX) / 2;
const centerX = targetX < sourceX ? targetX + xOffset : targetX - xOffset;
- let dAttr = `M${sourceX},${sourceY} C${sourceX},${centerY} ${targetX},${centerY} ${targetX},${targetY}`;
-
- const leftAndRight = [Position.Left, Position.Right];
-
- if (leftAndRight.includes(sourcePosition) && leftAndRight.includes(targetPosition)) {
- dAttr = `M${sourceX},${sourceY} C${centerX},${sourceY} ${centerX},${targetY} ${targetX},${targetY}`;
- } else if (leftAndRight.includes(targetPosition)) {
- dAttr = `M${sourceX},${sourceY} C${sourceX},${targetY} ${sourceX},${targetY} ${targetX},${targetY}`;
- } else if (leftAndRight.includes(sourcePosition)) {
- dAttr = `M${sourceX},${sourceY} C${targetX},${sourceY} ${targetX},${sourceY} ${targetX},${targetY}`;
- }
+ const path = getBezierPath({
+ centerX,
+ centerY,
+ sourceX,
+ sourceY,
+ sourcePosition,
+ targetX,
+ targetY,
+ targetPosition,
+ });
const text = label ? (
-
+
{text}
>
);
diff --git a/src/components/Edges/SmoothStepEdge.tsx b/src/components/Edges/SmoothStepEdge.tsx
index 7a93a2d5..8f61e7c8 100644
--- a/src/components/Edges/SmoothStepEdge.tsx
+++ b/src/components/Edges/SmoothStepEdge.tsx
@@ -32,6 +32,99 @@ const topRightCorner = (cornerX: number, cornerY: number, cornerSize: number): s
const rightTopCorner = (cornerX: number, cornerY: number, cornerSize: number): string =>
`L ${cornerX - cornerSize},${cornerY}Q ${cornerX},${cornerY} ${cornerX},${cornerY + cornerSize}`;
+interface GetSmoothStepPathParams {
+ xOffset: number;
+ yOffset: number;
+ centerX: number;
+ centerY: number;
+ sourceX: number;
+ sourceY: number;
+ sourcePosition?: Position;
+ targetX: number;
+ targetY: number;
+ targetPosition?: Position;
+}
+
+export function getSmoothStepPath({
+ xOffset,
+ yOffset,
+ centerX,
+ centerY,
+ sourceX,
+ sourceY,
+ sourcePosition = Position.Bottom,
+ targetX,
+ targetY,
+ targetPosition = Position.Top,
+}: GetSmoothStepPathParams): string {
+ const cornerWidth = Math.min(5, Math.abs(targetX - sourceX));
+ const cornerHeight = Math.min(5, Math.abs(targetY - sourceY));
+ const cornerSize = Math.min(cornerWidth, cornerHeight, xOffset, yOffset);
+
+ const leftAndRight = [Position.Left, Position.Right];
+
+ let firstCornerPath = null;
+ let secondCornerPath = null;
+
+ // default case: source and target positions are top or bottom
+ if (sourceX < targetX) {
+ firstCornerPath =
+ sourceY < targetY ? bottomLeftCorner(sourceX, centerY, cornerSize) : topLeftCorner(sourceX, centerY, cornerSize);
+ secondCornerPath =
+ sourceY < targetY
+ ? rightTopCorner(targetX, centerY, cornerSize)
+ : rightBottomCorner(targetX, centerY, cornerSize);
+ } else if (sourceX > targetX) {
+ firstCornerPath =
+ sourceY < targetY
+ ? bottomRightCorner(sourceX, centerY, cornerSize)
+ : topRightCorner(sourceX, centerY, cornerSize);
+ secondCornerPath =
+ sourceY < targetY ? leftTopCorner(targetX, centerY, cornerSize) : leftBottomCorner(targetX, centerY, cornerSize);
+ }
+
+ if (leftAndRight.includes(sourcePosition) && leftAndRight.includes(targetPosition)) {
+ if (sourceX < targetX) {
+ firstCornerPath =
+ sourceY < targetY
+ ? rightTopCorner(centerX, sourceY, cornerSize)
+ : rightBottomCorner(centerX, sourceY, cornerSize);
+ secondCornerPath =
+ sourceY < targetY
+ ? bottomLeftCorner(centerX, targetY, cornerSize)
+ : topLeftCorner(centerX, targetY, cornerSize);
+ }
+ } else if (leftAndRight.includes(sourcePosition) && !leftAndRight.includes(targetPosition)) {
+ if (sourceX < targetX) {
+ firstCornerPath =
+ sourceY < targetY
+ ? rightTopCorner(targetX, sourceY, cornerSize)
+ : rightBottomCorner(targetX, sourceY, cornerSize);
+ } else if (sourceX > targetX) {
+ firstCornerPath =
+ sourceY < targetY
+ ? bottomRightCorner(sourceX, targetY, cornerSize)
+ : topRightCorner(sourceX, targetY, cornerSize);
+ }
+ secondCornerPath = '';
+ } else if (!leftAndRight.includes(sourcePosition) && leftAndRight.includes(targetPosition)) {
+ if (sourceX < targetX) {
+ firstCornerPath =
+ sourceY < targetY
+ ? bottomLeftCorner(sourceX, targetY, cornerSize)
+ : topLeftCorner(sourceX, targetY, cornerSize);
+ } else if (sourceX > targetX) {
+ firstCornerPath =
+ sourceY < targetY
+ ? bottomRightCorner(sourceX, targetY, cornerSize)
+ : topRightCorner(sourceX, targetY, cornerSize);
+ }
+ secondCornerPath = '';
+ }
+
+ return `M ${sourceX},${sourceY}${firstCornerPath}${secondCornerPath}L ${targetX},${targetY}`;
+}
+
export default memo(
({
sourceX,
@@ -52,77 +145,18 @@ export default memo(
const xOffset = Math.abs(targetX - sourceX) / 2;
const centerX = targetX < sourceX ? targetX + xOffset : targetX - xOffset;
- const cornerWidth = Math.min(5, Math.abs(targetX - sourceX));
- const cornerHeight = Math.min(5, Math.abs(targetY - sourceY));
- const cornerSize = Math.min(cornerWidth, cornerHeight, xOffset, yOffset);
-
- const leftAndRight = [Position.Left, Position.Right];
-
- let path;
- let firstCornerPath = null;
- let secondCornerPath = null;
-
- // default case: source and target positions are top or bottom
- if (sourceX < targetX) {
- firstCornerPath =
- sourceY < targetY
- ? bottomLeftCorner(sourceX, centerY, cornerSize)
- : topLeftCorner(sourceX, centerY, cornerSize);
- secondCornerPath =
- sourceY < targetY
- ? rightTopCorner(targetX, centerY, cornerSize)
- : rightBottomCorner(targetX, centerY, cornerSize);
- } else if (sourceX > targetX) {
- firstCornerPath =
- sourceY < targetY
- ? bottomRightCorner(sourceX, centerY, cornerSize)
- : topRightCorner(sourceX, centerY, cornerSize);
- secondCornerPath =
- sourceY < targetY
- ? leftTopCorner(targetX, centerY, cornerSize)
- : leftBottomCorner(targetX, centerY, cornerSize);
- }
-
- if (leftAndRight.includes(sourcePosition) && leftAndRight.includes(targetPosition)) {
- if (sourceX < targetX) {
- firstCornerPath =
- sourceY < targetY
- ? rightTopCorner(centerX, sourceY, cornerSize)
- : rightBottomCorner(centerX, sourceY, cornerSize);
- secondCornerPath =
- sourceY < targetY
- ? bottomLeftCorner(centerX, targetY, cornerSize)
- : topLeftCorner(centerX, targetY, cornerSize);
- }
- } else if (leftAndRight.includes(sourcePosition) && !leftAndRight.includes(targetPosition)) {
- if (sourceX < targetX) {
- firstCornerPath =
- sourceY < targetY
- ? rightTopCorner(targetX, sourceY, cornerSize)
- : rightBottomCorner(targetX, sourceY, cornerSize);
- } else if (sourceX > targetX) {
- firstCornerPath =
- sourceY < targetY
- ? bottomRightCorner(sourceX, targetY, cornerSize)
- : topRightCorner(sourceX, targetY, cornerSize);
- }
- secondCornerPath = '';
- } else if (!leftAndRight.includes(sourcePosition) && leftAndRight.includes(targetPosition)) {
- if (sourceX < targetX) {
- firstCornerPath =
- sourceY < targetY
- ? bottomLeftCorner(sourceX, targetY, cornerSize)
- : topLeftCorner(sourceX, targetY, cornerSize);
- } else if (sourceX > targetX) {
- firstCornerPath =
- sourceY < targetY
- ? bottomRightCorner(sourceX, targetY, cornerSize)
- : topRightCorner(sourceX, targetY, cornerSize);
- }
- secondCornerPath = '';
- }
-
- path = `M ${sourceX},${sourceY}${firstCornerPath}${secondCornerPath}L ${targetX},${targetY}`;
+ const path = getSmoothStepPath({
+ xOffset,
+ yOffset,
+ centerX,
+ centerY,
+ sourceX,
+ sourceY,
+ sourcePosition,
+ targetX,
+ targetY,
+ targetPosition,
+ });
const text = label ? (
{
const yOffset = Math.abs(targetY - sourceY) / 2;
@@ -11,6 +23,8 @@ export default memo(
const xOffset = Math.abs(targetX - sourceX) / 2;
const centerX = targetX < sourceX ? targetX + xOffset : targetX - xOffset;
+ const path = getStepPath({ centerY, sourceX, sourceY, targetX, targetY });
+
const text = label ? (
-
+
{text}
>
);
diff --git a/src/types/index.ts b/src/types/index.ts
index 666a6e2d..47a27e58 100644
--- a/src/types/index.ts
+++ b/src/types/index.ts
@@ -167,8 +167,10 @@ export interface Connection {
}
export enum ConnectionLineType {
- Bezier = 'bezier',
+ Bezier = 'default',
Straight = 'straight',
+ Step = 'step',
+ SmoothStep = 'smoothstep',
}
export type OnConnectFunc = (connection: Connection) => void;