diff --git a/packages/typegpu-gl/src/glslGenerator.ts b/packages/typegpu-gl/src/glslGenerator.ts index 493becede8..ed41abaf50 100644 --- a/packages/typegpu-gl/src/glslGenerator.ts +++ b/packages/typegpu-gl/src/glslGenerator.ts @@ -1,5 +1,5 @@ import { NodeTypeCatalog as NODE } from 'tinyest'; -import type { Expression, Return, ObjectExpression, ObjectProperty } from 'tinyest'; +import type { Const, Expression, Return, ObjectExpression, ObjectProperty } from 'tinyest'; import { tgpu, d, type ShaderStage, std } from 'typegpu'; import { abstractInt, @@ -311,6 +311,20 @@ interface EntryFnState { */ const immutableOrigins: readonly Origin[] = ['uniform', 'readonly', 'handle']; +/** + * Adds every array and object in an expression tree to `out`, the node itself included. + * Identifiers are strings and stay out: evaluating one has no side-effects. + */ +function collectObjectNodes(node: unknown, out: Set): Set { + if (typeof node === 'object' && node !== null && !out.has(node)) { + out.add(node); + for (const child of Object.values(node)) { + collectObjectNodes(child, out); + } + } + return out; +} + function undecorateDataType(t: d.BaseData): d.BaseData { return d.isDecorated(t) ? t.inner : t; } @@ -414,6 +428,11 @@ export class GlslGenerator extends WgslGenerator { #functionType: ShaderStage | 'normal' | undefined; #entryFnState: EntryFnState | undefined; #vertexOutPropToVarMap: Record = {}; + /** + * The nodes of the right-hand side of the `const` statement being generated, and the snippets + * they evaluated to. See `_constStatement`. + */ + #constRhs: { nodes: Set; snippets: Map } | undefined; static { GlslGenerator.prototype.languageKey = 'glsl'; @@ -849,6 +868,41 @@ export class GlslGenerator extends WgslGenerator { return super.emitBinaryOp(lhs, op, rhs); } + /** + * `const x = ;` walks its right-hand side a second time in `_aliasConstStatement`, and + * resolves the result a third time. Comptime code in that expression (e.g. `tgpu.comptime` + * calls) must still run once, so every node of the right-hand side keeps the snippet it + * evaluated to the first time. Only those nodes are cached: a function body generated while + * evaluating them can be evaluated again with different argument types. + */ + protected override _constStatement(statement: Const): ResolvedStatement { + const eqNode = statement[2]; + if (eqNode === undefined) { + return super._constStatement(statement); + } + + const previous = this.#constRhs; + this.#constRhs = { nodes: collectObjectNodes(eqNode, new Set()), snippets: new Map() }; + try { + return super._constStatement(statement); + } finally { + this.#constRhs = previous; + } + } + + protected override _expression(expression: Expression): Snippet { + const rhs = this.#constRhs; + if (typeof expression !== 'object' || !rhs?.nodes.has(expression)) { + return super._expression(expression); + } + let snippet = rhs.snippets.get(expression); + if (snippet === undefined) { + snippet = super._expression(expression); + rhs.snippets.set(expression, snippet); + } + return snippet; + } + /** * GLSL has no pointers, so `const x = ;` cannot be turned into an implicit * pointer definition like it is in WGSL. Instead: diff --git a/packages/typegpu-gl/tests/implicitPointer.test.ts b/packages/typegpu-gl/tests/implicitPointer.test.ts index f126f2fc60..d8d85cf807 100644 --- a/packages/typegpu-gl/tests/implicitPointer.test.ts +++ b/packages/typegpu-gl/tests/implicitPointer.test.ts @@ -239,4 +239,140 @@ describe('implicit pointers in GLSL', () => { }" `); }); + + it('evaluates a comptime index of an alias once', () => { + let calls = 0; + const nextIndex = tgpu.comptime(() => calls++); + + const fn = () => { + 'use gpu'; + const values = d.arrayOf(d.vec2i, 3)([d.vec2i(10, 11), d.vec2i(20, 21), d.vec2i(30, 31)]); + const value = values[nextIndex()]!; + return value.x; + }; + + expect(tgpu.resolve([fn], glOptions())).toMatchInlineSnapshot(` + "int fn_1() { + ivec2 values[3] = ivec2[3](ivec2(10, 11), ivec2(20, 21), ivec2(30, 31)); + return values[0].x; + }" + `); + expect(calls).toBe(1); + }); + + it('evaluates a comptime part of a hoisted index once', () => { + let calls = 0; + const nextIndex = tgpu.comptime(() => calls++); + const boids = tgpu.privateVar(d.arrayOf(Boid, 16)); + + function bar(index: number) { + 'use gpu'; + const boid = boids.$[nextIndex() + index]!; + boid.pos.x = 1; + } + + function main() { + 'use gpu'; + bar(1); + } + + expect(tgpu.resolve([main], glOptions())).toMatchInlineSnapshot(` + "struct Boid { + vec3 pos; + vec3 vel; + }; + + Boid boids[16]; + + void bar(int index) { + int idx = (0 + index); + boids[idx].pos.x = 1.0; + } + + void main() { + bar(1); + }" + `); + expect(calls).toBe(1); + }); + + it('restores the outer right-hand side after a nested const in a function it calls', () => { + let calls = 0; + const nextIndex = tgpu.comptime(() => calls++); + const boids = tgpu.privateVar(d.arrayOf(Boid, 16)); + + function pick() { + 'use gpu'; + const values = d.arrayOf(d.i32, 2)([3, 4]); + const value = values[nextIndex()]!; + return value; + } + + function main() { + 'use gpu'; + const boid = boids.$[pick() + nextIndex()]!; + boid.pos.x = 1; + } + + expect(tgpu.resolve([main], glOptions())).toMatchInlineSnapshot(` + "int pick() { + int values[2] = int[2](3, 4); + int value = values[0]; + return value; + } + + struct Boid { + vec3 pos; + vec3 vel; + }; + + Boid boids[16]; + + void main() { + int idx = (pick() + 1); + boids[idx].pos.x = 1.0; + }" + `); + expect(calls).toBe(2); + }); + + it('evaluates each comptime part of a nested hoisted access once', () => { + let rowCalls = 0; + let colCalls = 0; + const nextRow = tgpu.comptime(() => rowCalls++); + const nextCol = tgpu.comptime(() => colCalls++); + const grid = tgpu.privateVar(d.arrayOf(d.arrayOf(Boid, 4), 4)); + + function bar(row: number, col: number) { + 'use gpu'; + const boid = grid.$[nextRow() + row]![nextCol() + col]!; + boid.pos.x = 1; + } + + function main() { + 'use gpu'; + bar(1, 2); + } + + expect(tgpu.resolve([main], glOptions())).toMatchInlineSnapshot(` + "struct Boid { + vec3 pos; + vec3 vel; + }; + + Boid grid[4][4]; + + void bar(int row, int col) { + int idx = (0 + row); + int idx_1 = (0 + col); + grid[idx][idx_1].pos.x = 1.0; + } + + void main() { + bar(1, 2); + }" + `); + expect(rowCalls).toBe(1); + expect(colCalls).toBe(1); + }); });