typegpu
Version:
A thin layer between JS and WebGPU/WGSL that improves development experience and allows for faster iteration.
1,391 lines • 69 kB
JavaScript
import * as tinyest from 'tinyest';
import { stitch } from "../core/resolve/stitch.js";
import { arrayOf } from "../data/array.js";
import { UnknownData, unptr } from "../data/dataTypes.js";
import { bool, i32, u32 } from "../data/numeric.js";
import { vec2u, vec3u, vec4u } from "../data/vector.js";
import { fallthroughCopyOrigin, isAlias, snip, } from "../data/snippet.js";
import * as wgsl from "../data/wgslTypes.js";
import { invariant, ResolutionError, WgslTypeError } from "../errors.js";
import { getName } from "../shared/meta.js";
import { $gpuCallable, $internal, $providing, isMarkedInternal } from "../shared/symbols.js";
import { safeStringify } from "../shared/stringify.js";
import { pow } from "../std/numeric.js";
import { add, div, mul, neg, sub } from "../std/operators.js";
import { eq, ne, lt, le, gt, ge, not } from "../std/boolean.js";
import { isGPUCallable, isKnownAtComptime, } from "../types.js";
import { convertStructValues, convertToCommonType, tryConvertSnippet } from "./conversion.js";
import { ArrayExpression, coerceToSnippet, concretize, numericLiteralToSnippet, } from "./generationHelpers.js";
import { accessIndex } from "./accessIndex.js";
import { accessProp } from "./accessProp.js";
import { resolveData } from "../core/resolve/resolveData.js";
import { createPtrFromOrigin, implicitFrom, ptrFn } from "../data/ptr.js";
import { _ref, RefOperator } from "../data/ref.js";
import { constant } from "../core/constant/tgpuConstant.js";
import { unroll, UnrollableIterable } from "../core/unroll/tgpuUnroll.js";
import { isGenericFn } from "../core/function/tgpuFn.js";
import { AutoStruct } from "../data/autoStruct.js";
import { mathToStd, supportedLogOps } from "./jsPolyfills.js";
import * as forOfUtils from "./forOfUtils.js";
import { isTgpuRange } from "../std/range.js";
import { stringifyNode } from "../shared/tseynit.js";
import { getAttributesString } from "../data/attributes.js";
import { validSelectBranchTypes } from "../std/boolean.js";
import { isInfixDispatch } from "./infixDispatch.js";
import { logger } from "../tgpuLogger.js";
const { NodeTypeCatalog: NODE } = tinyest;
const parenthesizedOps = [
'==',
'!=',
'===',
'!==',
'<',
'<=',
'>',
'>=',
'<<',
'>>',
'>>>',
'+',
'-',
'*',
'/',
'%',
'|',
'^',
'&',
];
const binaryLogicalOps = ['&&', '||', '==', '!=', '===', '!==', '<', '<=', '>', '>='];
const binaryRelationalOpToStdMap = {
'===': eq.toString(),
'!==': ne.toString(),
'<': lt.toString(),
'<=': le.toString(),
'>': gt.toString(),
'>=': ge.toString(),
};
const bitShiftOps = ['<<', '>>', '<<=', '>>=', '>>>', '>>>='];
const OP_MAP = {
//
// binary
//
'===': '==',
'!==': '!=',
'>>>': '>>',
get in() {
throw new Error('The `in` operator is unsupported in TypeGPU functions.');
},
get instanceof() {
throw new Error('The `instanceof` operator is unsupported in TypeGPU functions.');
},
get '|>'() {
throw new Error('The `|>` operator is unsupported in TypeGPU functions.');
},
//
// logical
//
get '??'() {
throw new Error('The `??` operator is unsupported in TypeGPU functions.');
},
//
// assignment
//
'>>>=': '>>=',
get '**='() {
throw new Error('The `**=` operator is unsupported in TypeGPU functions.');
},
get '??='() {
throw new Error('The `??=` operator is unsupported in TypeGPU functions.');
},
get '&&='() {
throw new Error('The `&&=` operator is unsupported in TypeGPU functions.');
},
get '||='() {
throw new Error('The `||=` operator is unsupported in TypeGPU functions.');
},
};
function operatorToType(lhs, op, rhs) {
if (!rhs) {
if (op === '!') {
return bool;
}
return lhs;
}
if (binaryLogicalOps.includes(op)) {
return bool;
}
if (op === '=') {
return rhs;
}
return lhs;
}
const unaryOpCodeToCodegen = {
'-': neg[$gpuCallable].call.bind(neg),
void: () => snip(undefined, wgsl.Void, 'constant', false),
'!': (ctx, [argExpr]) => {
if (argExpr === undefined) {
throw new Error('The unary operator `!` expects 1 argument, but 0 were provided.');
}
if (isKnownAtComptime(argExpr)) {
return snip(!argExpr.value, bool, 'constant', false);
}
const argStr = ctx.resolveSnippet(argExpr).value;
if (!wgsl.isBool(argExpr.dataType)) {
throw new WgslTypeError(`Unary operator ! requires boolean operand. Got ${String(argExpr.dataType)}.${wgsl.isVecBool(argExpr.dataType)
? ` For component-wise negation, use 'std.${not.toString()}'.`
: ''}`);
}
return snip(`!(${argStr})`, bool, 'runtime', argExpr.possibleSideEffects);
},
};
const binaryOpCodeToCodegen = {
'+': add[$gpuCallable].call.bind(add),
'-': sub[$gpuCallable].call.bind(sub),
'*': mul[$gpuCallable].call.bind(mul),
'/': div[$gpuCallable].call.bind(div),
'**': pow[$gpuCallable].call.bind(pow),
};
const usageToVarTemplateMap = {
private: 'private',
workgroup: 'workgroup',
uniform: 'uniform',
mutable: 'storage, read_write',
readonly: 'storage, read',
};
/**
* The block depth that we can expect when generating code in the function scope, not in any nested blocks.
*/
const functionInitialBlockDepth = 2;
export class WgslGenerator {
#ctx = undefined;
// used to detect `continue` and `break` nodes in loop body, as well as label
// unrolled blocks with comments
#unrollingChain = [];
static {
WgslGenerator.prototype.languageKey = 'wgsl';
}
initGenerator(ctx) {
if (this.#ctx !== undefined) {
throw new Error(`Cannot initialize shader generators twice. Create one generator per resolution.`);
}
this.#ctx = ctx;
}
get ctx() {
if (!this.#ctx) {
throw new Error('WGSL Generator has not yet been initialized. Please call initialize(ctx) before using the generator.');
}
return this.#ctx;
}
_block([_, statementNodes], allowInlining, externalMap) {
this.ctx.pushBlockScope();
try {
if (externalMap) {
const externals = Object.fromEntries(Object.entries(externalMap).map(([id, value]) => [id, coerceToSnippet(value)]));
this.ctx.setBlockExternals(externals);
}
let body = '';
/**
* True if any of the statements in the block define variables that would
* be scoped to the currently generated block. If not, we can safely inline it.
*/
let definesInNearestScope = false;
let endsWithControlFlow;
this.ctx.indent();
for (const statementNode of statementNodes) {
const statement = this._statement(statementNode);
if (statement.code.length > 0) {
body += `${statement.code}\n`;
}
definesInNearestScope ||= statement.definesInNearestScope ?? false;
if (statement.endsWithControlFlow !== undefined) {
endsWithControlFlow = statement.endsWithControlFlow;
break;
}
}
this.ctx.dedent();
const willInline = allowInlining && !definesInNearestScope;
// Omitting the 'return;' at the end of the statement list if
// the 'return;' would be placed in the function body outside
// of any nested block.
if (this.ctx.blockDepth === functionInitialBlockDepth) {
body = body.replace(/[ ]*return\s*;\s*$/u, '');
}
if (body === '') {
return { code: '', endsWithControlFlow, definesInNearestScope: false };
}
if (willInline) {
return {
code: this.ctx.getDedented(body.trim()),
endsWithControlFlow,
definesInNearestScope,
};
}
return {
code: `{\n${body}${this.ctx.pre}}`,
endsWithControlFlow,
// all defines will be scoped to the newly generated block
definesInNearestScope: false,
};
}
finally {
this.ctx.popBlockScope();
}
}
_blockStatement(block, externalMap) {
const { code, ...properties } = this._block(block, /* allowInlining */ true, externalMap);
if (code === '') {
return { ...properties, code: '' };
}
return { ...properties, code: `${this.ctx.pre}${code}` };
}
refVariable(id, dataType) {
const varName = this.ctx.makeUniqueIdentifier(id, 'block');
const ptrType = ptrFn(dataType);
const snippet = snip(new RefOperator(snip(varName, dataType, 'function', false), ptrType), ptrType, 'function', false);
this.ctx.defineVariable(id, snippet);
return varName;
}
/**
* Creates a variable declaration string.
* `keyword` may be a placeholder filled in later.
*/
_emitVarDecl(keyword, name, _dataType, rhsStr) {
return `${this.ctx.pre}${keyword} ${name} = ${rhsStr};`;
}
_identifier(id) {
if (!id) {
throw new Error('Cannot resolve an empty identifier');
}
if (id === 'undefined') {
return snip(undefined, wgsl.Void, 'constant', false);
}
const res = this.ctx.getById(id);
if (!res) {
throw new Error(`Identifier ${id} not found`);
}
return res;
}
_callShellless(callee, args) {
const isGeneric = isGenericFn(callee);
const slotPairs = isGeneric ? (callee[$providing]?.pairs ?? []) : [];
const callback = isGeneric ? callee[$internal].inner : callee;
const shelllessCall = this.ctx.withRenamed(callback, getName(callee), () => this.ctx.withSlots(slotPairs, () => {
const shellless = this.ctx.shelllessRepo.get(callback, args);
if (!shellless) {
return undefined;
}
const converted = args.map((s, idx) => {
const argType = shellless.argTypes[idx];
return tryConvertSnippet(this.ctx, s, argType, /* verbose */ false);
});
return this.ctx.withResetIndentLevel(() => {
const snippet = this.ctx.resolve(shellless);
return snip(stitch `${snippet.value}(${converted})`, snippet.dataType,
/* origin */ 'runtime');
});
}));
return shelllessCall;
}
/**
* A wrapper for `generateExpression` that updates `ctx.expectedType`
* and tries to convert the result when it does not match the expected type.
*/
_typedExpression(expression, expectedType) {
const prevExpectedType = this.ctx.expectedType;
this.ctx.expectedType = expectedType;
try {
const result = this._expression(expression);
if (expectedType instanceof AutoStruct) {
// We provide a certain AutoStruct object to later
// investigate what props were accessed. No need to
// convert the result.
return result;
}
return tryConvertSnippet(this.ctx, result, expectedType);
}
finally {
this.ctx.expectedType = prevExpectedType;
}
}
_expression(expression) {
if (typeof expression === 'string') {
return this._identifier(expression);
}
if (typeof expression === 'boolean') {
return snip(expression, bool, /* origin */ 'constant', false);
}
if (expression[0] === NODE.logicalExpr) {
const [_, lhs, op, rhs] = expression;
const lhsExpr = this._expression(lhs);
// Short Circuit Evaluation
if (isKnownAtComptime(lhsExpr)) {
const castToBool = wgsl.isBool(this.ctx.expectedType);
const evalRhs = op === '&&' ? lhsExpr.value : !lhsExpr.value;
if (!evalRhs) {
return castToBool
? snip(op === '||', bool, 'constant', false)
: coerceToSnippet(lhsExpr.value);
}
const rhsExpr = this._expression(rhs);
if (isKnownAtComptime(rhsExpr)) {
const rhsSnippet = coerceToSnippet(rhsExpr.value);
return castToBool ? tryConvertSnippet(this.ctx, rhsSnippet, bool, false) : rhsSnippet;
}
if (rhsExpr.dataType === UnknownData) {
throw new WgslTypeError(`Right-hand side of '${op}' is of unknown type`);
}
// we can skip lhs
return castToBool ? tryConvertSnippet(this.ctx, rhsExpr, bool, false) : rhsExpr;
}
const rhsExpr = this._expression(rhs);
// they are not known at comptime
if (lhsExpr.dataType === UnknownData) {
throw new WgslTypeError(`Left-hand side of '${op}' is of unknown type`);
}
if (!isKnownAtComptime(rhsExpr) && rhsExpr.dataType === UnknownData) {
throw new WgslTypeError(`Right-hand side of '${op}' is of unknown type`);
}
if (!wgsl.isBool(lhsExpr.dataType) || !wgsl.isBool(rhsExpr.dataType)) {
throw new WgslTypeError(`Logical expression '${op}' requires boolean operands. Got '${String(lhsExpr.dataType)}' and '${String(rhsExpr.dataType)}'.`);
}
const lhsStr = this.ctx.resolveSnippet(lhsExpr).value;
const rhsStr = this.ctx.resolveSnippet(rhsExpr).value;
// hardcoded parentheses - operators not present in `parenthesizedOps`
return snip(`(${lhsStr} ${op} ${rhsStr})`, bool, 'runtime', lhsExpr.possibleSideEffects || rhsExpr.possibleSideEffects);
}
if (expression[0] === NODE.binaryExpr || expression[0] === NODE.assignmentExpr) {
// Binary/Assignment Expression
const [exprType, lhs, op, rhs] = expression;
const lhsExpr = this._expression(lhs);
const rhsExpr = this._expression(rhs);
if (rhsExpr.value instanceof RefOperator) {
throw new WgslTypeError(stitch `Cannot assign a ref to an existing variable '${stringifyNode(lhs)}', define a new variable instead.`);
}
if (op === '==') {
throw new Error('Please use the === operator instead of ==');
}
if (op === '!=') {
throw new Error('Please use the !== operator instead of !=');
}
const stdBinaryRelationalOp = binaryRelationalOpToStdMap[op];
if (stdBinaryRelationalOp && isKnownAtComptime(lhsExpr) && isKnownAtComptime(rhsExpr)) {
const left = lhsExpr.value;
const right = rhsExpr.value;
switch (op) {
case '===':
return snip(left === right, bool, 'constant', false);
case '!==':
return snip(left !== right, bool, 'constant', false);
}
if (typeof left !== 'number' || typeof right !== 'number') {
const bothVectors = wgsl.isVec(lhsExpr.dataType) && wgsl.isVec(rhsExpr.dataType);
throw new WgslTypeError(`Comparison '${op}' requires numeric operands.${bothVectors
? ` For component-wise comparison, use 'std.${stdBinaryRelationalOp}'.`
: ''}`);
}
switch (op) {
case '<':
return snip(left < right, bool, 'constant', false);
case '<=':
return snip(left <= right, bool, 'constant', false);
case '>':
return snip(left > right, bool, 'constant', false);
case '>=':
return snip(left >= right, bool, 'constant', false);
}
}
if (lhsExpr.dataType === UnknownData) {
throw new WgslTypeError(`Left-hand side of '${op}' is of unknown type`);
}
if (rhsExpr.dataType === UnknownData) {
throw new WgslTypeError(`Right-hand side of '${op}' is of unknown type`);
}
const codegen = binaryOpCodeToCodegen[op];
if (codegen) {
return codegen(this.ctx, [lhsExpr, rhsExpr]);
}
let convLhs;
let convRhs;
if (bitShiftOps.includes(op)) {
const lhsDataType = lhsExpr.dataType;
if (!wgsl.isInteger(lhsDataType) && !wgsl.isIntegerVec(lhsDataType)) {
throw new WgslTypeError(`Expression: ${stringifyNode(expression)}\nLeft-hand side of '${op}' must be an integer or vector of integers.\nGot ${this.ctx.resolve(lhsDataType).value}.`);
}
const lhsPrimitive = wgsl.isVec(lhsDataType) ? lhsDataType.primitive : lhsDataType;
if (['>>>', '>>>='].includes(op) && lhsPrimitive.type !== 'u32') {
throw new WgslTypeError(`Expression: ${stringifyNode(expression)}\nLeft-hand side of '${op}' must be an unsigned integer or vector of unsigned integers.\nGot ${this.ctx.resolve(lhsDataType).value}.\nUse ${op.slice(1)} instead.`);
}
if (['>>', '>>='].includes(op) && lhsPrimitive.type === 'u32') {
logger.warn('deprecated', `\nExpression: ${stringifyNode(expression)}\nUsing u32 or vecN<u32> as left-hand side of ${op} is deprecated.\nUse >${op} instead.`);
}
// rhs must be u32 (or vecN<u32> for vector lhs) according to the WGSL spec
let rhsTarget;
if (wgsl.isVec(lhsDataType)) {
const cc = lhsDataType.componentCount;
rhsTarget = cc === 2 ? vec2u : cc === 3 ? vec3u : vec4u;
}
else {
rhsTarget = u32;
}
convRhs = tryConvertSnippet(this.ctx, rhsExpr, rhsTarget, false);
convLhs = lhsExpr;
}
else {
const forcedType = exprType === NODE.assignmentExpr ? [lhsExpr.dataType] : undefined;
[convLhs, convRhs] = convertToCommonType(this.ctx, [lhsExpr, rhsExpr], forcedType) ?? [
lhsExpr,
rhsExpr,
];
}
const type = operatorToType(convLhs.dataType, op, convRhs.dataType);
if (exprType === NODE.assignmentExpr) {
validateSnippetMutation(convLhs, expression);
this.tryMarkModified(lhs);
// Compound assignment operators are okay, e.g. +=, -=, *=, /=, ...
if (op === '=' && isAlias(rhsExpr) && !wgsl.isNaturallyEphemeral(rhsExpr.dataType)) {
throw new WgslTypeError(`'${stringifyNode(expression)}' is invalid, because references cannot be assigned.\n-----\nTry '${stringifyNode(lhs)} = ${this.ctx.resolve(unptr(rhsExpr.dataType)).value}(${stringifyNode(rhs)})' to copy the value instead.\n-----`);
}
}
if (stdBinaryRelationalOp) {
const equalityCheck = ['===', '!=='].includes(op);
const correctOperandTypes = (wgsl.isNumericSchema(convLhs.dataType) && wgsl.isNumericSchema(convRhs.dataType)) ||
(equalityCheck && wgsl.isBool(convLhs.dataType) && wgsl.isBool(convRhs.dataType));
if (!correctOperandTypes) {
const bothVectors = wgsl.isVec(convLhs.dataType) && wgsl.isVec(convRhs.dataType);
throw new WgslTypeError(`Comparison '${op}' requires numeric${equalityCheck ? ' or boolean' : ''} operands. Got '${String(convLhs.dataType)}' and '${String(convRhs.dataType)}'.${bothVectors
? ` For component-wise comparison, use 'std.${stdBinaryRelationalOp}'.`
: ''}`);
}
}
return snip(this.emitBinaryOp(convLhs, (OP_MAP[op] ?? op), convRhs), type,
// Result of an operation, so not a reference to anything
/* origin */ 'runtime', exprType === NODE.assignmentExpr ||
lhsExpr.possibleSideEffects ||
rhsExpr.possibleSideEffects);
}
if (expression[0] === NODE.postUpdate) {
throw new Error(`'${stringifyNode(expression)}' is invalid because update is only allowed as a statement.`);
}
if (expression[0] === NODE.unaryExpr) {
// Unary Expression
const [_, op, arg] = expression;
const argExpr = this._expression(arg);
const codegen = unaryOpCodeToCodegen[op];
if (codegen) {
return codegen(this.ctx, [argExpr]);
}
const argStr = this.ctx.resolveSnippet(argExpr).value;
const type = operatorToType(argExpr.dataType, op);
// Result of an operation, so not a reference to anything
return snip(`${op}${argStr}`, type, /* origin */ 'runtime', argExpr.possibleSideEffects);
}
if (expression[0] === NODE.memberAccess) {
// Member Access
const [_, targetNode, property] = expression;
const target = this._expression(targetNode);
const accessed = accessProp(target, property);
if (!accessed) {
throw new Error(`Property '${property}' not found on '${stringifyNode(targetNode)}'`);
}
return accessed;
}
if (expression[0] === NODE.indexAccess) {
// Index Access
const [_, targetNode, propertyNode] = expression;
const target = this._expression(targetNode);
const inProperty = this._expression(propertyNode);
const property = convertToCommonType(this.ctx, [inProperty], [u32, i32], /* verbose */ false)?.[0] ??
inProperty;
const accessed = accessIndex(target, property);
if (!accessed) {
throw new Error(`Index access '${stringifyNode(expression)}' is invalid. If the value is an array, to address this, consider one of the following approaches: (1) declare the array using 'tgpu.const', (2) store the array in a buffer, or (3) define the array within the GPU function scope.`);
}
return accessed;
}
if (expression[0] === NODE.numericLiteral) {
// Numeric Literal
const type = typeof expression[1] === 'string'
? numericLiteralToSnippet(parseNumericString(expression[1]))
: numericLiteralToSnippet(expression[1]);
invariant(type, `Expected ${stringifyNode(expression)} to be valid numeric literal`);
return type;
}
if (expression[0] === NODE.call) {
// Function Call
const [_, calleeNode, argNodes] = expression;
const _callee = this._expression(calleeNode);
const callee = mathToStd.has(_callee.value)
? snip(mathToStd.get(_callee.value), UnknownData, 'runtime', _callee.possibleSideEffects)
: _callee;
if (supportedLogOps().includes(callee.value)) {
return this.ctx.generateLog(callee.value, argNodes.map((arg) => this._expression(arg)));
}
if (wgsl.isWgslStruct(callee.value)) {
// Struct schema call.
if (argNodes.length > 1) {
throw new WgslTypeError('Struct schemas should always be called with at most 1 argument');
}
// No arguments `Struct()`, resolve struct name and return.
if (!argNodes[0]) {
// The schema becomes the data type.
return snip(`${this.ctx.resolve(callee.value).value}()`, callee.value,
// A new struct, so not a reference.
/* origin */ 'runtime', false);
}
const arg = this._typedExpression(argNodes[0], callee.value);
// Either `Struct({ x: 1, y: 2 })`, or `Struct(otherStruct)`.
// In both cases, we just let the argument resolve everything.
return snip(this.ctx.resolveSnippet(arg).value, callee.value,
// A new struct, so not a reference.
/* origin */ 'runtime', arg.possibleSideEffects);
}
if (wgsl.isWgslArray(callee.value)) {
// Array schema call.
if (argNodes.length > 1) {
throw new WgslTypeError('Array schemas should always be called with at most 1 argument');
}
// No arguments `array<...>()`, resolve array type and return.
if (!argNodes[0]) {
// The schema becomes the data type.
return this.typeInstantiation(callee.value, []);
}
const arg = this._typedExpression(argNodes[0], callee.value);
// `d.arrayOf(...)([...])`.
// We don't resolve the ArrayExpression object itself to
// avoid reference checks (we're copying so it's fine)
if (arg.value instanceof ArrayExpression) {
return this.typeInstantiation(callee.value, arg.value.elements);
}
// `d.arrayOf(...)(otherArr)`.
// We just let the argument resolve everything.
return snip(this.ctx.resolveSnippet(arg).value, callee.value,
// A new array, so not a reference.
/* origin */ 'runtime', arg.possibleSideEffects);
}
if (callee.value === constant) {
throw new Error('Constants cannot be defined within TypeGPU function scope. To address this, move the constant definition outside the function scope.');
}
if (isInfixDispatch(callee.value)) {
if (!argNodes[0]) {
throw new WgslTypeError(`An infix operator '${getName(callee.value.operator)}' was called without any arguments`);
}
const lhs = coerceToSnippet(callee.value.lhs);
const rhs = this._expression(argNodes[0]);
const callable = callee.value.operator[$gpuCallable];
return callable.call(this.ctx, [lhs, rhs]);
}
if ((callee.value === _ref || callee.value === unroll) && argNodes[0]) {
this.tryMarkModified(argNodes[0]);
}
if (isGPUCallable(callee.value)) {
const callable = callee.value[$gpuCallable];
const strictSignature = callable.strictSignature;
let convertedArguments;
if (strictSignature) {
// The function's signature does not depend on the context, so it can be used to
// give a hint to the argument expressions that a specific type is expected.
convertedArguments = argNodes.map((arg, i) => {
const argType = strictSignature.argTypes[i];
if (!argType) {
throw new WgslTypeError(`Call '${stringifyNode(expression)}' is invalid since the function expected fewer arguments`);
}
return this._typedExpression(arg, argType);
});
}
else {
convertedArguments = argNodes.map((arg) => this._expression(arg));
}
try {
return callable.call(this.ctx, convertedArguments);
}
catch (err) {
if (err instanceof ResolutionError) {
throw err;
}
throw new ResolutionError(err, [
{
toString: () => `fn:${getName(callee.value)}`,
},
]);
}
}
if (!isMarkedInternal(callee.value) || isGenericFn(callee.value)) {
const args = argNodes.map((arg) => this._expression(arg));
const result = this._callShellless(callee.value, args);
if (result) {
return result;
}
}
// try to throw a descriptive error
const maybeMathMethod = Object.getOwnPropertyNames(Math).find((prop) => Math[prop] === callee.value);
if (maybeMathMethod) {
throw new Error(`Unsupported Math functionality 'Math.${maybeMathMethod}()'. Use an std alternative, or implement the function manually.`);
}
const maybeConsoleMethod = Object.getOwnPropertyNames(console).find((prop) => console[prop] === callee.value);
if (maybeConsoleMethod) {
throw new Error(`Unsupported console functionality 'console.${maybeConsoleMethod}()'.`);
}
throw new Error(`Function '${getName(callee.value) ?? String(callee.value)}' is not marked with the 'use gpu' directive and cannot be used in a shader`);
}
if (expression[0] === NODE.objectExpr) {
// Object Literal
const obj = expression[1];
const structType = this.ctx.expectedType;
if (structType instanceof AutoStruct) {
const entries = Object.fromEntries(Object.entries(obj).map(([key, value]) => {
let accessed = structType.accessProp(key);
let expr;
if (accessed) {
// Generating the expression expecting a specific type
expr = this._typedExpression(value, accessed.type);
}
else {
// Generating the expression and inferring the type instead
expr = this._expression(value);
if (expr.dataType === UnknownData) {
throw new WgslTypeError(stitch `Property ${key} in object literal has a value of unknown type: '${expr}'`);
}
// Taking care of abstract numerics and implicit pointers
accessed = structType.provideProp(key, unptr(concretize(expr.dataType)));
}
return [accessed.prop, expr];
}));
const completeStruct = structType.completeStruct;
const convertedSnippets = convertStructValues(this.ctx, completeStruct, entries);
return snip(stitch `${this.ctx.resolve(structType).value}(${convertedSnippets})`, completeStruct,
/* origin */ 'runtime');
}
if (wgsl.isWgslStruct(structType)) {
const entries = Object.fromEntries(Object.entries(structType.propTypes).map(([key, value]) => {
const val = obj[key];
if (val === undefined) {
throw new WgslTypeError(`Missing property ${key} in object literal for struct ${structType}`);
}
const result = this._typedExpression(val, value);
return [key, result];
}));
const convertedSnippets = convertStructValues(this.ctx, structType, entries);
return snip(stitch `${this.ctx.resolve(structType).value}(${convertedSnippets})`, structType,
/* origin */ 'runtime', convertedSnippets.some((s) => s.possibleSideEffects));
}
throw new WgslTypeError(`No target type could be inferred for object '${stringifyNode(expression)}', please wrap the object in the corresponding schema.`);
}
if (expression[0] === NODE.arrayExpr) {
const [_, valueNodes] = expression;
// Array Expression
const arrType = this.ctx.expectedType;
let elemType;
let values;
if (wgsl.isWgslArray(arrType)) {
elemType = arrType.elementType;
// The array is typed, so its elements should be as well.
values = valueNodes.map((value) => this._typedExpression(value, elemType));
// Since it's an expected type, we enforce the length
if (values.length !== arrType.elementCount) {
throw new WgslTypeError(`Cannot create value of type '${arrType}' from an array of length: ${values.length}`);
}
}
else {
// The array is not typed, so we try to guess the types.
const valuesSnippets = valueNodes.map((value) => this._expression(value));
if (valuesSnippets.length === 0) {
throw new WgslTypeError('Cannot infer the type of an empty array literal.');
}
const converted = convertToCommonType(this.ctx, valuesSnippets);
if (!converted) {
throw new WgslTypeError(`Values '${stringifyNode(expression)}' cannot be automatically converted to a common type. Consider wrapping the array in an appropriate schema`);
}
values = converted;
elemType = concretize(values[0]?.dataType);
}
const arrayType = arrayOf(elemType, values.length);
const allConstant = values.every((value) => value.origin === 'constant');
return snip(new ArrayExpression(arrayType, values), arrayType,
/* origin */ allConstant ? 'constant' : 'runtime', values.some((v) => v.possibleSideEffects));
}
if (expression[0] === NODE.conditionalExpr) {
// ternary operator
const [_, testNode, consequentNode, alternativeNode] = expression;
const test = this._expression(testNode);
if (isKnownAtComptime(test)) {
return test.value ? this._expression(consequentNode) : this._expression(alternativeNode);
}
else {
const convertedTest = tryConvertSnippet(this.ctx, test, bool, false);
const consequent = this._expression(consequentNode);
const alternative = this._expression(alternativeNode);
const [con, alt] = convertToCommonType(this.ctx, [consequent, alternative], validSelectBranchTypes) ?? [];
if (!con ||
!alt ||
consequent.possibleSideEffects ||
alternative.possibleSideEffects ||
(isAlias(consequent) && !wgsl.isNaturallyEphemeral(consequent.dataType)) ||
(isAlias(alternative) && !wgsl.isNaturallyEphemeral(alternative.dataType))) {
throw new Error(`Ternary operator '${stringifyNode(expression)}' is invalid. For more complex branching, please use 'std.select' or if/else statements.`);
}
return snip(stitch `select(${alt}, ${con}, ${convertedTest})`, con.dataType, 'runtime',
// this select has side-effects only if the condition has side-effects
test.possibleSideEffects);
}
}
if (expression[0] === NODE.stringLiteral) {
return snip(expression[1], UnknownData, /* origin */ 'constant', false);
}
if (expression[0] === NODE.preUpdate) {
throw new Error('Cannot use pre-updates in TypeGPU functions.');
}
assertExhaustive(expression);
}
declareGlobalConst(options) {
const resolvedDataType = this.ctx.resolve(options.dataType).value;
const resolvedValue = this.ctx.resolveSnippet(options.init).value;
this.ctx.addDeclaration(`const ${options.id}: ${resolvedDataType} = ${resolvedValue};`, options.id);
return snip(options.id, options.dataType, 'constant-immutable-def');
}
declareGlobalVar(options) {
let pre = '';
if (options.group !== undefined) {
pre += `@group(${options.group}) `;
}
if (options.binding !== undefined) {
pre += `@binding(${options.binding}) `;
}
if (options.scope in usageToVarTemplateMap) {
pre += `var<${usageToVarTemplateMap[options.scope]}> `;
}
else {
pre += `var `;
}
pre += `${options.id}: ${this.ctx.resolve(options.dataType).value}`;
this.ctx.addDeclaration(options.init ? `${pre} = ${this.ctx.resolveSnippet(options.init).value};` : `${pre};`, options.id);
return snip(options.id, options.dataType, options.scope);
}
functionDefinition(options) {
// Function body
invariant(this.ctx.blockDepth === functionInitialBlockDepth - 1, `Expecting exactly ${functionInitialBlockDepth - 1} block(s) before going into the first function block scope`);
let body = this._block(options.body, /* allowInlining */ false);
const scope = this.ctx.topFunctionScope;
invariant(scope, 'Expected function scope to be present');
const replacements = Object.fromEntries([...scope.placeholderForVariable.entries()].map(([variable, placeholder]) => [
placeholder,
scope.modifiedVariables.has(variable) ? 'var' : 'let',
]));
if (Object.keys(replacements).length > 0) {
const regex = new RegExp(Object.keys(replacements).join('|'), 'gi');
body.code = body.code.replace(regex, (match) => replacements[match] ?? '#ERR');
}
// Only after generating the body can we determine the return type
const returnType = options.determineReturnType();
const argList = options.args
// Stripping out unused arguments in entry functions
.filter((arg) => arg.used || options.functionType === 'normal')
.map((arg) => {
return `${getAttributesString(arg.decoratedType)}${arg.name}: ${this.ctx.resolve(arg.decoratedType).value}`;
})
.join(', ');
const head = returnType.type !== 'void'
? `(${argList}) -> ${getAttributesString(returnType)}${this.ctx.resolve(returnType).value} `
: `(${argList}) `;
let attributes = '';
if (options.functionType === 'compute') {
if (!options.workgroupSize) {
throw new Error('Compute shaders must have a workgroup size');
}
attributes = `@compute @workgroup_size(${options.workgroupSize.join(', ')}) `;
}
else if (options.functionType === 'vertex') {
attributes = `@vertex `;
}
else if (options.functionType === 'fragment') {
attributes = `@fragment `;
}
return `${attributes}fn ${options.name}${head}${body.code || '{}'}`;
}
/**
* Generates a WGSL type string for the given data type, and adds necessary
* definitions to the shader preamble. This shouldn't be called directly, only
* through `ctx.resolve` to properly cache the result.
*/
emitTypeAnnotation(data) {
return resolveData(this.ctx, data);
}
typeInstantiation(schema, args) {
if (args.length === 1 && args[0]?.dataType === schema) {
// Already of the desired type, e.g. `bool(false)` or `vec3f(vec3f(1, 2, 3))`
// We can make this snippet ephemeral, as we know it will be deep copied in JS
return snip(stitch `${args[0]}`, schema, fallthroughCopyOrigin(args[0].origin), args[0].possibleSideEffects);
}
// Creating a 'runtime' snippet, since it's instantiating a new value
return snip(stitch `${this.ctx.resolve(schema).value}(${args})`, schema, 'runtime', args.some((s) => s.possibleSideEffects));
}
numericLiteral(value, schema) {
if (!Number.isFinite(value)) {
throw new Error(`Value '${value}' (${schema.type}) cannot be resolved due to WGSL's Finite Math Assumption (see: https://www.w3.org/TR/WGSL/#finite-math-assumption). This value might be a result of a comptime-evaluated operation.`);
}
if (schema.type === 'abstractInt') {
return snip(`${value}`, schema, /* origin */ 'constant', false);
}
if (schema.type === 'u32') {
return snip(`${value}u`, schema, /* origin */ 'constant', false);
}
if (schema.type === 'i32') {
return snip(`${value}i`, schema, /* origin */ 'constant', false);
}
const exp = value.toExponential();
const decimal = schema.type === 'abstractFloat' && Number.isInteger(value) ? `${value}.` : `${value}`;
// Just picking the shorter one
const base = exp.length < decimal.length ? exp : decimal;
if (schema.type === 'f32') {
return snip(`${base}f`, schema, /* origin */ 'constant', false);
}
if (schema.type === 'f16') {
return snip(`${base}h`, schema, /* origin */ 'constant', false);
}
return snip(base, schema, /* origin */ 'constant', false);
}
emitCall(name, templateParams, args) {
const resolvedTemplateParams = templateParams
.map((arg) => this.ctx.resolveSnippet(arg).value)
.join(', ');
const resolvedArgs = args.map((arg) => this.ctx.resolveSnippet(arg).value).join(', ');
if (resolvedTemplateParams.length > 0) {
return `${name}<${resolvedTemplateParams}>(${resolvedArgs})`;
}
return `${name}(${resolvedArgs})`;
}
emitBinaryOp(lhs, op, rhs) {
const lhsStr = this.ctx.resolveSnippet(lhs).value;
const rhsStr = this.ctx.resolveSnippet(rhs).value;
return parenthesizedOps.includes(op)
? `(${lhsStr} ${op} ${rhsStr})`
: `${lhsStr} ${op} ${rhsStr}`;
}
_return(statement) {
const returnNode = statement[1];
if (returnNode !== undefined) {
const expectedReturnType = this.ctx.topFunctionReturnType;
let returnSnippet = expectedReturnType
? this._typedExpression(returnNode, expectedReturnType)
: this._expression(returnNode);
if (returnSnippet.value === undefined && wgsl.isVoid(returnSnippet.dataType)) {
this.ctx.reportReturnType(wgsl.Void);
return `${this.ctx.pre}return;`;
}
if (returnSnippet.value instanceof RefOperator) {
throw new WgslTypeError(`Cannot return '${stringifyNode(returnNode)}' because it is a d.ref`);
}
// Arguments cannot be returned from functions without copying. A simple example why is:
// const identity = (x) => {
// 'use gpu';
// return x;
// };
//
// const foo = (arg: d.v3f) => {
// 'use gpu';
// const marg = identity(arg);
// marg.x = 1; // 'marg's origin would be 'runtime', so we wouldn't be able to track this misuse.
// };
if (returnSnippet.origin === 'argument' &&
!wgsl.isNaturallyEphemeral(returnSnippet.dataType) &&
// Only restricting this use in non-entry functions, as the function
// is giving up ownership of all references anyway.
this.ctx.topFunctionScope?.functionType === 'normal') {
throw new WgslTypeError(`'${stringifyNode(statement)}' is invalid, cannot return references to arguments. Copy the argument before returning it.`);
}
if (
// The existence of `expectedReturnType` implies a function shell, which in turn implies that the
// value will be copied on return anyway
!expectedReturnType &&
isAlias(returnSnippet) &&
!wgsl.isNaturallyEphemeral(returnSnippet.dataType) &&
returnSnippet.origin !== 'local-def') {
const str = stringifyNode(returnNode);
const typeStr = this.ctx.resolve(unptr(returnSnippet.dataType)).value;
throw new WgslTypeError(`'return ${str};' is invalid, cannot return references.
-----
Try 'return ${typeStr}(${str});' instead.
-----`);
}
returnSnippet = tryConvertSnippet(this.ctx, returnSnippet, unptr(returnSnippet.dataType), false);
invariant(returnSnippet.dataType !== UnknownData, 'Return type should be known');
this.ctx.reportReturnType(returnSnippet.dataType);
return stitch `${this.ctx.pre}return ${returnSnippet};`;
}
this.ctx.reportReturnType(wgsl.Void);
return `${this.ctx.pre}return;`;
}
_letStatement(statement) {
const [_, rawId, eqNode] = statement;
if (eqNode === undefined) {
throw new Error(`'${stringifyNode(statement)}' is invalid because all variables need initializers.`);
}
const eq = this._expression(eqNode);
if (eq.value instanceof RefOperator) {
const rhsStr = stringifyNode(eqNode);
throw new WgslTypeError(`'let ${rawId} = ${rhsStr}' is invalid, cannot initialize 'let' variables with d.ref()
-----
- Try 'const ${rawId} = ${rhsStr}'.
-----`);
}
const definitionDataType = eq.dataType;
if (definitionDataType === UnknownData) {
const rhsStr = stringifyNode(eqNode);
throw new WgslTypeError(`'let ${rawId} = ${rhsStr}' is invalid, cannot determine WGSL type of '${rhsStr}'
-----
- Try using or defining a schema that matches your desired value the most, and wrap the value with it: 'let ${rawId} = Schema(${rhsStr})'
-----`);
}
if (isAlias(eq) && !wgsl.isNaturallyEphemeral(eq.dataType)) {
// `let` declarations cannot store references
const rhsStr = stringifyNode(eqNode);
const rhsTypeStr = this.ctx.resolve(unptr(eq.dataType)).value;
throw new WgslTypeError(`'let ${rawId} = ${rhsStr}' is invalid, because references cannot be assigned to 'let' variable declarations.
-----
- Try 'let ${rawId} = ${rhsTypeStr}(${rhsStr})' if you need to reassign '${rawId}' later
- Try 'const ${rawId} = ${rhsStr}' if you won't reassign '${rawId}' later.
-----`);
}
const concreteType = concretize(definitionDataType);
const snippet = snip(this.ctx.makeUniqueIdentifier(rawId, 'block'), concreteType,
/* origin */ 'local-def', false);
this.ctx.defineVariable(rawId, snippet);
const rhsSnippet = tryConvertSnippet(this.ctx, eq, definitionDataType, false);
const rhsStr = this.ctx.resolveSnippet(rhsSnippet).value;
// Even though the user defined a 'let' (expecting it to be reassigned), the
// reassignment might happen in a pruned branch, in which case we can generate
// more optimised code by emitting 'let' or 'const' instead of 'var'.
const scope = this.ctx.topFunctionScope;
invariant(scope, `Expected function scope to be present for ${rawId}`);
const emittedVarType = `#VAR_${scope.placeholderForVariable.size}#`;
scope.placeholderForVariable.set(snippet, emittedVarType);
return {
code: this._emitVarDecl(emittedVarType, snippet.value, concreteType, rhsStr),
definesInNearestScope: true,
};
}
_constStatement(statement) {
const [_, rawId, eqNode] = statement;
if (eqNode === undefined) {
throw new Error(`'${stringifyNode(statement)}' is invalid because all variables need initializers.`);
}
const eq = this._expression(eqNode);
if (eq.value instanceof RefOperator) {
// We're assigning a newly created `d.ref()`
if (eq.dataType !== UnknownData) {
throw new WgslTypeError(`Cannot store d.ref() in a variable if it references another value. Copy the value passed into d.ref() instead.`);
}
const refSnippet = eq.value.snippet;
const varName = this.refVariable(rawId, concretize(refSnippet.dataType));
return {
code: stitch `${this.ctx.pre}var ${varName} = ${tryConvertSnippet(this.ctx, refSnippet, refSnippet.dataType, false)};`,
definesInNearestScope: true,
};
}
const rhsNaturallyEphemeral = wgsl.isNaturallyEphemeral(eq.dataType);
let varOrigin = 'local-def';
let varType = '<deferred>';
let definitionDataType = eq.dataType;
if (definitionDataType === UnknownData) {
const rhsStr = stringifyNode(eqNode);
throw new WgslTypeError(`'const ${rawId} = ${rhsStr}' is invalid, cannot determine WGSL type of '${rhsStr}'
-----
- Try using or defining a schema that matches your desired value the most, and wrap the value with it: 'const ${rawId} = Schema(${rhsStr})'
-----`);
}
if (eq.origin === 'argument') {
// Arguments are immutable, so we 'let' them be (kill me)
varType = 'let';
// When we declare a new variable with a naturally ephemeral value (e.g. a scalar)
// the variable now loses the restrictions of an argument, and becomes just a regular
// variable. For vectors and other non-naturally ephemeral values, the restrictions of
// arguments are kept.
varOrigin = rhsNaturallyEphemeral ? 'local-def' : 'argument';
}
else if (eq.origin === 'constant-immutable-def') {
varType = 'const';
varOrigin = 'constant-immutable-def';
}
else if (eq.origin === 'runtime-immutable-def') {
varType = 'let';
varOrigin = 'runtime-immutable-def';
}
else if (rhsNaturallyEphemeral) {
varType = eq.origin === 'constant' ? 'const' : 'let';
// Constants are also local declarations. We lose some information here, meaning
// when we look at a variable's snippet, we cannot tell if it's a constant or not.
// This is mostly because we plan to determine this fact later, after all of the
// function code has been processed, so at least currently, we lose that info.
varOrigin = 'local-def';
}
else if (!isAlias(eq)) {
// Not a reference, but also not naturally ephemeral, so we cannot guarantee it won't be mutated.
// We defer the decision for now.
varType = '<deferred>';
varOrigin = 'local-def';
}
else {
return this._aliasConstStatement(rawId, eqNode, eq);
}
const concreteType = concretize(definitionDataType);
const snippet = snip(this.ctx.makeUniqueIdentifier(rawId, 'block'), concreteType,
/* origin */ varOrigin, false);
this.ctx.defineVariable(rawId, snippet);
const rhsSnippet = tryConvertSnippet(this.ctx, eq, definitionDataType, false);
const rhsStr = this.ctx.resolveSnippet(rhsSnippet).value;
let emittedVarType;
if (varType === '<deferred>') {
const scope = this.ctx.topFunctionScope;
invariant(scope, `Expected function scope to be present for ${rawId}`);
emittedVarType = `#VAR_${scope.placeholderForVariable.size}#`;
scope.placeholderForVariable.set(snippet, emittedVarType);
}
else {
emittedVarType = varType;
}
return {
code: this._emitVarDecl(emittedVarType, snippet.value, concreteType, rhsStr),
definesInNearestScope: true,
};
}
/**
* Handles `const x = <rhs>;` declarations in which the right-hand side aliases memory
* that outlives the expression (a buffer, a local variable, an array element, ...).
*
* In WGSL we store an *implicit* pointer to that memory, so mutations done through `x`
* affect the original. Languages without pointers (e.g. GLSL) override this.
*/
_aliasConstStatement(rawId, eqNode, eq) {
// Assigning a reference to a `const` variable means we store the pointer
// of the rhs.
let definitionDataType = eq.dataType;
if (!wgsl.isPtr(definitionDataType)) {
const ptrType = createPtrFromOrigin(eq.origin, concretize(definitionDataType));
invariant(ptrType !== undefined, `Creating pointer type from origin ${eq.origin}`);
definitionDataType = ptrType;
}
// Making the pointer implicit, meaning the fact it's a pointer isn't
// reflected in the JS source code.
definitionDataType = implicitFrom(definitionDataType);
this.tryMarkModified(eqNode);
const concreteType = concretize(definitionDataType);
const snippet = snip(this.ctx.makeUniqueIdentifier(rawId, 'block'), concreteType,
// we pass on the origin
/* origin */ eq.origin, false);
this.ctx.defineVariable(rawId, snippet);
const rhsSnippet = tryConvertSnippet(this.ctx, eq, definitionDataType, false);
const rhsStr = this.ctx.resolveSnippet(rhsSnippet).value;
return {
code: this._emitVarDecl('let', snippet.value, concreteType, rhsStr),
definesInNearestScope: true,
};
}
_statement(statement) {
if (typeof statement === 'string') {
const id = this._identifier(statement);
const resolved = id.value !== undefined && id.value !== null ? this.ctx.resolveSnippet(id).value : '';
return { code: resolved ? `${this.ctx.pre}${resolved};` : '', definesInNearestScope: false };
}
if (typeof statement === 'boolean') {
return {
code: `${this.ctx.pre}${statement ? 'true' : 'false'};`,
definesInNearestScope: false,
};
}
if (statement[0] === NODE.return) {
return {
code: this._return(statement),
endsWithControlFlow: 'return',
definesInNearestScope: false,
};
}
if (statement[0] === NODE.if) {
const [_, condNode, consNode, altNode] = statement;
const condition = this._typedExpression(condNode, bool);
if (typeof condition.value === 'boolean') {
// the condition is known at comptime
let node = condition.value ? consNode : altNode;
if (node === undefined) {
return { code: '', definesInNearestScope: false };
}
if (!Array.isArray(node)) {
node = blockifySingleStatement(node);
}
if (node[0] === NODE.block && node[1].length === 1 && node[1][0][0] === NODE.if) {
// simplify 'if (true) { if (A) {B} } else {C}' to 'if (A) {B}'
return this._statement(node[1][0]);
}
if (node[0] === NODE.if) {
// simplify 'if (false) {A} else if (B) {C}' to 'if (B) {C}'
return this._statement(node);
}
// simplify 'if (true) {A} else {B}' to '{A}'
return this._blockStatement(blockifySingleStatement(node));
}
const consequent = this._block(blockifySingleStatement(consNode), /* allowInlining */ false);
const alternate = !altNode
? undefined
: this._block(blockifySingleStatement(altNode), /* allowInlining */ false).code;
if (!alternate) {
return {
code: stitch `${this.ctx.pre}if (${condition}) ${consequent.code || '{}'}`,
definesInNearestScope: false,
};
}
return {
code: stitch `\
${this.ctx.pre}if (${condition}) ${consequent.code || '{}'}
${this.ctx.pre}else ${alternate}`,
definesInNearestScope: false,
};
}
if (statement[0] === NODE.let) {
return this._letStatement(statement);
}
if (statement[0] === NODE.const) {
return this._constStatement(statement);
}
if (statement[0] === NODE.block) {
return this._blockStatement(statement);
}
if (statement[0] === NODE.for) {
const [_, init, condition, update, body] = statement;
const prevUnrollingChain = this.#unrollingChain;
this.#unrollingChain = [];
try {
this.ctx.pushBlockScope();
const [initStatement, conditionExpr, updateStatement] = this.ctx.withResetIndentLevel(() => [
init ? this._statement(init).code : undefined,
condition ? this._typedExpression(condition, bool) : undefined,
update ? this._statement(update).code : undefined,
]);
const initStr = initStatement ? initStatement.slice(0, -1) : '';
const updateStr = updateStatement ? updateStatement.slice(0, -1) : '';
const bodyStr = this._block(blockifySingleStatement(body), /* allowInlining */ false).code;
return {
code: stitch `${this.ctx.pre}for (${initStr}; ${conditionExpr}; ${updateStr}) ${bodyStr || '{}'}`,
definesInNearestScope: false,
};
}
finally {
this.#unrollingChain = prevUnrollingChain;
this.ctx.popBlockScope();
}
}
if (statement[0] === NODE.while) {
const prevUnrollingChain = this.#unrollingChain;
this.#unrollingChain = [];
try {
const [_, condition, body] = statement;
const condSnippet = this._typedExpression(condition, bool);
const conditionStr = this.ctx.resolveSnippet(condSnippet).value;
const bodyStr = this._block(blockifySingleStatement(body), /* allowInlining */ false).code;
return {
code: `${this.ctx.pre}while (${conditionStr}) ${bodyStr || '{}'}`,
definesInNearestScope: false,
};
}
finally {
this.#unrollingChain = prevUnrollingChain;
}
}
if (statement[0] === NODE.forOf) {
const [_, loopVar, iterable, body] = statement;
if (loopVar[0] !== NODE.const) {
throw new WgslTypeError('Only `for (const ... of ... )` loops are supported');
}
this.tryMarkModified(iterable); // overly-defensive, but let's not tempt fate
let ctxIndent = false;
const prevUnrollingChain = this.#unrollingChain;
try {
this.ctx.pushBlockScope();
const iterableExpr = this._expression(iterable);
const shouldUnroll = iterableExpr.value instanceof UnrollableIterable;
const iterableSnippet = shouldUnroll ? iterableExpr.value.snippet : iterableExpr;
const range = forOfUtils.getRangeSnippets(this.ctx, iterableSnippet, shouldUnroll);
const originalLoopVarName = loopVar[1];
const blockified = blockifySingleStatement(body);
if (shouldUnroll) {
if (!isKnownAtComptime(range.end)) {
throw new Error('Cannot unroll loop. Length of iterable is unknown at comptime.');
}
const length = range.end.value;
if (length === 0) {
return { code: '', definesInNearestScope: false };
}
const { value } = iterableSnippet;
const elements = isTgpuRange(value)
? value.map((i) => coerceToSnippet(i))
: value instanceof ArrayExpression
? value.elements
: Array.from({ length }, (_, i) => forOfUtils.getElementSnippet(iterableSnippet, snip(i, u32, 'constant')));
const firstElement = elements[0];
if (!isAlias(firstElement) && !wgsl.isNaturallyEphemeral(firstElement.dataType)) {
throw new WgslTypeError(`Cannot unroll '${stringifyNode(iterable)}'. The elements of iterable are constructed in place but are not value types.`);
}
let blocksCode = '';
let endsWithControlFlow;
let definesInNearestScope = false;
for (let i = 0; i < elements.length; i++) {
const e = elements[i];
this.#unrollingChain = [...prevUnrollingChain, i];
const resolvedBlock = this._blockStatement(blockified, {
[originalLoopVarName]: e,
});
definesInNearestScope ||= resolvedBlock.definesInNearestScope;
blocksCode += `${this.ctx.pre}// unrolled iteration ${this.#unrollingChain.map((idx) => `#${idx}`).join(' / ')}\n${resolvedBlock.code}\n`;
if (resolvedBlock.endsWithControlFlow !== undefined) {
endsWithControlFlow = resolvedBlock.endsWithControlFlow;
break;
}
}
return {
code: `${blocksCode}${this.ctx.pre}// ---`,
endsWithControlFlow,
definesInNearestScope,
};
}
this.#unrollingChain = [];
const index = this.ctx.makeUniqueIdentifier('i', 'block');
const forHeaderStr = stitch `${this.ctx.pre}for (var ${index} = ${range.start}; ${index} ${range.comparison} ${range.end}; ${index} += ${range.step})`;
let bodyStr = '';
if (isTgpuRange(iterableSnippet.value)) {
bodyStr = this._block(blockified, /* allowInlining */ false, {
[originalLoopVarName]: snip(index, range.start.dataType, 'runtime', false), // range.start, .end , .step have the same dataType
}).code;
}
else {
this.ctx.indent();
ctxIndent = true;
const loopVarName = this.ctx.makeUniqueIdentifier(originalLoopVarName, 'block');
const elementSnippet = forOfUtils.getElementSnippet(iterableSnippet, snip(index, u32, 'runtime'));
const loopVarKind = forOfUtils.getLoopVarKind(elementSnippet);
const elementType = forOfUtils.getElementType(elementSnippet, iterableSnippet);
const loopVarDeclStr = stitch `${this.ctx.pre}${loopVarKind} ${loopVarName} = ${tryConvertSnippet(this.ctx, elementSnippet, elementType, false)};`;
bodyStr = `{\n${loopVarDeclStr}\n${this._blockStatement(blockified, {
[originalLoopVarName]: snip(loopVarName, elementType, elementSnippet.origin, false),
}).code}\n`;
this.ctx.dedent();
bodyStr += `${this.ctx.pre}}`;
ctxIndent = false;
}
return {
code: stitch `${forHeaderStr} ${bodyStr.trim() || '{}'}`,
definesInNearestScope: false,
};
}
finally {
if (ctxIndent) {
this.ctx.dedent();
}
this.#unrollingChain = prevUnrollingChain;
this.ctx.popBlockScope();
}
}
if (statement[0] === NODE.postUpdate) {
// Post-update statement
const [_, op, arg] = statement;
const argExpr = this._expression(arg);
const argStr = this.ctx.resolveSnippet(argExpr).value;
validateSnippetMutation(argExpr, statement);
this.tryMarkModified(arg);
return { code: `${this.ctx.pre}${argStr}${op};`, definesInNearestScope: false };
}
if (statement[0] === NODE.continue) {
if (this.#unrollingChain.length > 0) {
throw new WgslTypeError('Cannot unroll loop containing `continue`');
}
return {
code: `${this.ctx.pre}continue;`,
endsWithControlFlow: 'continue',
definesInNearestScope: false,
};
}
if (statement[0] === NODE.break) {
if (this.#unrollingChain.length > 0) {
throw new WgslTypeError('Cannot unroll loop containing `break`');
}
return {
code: `${this.ctx.pre}break;`,
endsWithControlFlow: 'break',
definesInNearestScope: false,
};
}
const expr = this._expression(statement);
const resolved = expr.value !== undefined && expr.value !== null ? this.ctx.resolveSnippet(expr).value : '';
return { code: resolved ? `${this.ctx.pre}${resolved};` : '', definesInNearestScope: false };
}
/**
* Attempts a member access lookup to mark a variable as modified.
* @example
* // given `let a; a = 1;`
* tryMarkModified('a') // `a` is marked in the function scope
*
* // given `const obj; obj.prop = 1;`
* tryMarkModified('obj.prop') // `obj` is marked in the function scope
*
* // given `this.buffer.$;`
* tryMarkModified('this.buffer.$') // `this` is not marked, since there is no placeholder for it
*/
tryMarkModified(expr) {
if (!expr) {
return;
}
const maybeObject = extractObject(expr);
if (maybeObject !== undefined) {
const snippet = this.ctx.getById(maybeObject);
const scope = this.ctx.topFunctionScope;
if (snippet && scope && scope.placeholderForVariable.has(snippet)) {
scope.modifiedVariables.add(snippet);
}
}
}
}
function validateSnippetMutation(mutated, expr) {
if (mutated.origin === 'constant' ||
mutated.origin === 'constant-immutable-def' ||
mutated.origin === 'runtime-immutable-def') {
if (isKnownAtComptime(mutated)) {
throw new WgslTypeError(`'${stringifyNode(expr)}' is invalid, because the left side is defined outside of the shader, and therefore is immutable during its execution. Try using tgpu.privateVar or buffers.`);
}
throw new WgslTypeError(`'${stringifyNode(expr)}' is invalid, because the left side is a constant.`);
}
if (mutated.origin === 'uniform') {
throw new WgslTypeError(`'${stringifyNode(expr)}' is invalid, because uniform buffers cannot be mutated.`);
}
if (mutated.origin === 'readonly') {
throw new WgslTypeError(`'${stringifyNode(expr)}' is invalid, because readonly buffers cannot be mutated.`);
}
if (mutated.origin === 'argument') {
throw new WgslTypeError(`'${stringifyNode(expr)}' is invalid, because non-pointer arguments cannot be mutated.`);
}
}
function assertExhaustive(value) {
throw new Error(`'${safeStringify(value)}' was not handled by the WGSL generator.`);
}
function parseNumericString(str) {
// Hex literals
if (/^0x[0-9a-f]+$/i.test(str)) {
return Number.parseInt(str);
}
// Binary literals
if (/^0b[01]+$/i.test(str)) {
return Number.parseInt(str.slice(2), 2);
}
return Number.parseFloat(str);
}
function blockifySingleStatement(statement) {
return typeof statement !== 'object' || statement[0] !== NODE.block
? [NODE.block, [statement]]
: statement;
}
function extractObject(expr) {
let object = expr;
while (Array.isArray(object) &&
(object[0] === NODE.memberAccess || object[0] === NODE.indexAccess)) {
object = object[1];
}
if (typeof object === 'string') {
return object;
}
}