Skip to content

Commit dbf8367

Browse files
authored
Merge pull request #8576 from processing/fix/strands-set-branch
Handle strands set() calls in branches and loops
2 parents 8e87a23 + d8ba7bc commit dbf8367

5 files changed

Lines changed: 361 additions & 13 deletions

File tree

‎src/strands/ir_builders.js‎

Lines changed: 11 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -86,8 +86,17 @@ export function binaryOpNode(strandsContext, leftStrandsNode, rightArg, opCode)
8686
let finalRightNodeID = rightStrandsNode.id;
8787

8888
// Check if we have to cast either node
89-
const leftType = DAG.extractNodeTypeInfo(dag, leftStrandsNode.id);
90-
const rightType = DAG.extractNodeTypeInfo(dag, rightStrandsNode.id);
89+
let leftType = DAG.extractNodeTypeInfo(dag, leftStrandsNode.id);
90+
let rightType = DAG.extractNodeTypeInfo(dag, rightStrandsNode.id);
91+
92+
// Update ASSIGN_ON_USE nodes to match the type of the other operand
93+
if (leftType.baseType === BaseType.ASSIGN_ON_USE && rightType.baseType !== BaseType.ASSIGN_ON_USE) {
94+
DAG.propagateTypeToAssignOnUse(dag, leftStrandsNode.id, rightType.baseType, rightType.dimension);
95+
leftType = DAG.extractNodeTypeInfo(dag, leftStrandsNode.id);
96+
} else if (rightType.baseType === BaseType.ASSIGN_ON_USE && leftType.baseType !== BaseType.ASSIGN_ON_USE) {
97+
DAG.propagateTypeToAssignOnUse(dag, rightStrandsNode.id, leftType.baseType, leftType.dimension);
98+
rightType = DAG.extractNodeTypeInfo(dag, rightStrandsNode.id);
99+
}
91100
const cast = { node: null, toType: leftType };
92101
const bothDeferred = leftType.baseType === rightType.baseType && leftType.baseType === BaseType.DEFER;
93102
if (bothDeferred) {

‎src/strands/ir_dag.js‎

Lines changed: 31 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
import { NodeTypeRequiredFields, NodeTypeToName, BasePriority, StatementType } from './ir_types';
1+
import { NodeTypeRequiredFields, NodeTypeToName, BasePriority, StatementType, BaseType } from './ir_types';
22
import * as FES from './strands_FES';
33

44
/////////////////////////////////
@@ -81,6 +81,36 @@ export function extractNodeTypeInfo(dag, nodeID) {
8181
};
8282
}
8383

84+
// Propagate a known type to an ASSIGN_ON_USE node and all its ASSIGN_ON_USE dependencies
85+
export function propagateTypeToAssignOnUse(dag, nodeId, baseType, dimension, visited = new Set()) {
86+
// Avoid infinite loops
87+
if (visited.has(nodeId)) {
88+
return;
89+
}
90+
visited.add(nodeId);
91+
92+
const node = getNodeDataFromID(dag, nodeId);
93+
94+
// Only update if this node is ASSIGN_ON_USE
95+
if (node.baseType !== BaseType.ASSIGN_ON_USE) {
96+
return;
97+
}
98+
99+
// Update this node's type
100+
dag.baseTypes[nodeId] = baseType;
101+
dag.dimensions[nodeId] = dimension;
102+
103+
// Recursively propagate to any ASSIGN_ON_USE dependencies
104+
if (node.dependsOn && node.dependsOn.length > 0) {
105+
for (const depId of node.dependsOn) {
106+
const dep = getNodeDataFromID(dag, depId);
107+
if (dep.baseType === BaseType.ASSIGN_ON_USE) {
108+
propagateTypeToAssignOnUse(dag, depId, baseType, dimension, visited);
109+
}
110+
}
111+
}
112+
}
113+
84114
/////////////////////////////////
85115
// Private functions
86116
/////////////////////////////////

‎src/strands/strands_phi_utils.js‎

