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;