Lines changed: 20 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -11,12 +11,27 @@ export function createPhiNode(strandsContext, phiInputs, varName) {
1111

1212
// Get dimension and baseType from first valid input, skipping ASSIGN_ON_USE nodes
1313
const inputNodes = validInputs.map((input) => DAG.getNodeDataFromID(strandsContext.dag, input.value.id));
14-
let firstInput = inputNodes.find((input) => input.baseType !== BaseType.ASSIGN_ON_USE && input.dimension) ??
15-
inputNodes.find((input) => input.baseType !== BaseType.ASSIGN_ON_USE) ??
16-
inputNodes[0];
1714

18-
const dimension = firstInput.dimension;
19-
const baseType = firstInput.baseType;
15+
// Find first non-ASSIGN_ON_USE input to determine type
16+
let typeSource = inputNodes.find((input) => input.baseType !== BaseType.ASSIGN_ON_USE && input.dimension) ??
17+
inputNodes.find((input) => input.baseType !== BaseType.ASSIGN_ON_USE);
18+
19+
// If all are ASSIGN_ON_USE, fall back to first input
20+
if (!typeSource) {
21+
typeSource = inputNodes[0];
22+
}
23+
24+
const dimension = typeSource.dimension;
25+
const baseType = typeSource.baseType;
26+
27+
// Propagate the type to all ASSIGN_ON_USE inputs
28+
if (baseType !== BaseType.ASSIGN_ON_USE) {
29+
for (const input of inputNodes) {
30+
if (input.baseType === BaseType.ASSIGN_ON_USE) {
31+
DAG.propagateTypeToAssignOnUse(strandsContext.dag, input.id, baseType, dimension);
32+
}
33+
}
34+
}
2035

2136
const nodeData = {
2237
nodeType: NodeType.PHI,

‎src/strands/strands_transpiler.js‎

Lines changed: 229 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -26,7 +26,7 @@ function replaceBinaryOperator(codeSource) {
2626
}
2727
}
2828
function nodeIsUniform(ancestor) {
29-
return ancestor.type === 'CallExpression'
29+
return ancestor && ancestor.type === 'CallExpression'
3030
&& (
3131
(
3232
// Global mode
@@ -41,7 +41,7 @@ function nodeIsUniform(ancestor) {
4141
}
4242

4343
function nodeIsVarying(node) {
44-
return node?.type === 'CallExpression'
44+
return node && node.type === 'CallExpression'
4545
&& (
4646
(
4747
// Global mode
@@ -1286,6 +1286,226 @@ function transformHelperFunction(functionNode) {
12861286
functionNode.body.body.push(finalReturn);
12871287
}
12881288

1289+
// Helper function to check if a function body contains .set() calls in control flow
1290+
function functionHasSetInControlFlow(functionNode) {
1291+
let hasSetInControlFlow = false;
1292+
let inControlFlow = 0;
1293+
1294+
const checkForSetCalls = {
1295+
IfStatement(node, state, c) {
1296+
inControlFlow++;
1297+
if (node.test) c(node.test, state);
1298+
if (node.consequent) c(node.consequent, state);
1299+
if (node.alternate) c(node.alternate, state);
1300+
inControlFlow--;
1301+
},
1302+
ForStatement(node, state, c) {
1303+
inControlFlow++;
1304+
if (node.init) c(node.init, state);
1305+
if (node.test) c(node.test, state);
1306+
if (node.update) c(node.update, state);
1307+
if (node.body) c(node.body, state);
1308+
inControlFlow--;
1309+
},
1310+
CallExpression(node) {
1311+
// Check if this is a .set() call
1312+
if (inControlFlow > 0 &&
1313+
node.callee?.type === 'MemberExpression' &&
1314+
node.callee?.property?.name === 'set') {
1315+
hasSetInControlFlow = true;
1316+
}
1317+
}
1318+
};
1319+
1320+
if (functionNode.body && functionNode.body.type === 'BlockStatement') {
1321+
recursive(functionNode.body, {}, checkForSetCalls);
1322+
}
1323+
1324+
return hasSetInControlFlow;
1325+
}
1326+
1327+
// Transform a function to use __setValue pattern instead of .set() calls in branches/loops
1328+
function transformFunctionSetCalls(functionNode) {
1329+
if (!functionNode.body || functionNode.body.type !== 'BlockStatement') {
1330+
return; // Can't transform arrow functions with expression bodies
1331+
}
1332+
1333+
// Track which hooks have .set() calls, mapping expression string to the actual AST node
1334+
const hooksWithSetCalls = new Map(); // exprString -> hookObjectNode
1335+
1336+
// First pass: find all hooks that have .set() calls in control flow
1337+
const findSetCalls = {
1338+
CallExpression(node) {
1339+
if (node.callee?.type === 'MemberExpression' &&
1340+
node.callee?.property?.name === 'set' &&
1341+
node.callee?.object) {
1342+
// This is something like filterColor.set(...) or myp5.filterColor.set(...)
1343+
const hookObjectNode = node.callee.object;
1344+
const exprString = escodegen.generate(hookObjectNode);
1345+
if (!hooksWithSetCalls.has(exprString)) {
1346+
hooksWithSetCalls.set(exprString, hookObjectNode);
1347+
}
1348+
}
1349+
}
1350+
};
1351+
1352+
recursive(functionNode.body, {}, findSetCalls);
1353+
1354+
if (hooksWithSetCalls.size === 0) {
1355+
return; // No .set() calls to transform
1356+
}
1357+
1358+
// For each hook with .set() calls, add intermediate variable and transform
1359+
for (const [exprString, hookObjectNode] of hooksWithSetCalls) {
1360+
// Create a safe variable name from the expression
1361+
const safeVarName = exprString.replace(/[^a-zA-Z0-9_]/g, '_');
1362+
const intermediateVarName = `__${safeVarName}_value`;
1363+
1364+
// 1. Find the .begin() call and insert intermediate variable right after it
1365+
const intermediateVarDecl = {
1366+
type: 'VariableDeclaration',
1367+
declarations: [{
1368+
type: 'VariableDeclarator',
1369+
id: { type: 'Identifier', name: intermediateVarName },
1370+
init: null
1371+
}],
1372+
kind: 'let'
1373+
};
1374+
1375+
let beginCallIndex = -1;
1376+
for (let i = 0; i < functionNode.body.body.length; i++) {
1377+
const stmt = functionNode.body.body[i];
1378+
if (stmt.type === 'ExpressionStatement' &&
1379+
stmt.expression?.type === 'CallExpression' &&
1380+
stmt.expression?.callee?.type === 'MemberExpression' &&
1381+
stmt.expression?.callee?.property?.name === 'begin') {
1382+
const beginExprString = escodegen.generate(stmt.expression.callee.object);
1383+
if (beginExprString === exprString) {
1384+
beginCallIndex = i;
1385+
break;
1386+
}
1387+
}
1388+
}
1389+
1390+
// Insert intermediate variable after .begin() if found, otherwise at the start
1391+
if (beginCallIndex !== -1) {
1392+
functionNode.body.body.splice(beginCallIndex + 1, 0, intermediateVarDecl);
1393+
} else {
1394+
functionNode.body.body.unshift(intermediateVarDecl);
1395+
}
1396+
1397+
// 2. Transform all .set() calls to assignments
1398+
const transformSetToAssignment = {
1399+
CallExpression(node, state, ancestors) {
1400+
// Check if this is a .set() call for this hook
1401+
if (node.callee?.type === 'MemberExpression' &&
1402+
node.callee?.property?.name === 'set' &&
1403+
node.callee?.object) {
1404+
const currentExprString = escodegen.generate(node.callee.object);
1405+
if (currentExprString === exprString && node.arguments.length > 0) {
1406+
// Find the parent statement
1407+
let parentStmt = null;
1408+
for (let i = ancestors.length - 1; i >= 0; i--) {
1409+
if (ancestors[i].type === 'ExpressionStatement') {
1410+
parentStmt = ancestors[i];
1411+
break;
1412+
}
1413+
}
1414+
1415+
if (parentStmt) {
1416+
// Replace the .set() call with an assignment
1417+
parentStmt.type = 'ExpressionStatement';
1418+
parentStmt.expression = {
1419+
type: 'AssignmentExpression',
1420+
operator: '=',
1421+
left: { type: 'Identifier', name: intermediateVarName },
1422+
right: node.arguments[0]
1423+
};
1424+
}
1425+
}
1426+
}
1427+
}
1428+
};
1429+
1430+
ancestor(functionNode.body, transformSetToAssignment);
1431+
1432+
// 3. Find the .end() call and insert final .set() call right before it
1433+
const finalSetCall = {
1434+
type: 'ExpressionStatement',
1435+
expression: {
1436+
type: 'CallExpression',
1437+
callee: {
1438+
type: 'MemberExpression',
1439+
object: JSON.parse(JSON.stringify(hookObjectNode)), // Deep copy the original node
1440+
property: { type: 'Identifier', name: 'set' },
1441+
computed: false
1442+
},
1443+
arguments: [{ type: 'Identifier', name: intermediateVarName }]
1444+
}
1445+
};
1446+
1447+
// Find the .end() call for this hook
1448+
let endCallIndex = -1;
1449+
for (let i = 0; i < functionNode.body.body.length; i++) {
1450+
const stmt = functionNode.body.body[i];
1451+
if (stmt.type === 'ExpressionStatement' &&
1452+
stmt.expression?.type === 'CallExpression' &&
1453+
stmt.expression?.callee?.type === 'MemberExpression' &&
1454+
stmt.expression?.callee?.property?.name === 'end') {
1455+
const endExprString = escodegen.generate(stmt.expression.callee.object);
1456+
if (endExprString === exprString) {
1457+
endCallIndex = i;
1458+
break;
1459+
}
1460+
}
1461+
}
1462+
1463+
// Insert the final .set() call before .end() if found, otherwise at the end
1464+
if (endCallIndex !== -1) {
1465+
functionNode.body.body.splice(endCallIndex, 0, finalSetCall);
1466+
} else {
1467+
// If no .end() found, insert before return statement or at the end
1468+
const lastStatement = functionNode.body.body[functionNode.body.body.length - 1];
1469+
if (lastStatement && lastStatement.type === 'ReturnStatement') {
1470+
functionNode.body.body.splice(functionNode.body.body.length - 1, 0, finalSetCall);
1471+
} else {
1472+
functionNode.body.body.push(finalSetCall);
1473+
}
1474+
}
1475+
}
1476+
}
1477+
1478+
// Main transformation pass: find and transform functions with .set() calls in control flow
1479+
function transformSetCallsInControlFlow(ast) {
1480+
const functionsToTransform = [];
1481+
1482+
// Collect functions that have .set() calls in control flow
1483+
const collectFunctions = {
1484+
ArrowFunctionExpression(node, ancestors) {
1485+
if (functionHasSetInControlFlow(node)) {
1486+
functionsToTransform.push(node);
1487+
}
1488+
},
1489+
FunctionExpression(node, ancestors) {
1490+
if (functionHasSetInControlFlow(node)) {
1491+
functionsToTransform.push(node);
1492+
}
1493+
},
1494+
FunctionDeclaration(node, ancestors) {
1495+
if (functionHasSetInControlFlow(node)) {
1496+
functionsToTransform.push(node);
1497+
}
1498+
}
1499+
};
1500+
1501+
ancestor(ast, collectFunctions);
1502+
1503+
// Transform each collected function
1504+
for (const funcNode of functionsToTransform) {
1505+
transformFunctionSetCalls(funcNode);
1506+
}
1507+
}
1508+
12891509
// Main transformation pass: find and transform helper functions with early returns
12901510
function transformHelperFunctionEarlyReturns(ast) {
12911511
const helperFunctionsToTransform = [];
@@ -1329,16 +1549,20 @@ export function transpileStrandsToJS(p5, sourceString, srcLocations, scope) {
13291549
ecmaVersion: 2021,
13301550
locations: srcLocations
13311551
});
1332-
// First pass: transform everything except if/for statements using normal ancestor traversal
1552+
1553+
// First pass: transform .set() calls in control flow to use intermediate variables
1554+
transformSetCallsInControlFlow(ast);
1555+
1556+
// Second pass: transform everything except if/for statements using normal ancestor traversal
13331557
const nonControlFlowCallbacks = { ...ASTCallbacks };
13341558
delete nonControlFlowCallbacks.IfStatement;
13351559
delete nonControlFlowCallbacks.ForStatement;
13361560
ancestor(ast, nonControlFlowCallbacks, undefined, { varyings: {} });
13371561

1338-
// Second pass: transform helper functions with early returns to use __returnValue pattern
1562+
// Third pass: transform helper functions with early returns to use __returnValue pattern
13391563
transformHelperFunctionEarlyReturns(ast);
13401564

1341-
// Third pass: transform if/for statements in post-order using recursive traversal
1565+
// Fourth pass: transform if/for statements in post-order using recursive traversal
13421566
const postOrderControlFlowTransform = {
13431567
IfStatement(node, state, c) {
13441568
state.inControlFlow++;

0 commit comments

Comments
 (0)