diff --git a/language-extensions/index.d.ts b/language-extensions/index.d.ts index e5d3096a3..ad3692a97 100644 --- a/language-extensions/index.d.ts +++ b/language-extensions/index.d.ts @@ -34,6 +34,12 @@ declare type LuaMultiReturn = T & LuaExtension<"__luaMultiRetur declare const $range: ((start: number, limit: number, step?: number) => Iterable) & LuaExtension<"__luaRangeFunctionBrand">; +/** + * Transpiles to the global vararg (`...`) + * For more information see: https://typescripttolua.github.io/docs/advanced/language-extensions + */ +declare const $vararg: string[] & LuaExtension<"__luaVarargConstantBrand">; + /** * Represents a Lua-style iterator which is returned from a LuaIterable. * For simple iterators (with no state), this is just a function. diff --git a/src/lualib/ArrayReduce.ts b/src/lualib/ArrayReduce.ts index edb6825e6..c8ea2de50 100644 --- a/src/lualib/ArrayReduce.ts +++ b/src/lualib/ArrayReduce.ts @@ -3,7 +3,7 @@ function __TS__ArrayReduce( this: void, arr: T[], callbackFn: (accumulator: T, currentValue: T, index: number, array: T[]) => T, - ...initial: Vararg + ...initial: T[] ): T { const len = arr.length; diff --git a/src/lualib/ArrayReduceRight.ts b/src/lualib/ArrayReduceRight.ts index 8e5336f61..a0d5933d2 100644 --- a/src/lualib/ArrayReduceRight.ts +++ b/src/lualib/ArrayReduceRight.ts @@ -3,7 +3,7 @@ function __TS__ArrayReduceRight( this: void, arr: T[], callbackFn: (accumulator: T, currentValue: T, index: number, array: T[]) => T, - ...initial: Vararg + ...initial: T[] ): T { const len = arr.length; diff --git a/src/lualib/ArraySplice.ts b/src/lualib/ArraySplice.ts index 1b0bf9cff..0fae685e8 100644 --- a/src/lualib/ArraySplice.ts +++ b/src/lualib/ArraySplice.ts @@ -1,5 +1,5 @@ // https://www.ecma-international.org/ecma-262/9.0/index.html#sec-array.prototype.splice -function __TS__ArraySplice(this: void, list: T[], ...args: Vararg): T[] { +function __TS__ArraySplice(this: void, list: T[], ...args: unknown[]): T[] { const len = list.length; const actualArgumentCount = select("#", ...args); diff --git a/src/lualib/Generator.ts b/src/lualib/Generator.ts index a71740761..1573f5ec3 100644 --- a/src/lualib/Generator.ts +++ b/src/lualib/Generator.ts @@ -8,7 +8,7 @@ function __TS__GeneratorIterator(this: GeneratorIterator) { return this; } -function __TS__GeneratorNext(this: GeneratorIterator, ...args: Vararg) { +function __TS__GeneratorNext(this: GeneratorIterator, ...args: any[]) { const co = this.____coroutine; if (coroutine.status(co) === "dead") return { done: true }; @@ -19,7 +19,7 @@ function __TS__GeneratorNext(this: GeneratorIterator, ...args: Vararg) { } function __TS__Generator(this: void, fn: (this: void, ...args: any[]) => any) { - return function (this: void, ...args: Vararg): GeneratorIterator { + return function (this: void, ...args: any[]): GeneratorIterator { const argsLength = select("#", ...args); return { // Using explicit this there, since we don't pass arguments after the first nil and context is likely to be nil diff --git a/src/lualib/New.ts b/src/lualib/New.ts index 507ba120c..8b33ec837 100644 --- a/src/lualib/New.ts +++ b/src/lualib/New.ts @@ -1,4 +1,4 @@ -function __TS__New(this: void, target: LuaClass, ...args: Vararg): any { +function __TS__New(this: void, target: LuaClass, ...args: any[]): any { const instance: any = setmetatable({}, target.prototype); instance.____constructor(...args); return instance; diff --git a/src/lualib/declarations/tstl.d.ts b/src/lualib/declarations/tstl.d.ts index 4a718a31a..f68237190 100644 --- a/src/lualib/declarations/tstl.d.ts +++ b/src/lualib/declarations/tstl.d.ts @@ -1,8 +1,5 @@ /** @noSelfInFile */ -/** @vararg */ -type Vararg = T & { __luaVararg?: never }; - interface Metatable { _descriptors?: Record; __index?: any; diff --git a/src/transformation/utils/diagnostics.ts b/src/transformation/utils/diagnostics.ts index 4086ee0f2..422a93cac 100644 --- a/src/transformation/utils/diagnostics.ts +++ b/src/transformation/utils/diagnostics.ts @@ -64,6 +64,10 @@ export const invalidForRangeCall = createErrorDiagnosticFactory( export const invalidRangeUse = createErrorDiagnosticFactory("$range can only be used in a for...of loop."); +export const invalidVarargUse = createErrorDiagnosticFactory( + "$vararg can only be used in a spread element ('...$vararg') in global scope." +); + export const invalidRangeControlVariable = createErrorDiagnosticFactory( "For loop using $range must declare a single control variable." ); diff --git a/src/transformation/utils/language-extensions.ts b/src/transformation/utils/language-extensions.ts index cb17987ad..2991204c2 100644 --- a/src/transformation/utils/language-extensions.ts +++ b/src/transformation/utils/language-extensions.ts @@ -5,6 +5,7 @@ export enum ExtensionKind { MultiFunction = "MultiFunction", MultiType = "MultiType", RangeFunction = "RangeFunction", + VarargConstant = "VarargConstant", IterableType = "IterableType", AdditionOperatorType = "AdditionOperatorType", AdditionOperatorMethodType = "AdditionOperatorMethodType", @@ -49,15 +50,17 @@ export enum ExtensionKind { TableSetMethodType = "TableSetMethodType", } -const extensionKindToFunctionName: { [T in ExtensionKind]?: string } = { +const extensionKindToValueName: { [T in ExtensionKind]?: string } = { [ExtensionKind.MultiFunction]: "$multi", [ExtensionKind.RangeFunction]: "$range", + [ExtensionKind.VarargConstant]: "$vararg", }; const extensionKindToTypeBrand: { [T in ExtensionKind]: string } = { [ExtensionKind.MultiFunction]: "__luaMultiFunctionBrand", [ExtensionKind.MultiType]: "__luaMultiReturnBrand", [ExtensionKind.RangeFunction]: "__luaRangeFunctionBrand", + [ExtensionKind.VarargConstant]: "__luaVarargConstantBrand", [ExtensionKind.IterableType]: "__luaIterableBrand", [ExtensionKind.AdditionOperatorType]: "__luaAdditionBrand", [ExtensionKind.AdditionOperatorMethodType]: "__luaAdditionMethodBrand", @@ -107,13 +110,13 @@ export function isExtensionType(type: ts.Type, extensionKind: ExtensionKind): bo return typeBrand !== undefined && type.getProperty(typeBrand) !== undefined; } -export function isExtensionFunction( +export function isExtensionValue( context: TransformationContext, symbol: ts.Symbol, extensionKind: ExtensionKind ): boolean { return ( - symbol.getName() === extensionKindToFunctionName[extensionKind] && + symbol.getName() === extensionKindToValueName[extensionKind] && symbol.declarations.some(d => isExtensionType(context.checker.getTypeAtLocation(d), extensionKind)) ); } diff --git a/src/transformation/utils/lua-ast.ts b/src/transformation/utils/lua-ast.ts index 91e81bc9a..5f3a21e4a 100644 --- a/src/transformation/utils/lua-ast.ts +++ b/src/transformation/utils/lua-ast.ts @@ -62,6 +62,8 @@ export function getNumberLiteralValue(expression?: lua.Expression) { return undefined; } +// Prefer use of transformToImmediatelyInvokedFunctionExpression to maintain correct scope. If you use this directly, +// ensure you push/pop a function scope appropriately to avoid incorrect vararg optimization. export function createImmediatelyInvokedFunctionExpression( statements: lua.Statement[], result: lua.Expression | lua.Expression[], diff --git a/src/transformation/utils/scope.ts b/src/transformation/utils/scope.ts index 2c9f267dd..200b65860 100644 --- a/src/transformation/utils/scope.ts +++ b/src/transformation/utils/scope.ts @@ -3,7 +3,7 @@ import * as lua from "../../LuaAST"; import { assert, getOrUpdate, isNonNull } from "../../utils"; import { TransformationContext } from "../context"; import { getSymbolInfo } from "./symbols"; -import { getFirstDeclarationInFile } from "./typescript"; +import { findFirstNodeAbove, getFirstDeclarationInFile } from "./typescript"; export enum ScopeType { File = 1 << 0, @@ -24,6 +24,7 @@ interface FunctionDefinitionInfo { export interface Scope { type: ScopeType; id: number; + node?: ts.Node; referencedSymbols?: Map; variableDeclarations?: lua.VariableDeclarationStatement[]; functionDefinitions?: Map; @@ -91,6 +92,45 @@ export function popScope(context: TransformationContext): Scope { return scope; } +function isDeclaredInScope(symbol: ts.Symbol, scopeNode: ts.Node) { + return symbol?.declarations.some(d => findFirstNodeAbove(d, (n): n is ts.Node => n === scopeNode)); +} + +// Checks for references to local functions which haven't been defined yet, +// and thus will be hoisted above the current position. +export function hasReferencedUndefinedLocalFunction(context: TransformationContext, scope: Scope) { + if (!scope.referencedSymbols || !scope.node) { + return false; + } + for (const [symbolId, nodes] of scope.referencedSymbols) { + const type = context.checker.getTypeAtLocation(nodes[0]); + if ( + !scope.functionDefinitions?.has(symbolId) && + type.getCallSignatures().length > 0 && + isDeclaredInScope(type.symbol, scope.node) + ) { + return true; + } + } + return false; +} + +export function hasReferencedSymbol(context: TransformationContext, scope: Scope, symbol: ts.Symbol) { + if (!scope.referencedSymbols) { + return; + } + for (const nodes of scope.referencedSymbols.values()) { + if (nodes.some(node => context.checker.getSymbolAtLocation(node) === symbol)) { + return true; + } + } + return false; +} + +export function isFunctionScopeWithDefinition(scope: Scope): scope is Scope & { node: ts.SignatureDeclaration } { + return scope.node !== undefined && ts.isFunctionLike(scope.node); +} + export function performHoisting(context: TransformationContext, statements: lua.Statement[]): lua.Statement[] { const scope = peekScope(context); let result = statements; diff --git a/src/transformation/utils/symbols.ts b/src/transformation/utils/symbols.ts index c99b41a25..f800eea86 100644 --- a/src/transformation/utils/symbols.ts +++ b/src/transformation/utils/symbols.ts @@ -2,6 +2,7 @@ import * as ts from "typescript"; import * as lua from "../../LuaAST"; import { getOrUpdate } from "../../utils"; import { TransformationContext } from "../context"; +import { isOptimizedVarArgSpread } from "../visitors/spread"; import { markSymbolAsReferencedInCurrentScopes } from "./scope"; const symbolIdCounters = new WeakMap(); @@ -44,7 +45,11 @@ export function trackSymbolReference( symbolInfo.set(symbolId, { symbol, firstSeenAtPos: identifier.pos }); } - markSymbolAsReferencedInCurrentScopes(context, symbolId, identifier); + // If isOptimizedVarArgSpread returns true, the identifier will not appear in the resulting Lua. + // Only the optimized ellipses (...) will be used. + if (!isOptimizedVarArgSpread(context, symbol, identifier)) { + markSymbolAsReferencedInCurrentScopes(context, symbolId, identifier); + } return symbolId; } diff --git a/src/transformation/utils/transform.ts b/src/transformation/utils/transform.ts new file mode 100644 index 000000000..215a018d5 --- /dev/null +++ b/src/transformation/utils/transform.ts @@ -0,0 +1,22 @@ +import * as ts from "typescript"; +import * as lua from "../../LuaAST"; +import { castArray } from "../../utils"; +import { TransformationContext } from "../context"; +import { createImmediatelyInvokedFunctionExpression } from "./lua-ast"; +import { ScopeType, pushScope, popScope } from "./scope"; + +export interface ImmediatelyInvokedFunctionParameters { + statements: lua.Statement | lua.Statement[]; + result: lua.Expression | lua.Expression[]; +} + +export function transformToImmediatelyInvokedFunctionExpression( + context: TransformationContext, + transformFunction: () => ImmediatelyInvokedFunctionParameters, + tsOriginal?: ts.Node +): lua.CallExpression { + pushScope(context, ScopeType.Function); + const { statements, result } = transformFunction(); + popScope(context); + return createImmediatelyInvokedFunctionExpression(castArray(statements), result, tsOriginal); +} diff --git a/src/transformation/visitors/binary-expression/assignments.ts b/src/transformation/visitors/binary-expression/assignments.ts index df99a2a7d..0b1b7fe70 100644 --- a/src/transformation/visitors/binary-expression/assignments.ts +++ b/src/transformation/visitors/binary-expression/assignments.ts @@ -5,13 +5,18 @@ import { TransformationContext } from "../../context"; import { isTupleReturnCall } from "../../utils/annotations"; import { validateAssignment } from "../../utils/assignment-validation"; import { createExportedIdentifier, getDependenciesOfSymbol, isSymbolExported } from "../../utils/export"; -import { createImmediatelyInvokedFunctionExpression, createUnpackCall, wrapInTable } from "../../utils/lua-ast"; +import { createUnpackCall, wrapInTable } from "../../utils/lua-ast"; import { LuaLibFeature, transformLuaLibFunction } from "../../utils/lualib"; import { isArrayType, isDestructuringAssignment } from "../../utils/typescript"; import { transformElementAccessArgument } from "../access"; import { transformLuaTablePropertyAccessInAssignment } from "../lua-table"; import { isArrayLength, transformDestructuringAssignment } from "./destructuring-assignments"; import { isMultiReturnCall } from "../language-extensions/multi"; +import { popScope, pushScope, ScopeType } from "../../utils/scope"; +import { + ImmediatelyInvokedFunctionParameters, + transformToImmediatelyInvokedFunctionExpression, +} from "../../utils/transform"; export function transformAssignmentLeftHandSideExpression( context: TransformationContext, @@ -68,6 +73,25 @@ export function transformAssignment( ]; } +function transformDestructuredAssignmentExpression( + context: TransformationContext, + expression: ts.DestructuringAssignment +): ImmediatelyInvokedFunctionParameters { + const rootIdentifier = lua.createAnonymousIdentifier(expression.left); + + let right = context.transformExpression(expression.right); + if (isTupleReturnCall(context, expression.right) || isMultiReturnCall(context, expression.right)) { + right = wrapInTable(right); + } + + const statements = [ + lua.createVariableDeclarationStatement(rootIdentifier, right), + ...transformDestructuringAssignment(context, expression, rootIdentifier), + ]; + + return { statements, result: rootIdentifier }; +} + export function transformAssignmentExpression( context: TransformationContext, expression: ts.AssignmentExpression @@ -89,19 +113,11 @@ export function transformAssignmentExpression( } if (isDestructuringAssignment(expression)) { - const rootIdentifier = lua.createAnonymousIdentifier(expression.left); - - let right = context.transformExpression(expression.right); - if (isTupleReturnCall(context, expression.right) || isMultiReturnCall(context, expression.right)) { - right = wrapInTable(right); - } - - const statements = [ - lua.createVariableDeclarationStatement(rootIdentifier, right), - ...transformDestructuringAssignment(context, expression, rootIdentifier), - ]; - - return createImmediatelyInvokedFunctionExpression(statements, rootIdentifier, expression); + return transformToImmediatelyInvokedFunctionExpression( + context, + () => transformDestructuredAssignmentExpression(context, expression), + expression + ); } if (ts.isPropertyAccessExpression(expression.left) || ts.isElementAccessExpression(expression.left)) { @@ -120,6 +136,7 @@ export function transformAssignmentExpression( indexParameter, valueParameter, ]); + pushScope(context, ScopeType.Function); const objExpression = context.transformExpression(expression.left.expression); let indexExpression: lua.Expression; if (ts.isPropertyAccessExpression(expression.left)) { @@ -134,15 +151,19 @@ export function transformAssignmentExpression( } const args = [objExpression, indexExpression, context.transformExpression(expression.right)]; + popScope(context); return lua.createCallExpression(iife, args, expression); } else { - // Simple assignment - // (function() ${left} = ${right}; return ${left} end)() - const left = context.transformExpression(expression.left); - const right = context.transformExpression(expression.right); - return createImmediatelyInvokedFunctionExpression( - transformAssignment(context, expression.left, right), - left, + return transformToImmediatelyInvokedFunctionExpression( + context, + () => { + // Simple assignment + // (function() ${left} = ${right}; return ${left} end)() + const left = context.transformExpression(expression.left); + const right = context.transformExpression(expression.right); + const statements = transformAssignment(context, expression.left, right); + return { statements, result: left }; + }, expression ); } diff --git a/src/transformation/visitors/binary-expression/compound.ts b/src/transformation/visitors/binary-expression/compound.ts index 3b92ade58..7a11454c8 100644 --- a/src/transformation/visitors/binary-expression/compound.ts +++ b/src/transformation/visitors/binary-expression/compound.ts @@ -2,7 +2,10 @@ import * as ts from "typescript"; import * as lua from "../../../LuaAST"; import { cast, assertNever } from "../../../utils"; import { TransformationContext } from "../../context"; -import { createImmediatelyInvokedFunctionExpression } from "../../utils/lua-ast"; +import { + ImmediatelyInvokedFunctionParameters, + transformToImmediatelyInvokedFunctionExpression, +} from "../../utils/transform"; import { isArrayType, isExpressionWithEvaluationEffect } from "../../utils/typescript"; import { transformBinaryOperation } from "../binary-expression"; import { transformAssignment } from "./assignments"; @@ -75,15 +78,14 @@ export const isCompoundAssignmentToken = (token: ts.BinaryOperator): token is ts export const unwrapCompoundAssignmentToken = (token: ts.CompoundAssignmentOperator): CompoundAssignmentToken => compoundToAssignmentTokens[token]; -export function transformCompoundAssignmentExpression( +export function transformCompoundAssignment( context: TransformationContext, expression: ts.Expression, - // TODO: Change type to ts.LeftHandSideExpression? lhs: ts.Expression, rhs: ts.Expression, operator: CompoundAssignmentToken, isPostfix: boolean -): lua.CallExpression { +): ImmediatelyInvokedFunctionParameters { const left = cast(context.transformExpression(lhs), lua.isAssignmentLeftHandSideExpression); const right = context.transformExpression(rhs); @@ -116,11 +118,7 @@ export function transformCompoundAssignmentExpression( assignStatement = lua.createAssignmentStatement(accessExpression, tmp); } // return ____tmp - return createImmediatelyInvokedFunctionExpression( - [objAndIndexDeclaration, tmpDeclaration, assignStatement], - tmp, - expression - ); + return { statements: [objAndIndexDeclaration, tmpDeclaration, assignStatement], result: tmp }; } else if (isPostfix) { // Postfix expressions need to cache original value in temp // local ____tmp = ${left}; @@ -130,11 +128,7 @@ export function transformCompoundAssignmentExpression( const tmpDeclaration = lua.createVariableDeclarationStatement(tmpIdentifier, left); const operatorExpression = transformBinaryOperation(context, tmpIdentifier, right, operator, expression); const assignStatements = transformAssignment(context, lhs, operatorExpression); - return createImmediatelyInvokedFunctionExpression( - [tmpDeclaration, ...assignStatements], - tmpIdentifier, - expression - ); + return { statements: [tmpDeclaration, ...assignStatements], result: tmpIdentifier }; } else if (ts.isPropertyAccessExpression(lhs) || ts.isElementAccessExpression(lhs)) { // Simple property/element access expressions need to cache in temp to avoid double-evaluation // local ____tmp = ${left} ${replacementOperator} ${right}; @@ -146,27 +140,39 @@ export function transformCompoundAssignmentExpression( const assignStatements = transformAssignment(context, lhs, tmpIdentifier); if (isSetterSkippingCompoundAssignmentOperator(operator)) { - return createImmediatelyInvokedFunctionExpression( - [tmpDeclaration, ...transformSetterSkippingCompoundAssignment(context, tmpIdentifier, operator, rhs)], - tmpIdentifier, - expression - ); + const statements = [ + tmpDeclaration, + ...transformSetterSkippingCompoundAssignment(context, tmpIdentifier, operator, rhs), + ]; + return { statements, result: tmpIdentifier }; } - return createImmediatelyInvokedFunctionExpression( - [tmpDeclaration, ...assignStatements], - tmpIdentifier, - expression - ); + return { statements: [tmpDeclaration, ...assignStatements], result: tmpIdentifier }; } else { // Simple expressions // ${left} = ${right}; return ${right} const operatorExpression = transformBinaryOperation(context, left, right, operator, expression); - const assignStatements = transformAssignment(context, lhs, operatorExpression); - return createImmediatelyInvokedFunctionExpression(assignStatements, left, expression); + const statements = transformAssignment(context, lhs, operatorExpression); + return { statements, result: left }; } } +export function transformCompoundAssignmentExpression( + context: TransformationContext, + expression: ts.Expression, + // TODO: Change type to ts.LeftHandSideExpression? + lhs: ts.Expression, + rhs: ts.Expression, + operator: CompoundAssignmentToken, + isPostfix: boolean +): lua.CallExpression { + return transformToImmediatelyInvokedFunctionExpression( + context, + () => transformCompoundAssignment(context, expression, lhs, rhs, operator, isPostfix), + expression + ); +} + export function transformCompoundAssignmentStatement( context: TransformationContext, node: ts.Node, diff --git a/src/transformation/visitors/binary-expression/index.ts b/src/transformation/visitors/binary-expression/index.ts index d9e4746e0..26e90ad67 100644 --- a/src/transformation/visitors/binary-expression/index.ts +++ b/src/transformation/visitors/binary-expression/index.ts @@ -3,7 +3,7 @@ import * as lua from "../../../LuaAST"; import { FunctionVisitor, TransformationContext } from "../../context"; import { AnnotationKind, getTypeAnnotations } from "../../utils/annotations"; import { luaTableInvalidInstanceOf } from "../../utils/diagnostics"; -import { createImmediatelyInvokedFunctionExpression, wrapInToStringForConcat } from "../../utils/lua-ast"; +import { wrapInToStringForConcat } from "../../utils/lua-ast"; import { LuaLibFeature, transformLuaLibFunction } from "../../utils/lualib"; import { isStandardLibraryType, isStringType, typeCanSatisfy } from "../../utils/typescript"; import { transformTypeOfBinaryExpression } from "../typeof"; @@ -16,6 +16,7 @@ import { unwrapCompoundAssignmentToken, } from "./compound"; import { assert } from "../../../utils"; +import { transformToImmediatelyInvokedFunctionExpression } from "../../utils/transform"; type SimpleOperator = | ts.AdditiveOperatorOrHigher @@ -121,9 +122,12 @@ export const transformBinaryExpression: FunctionVisitor = ( } case ts.SyntaxKind.CommaToken: { - return createImmediatelyInvokedFunctionExpression( - context.transformStatements(ts.createExpressionStatement(node.left)), - context.transformExpression(node.right), + return transformToImmediatelyInvokedFunctionExpression( + context, + () => ({ + statements: context.transformStatements(ts.createExpressionStatement(node.left)), + result: context.transformExpression(node.right), + }), node ); } diff --git a/src/transformation/visitors/call.ts b/src/transformation/visitors/call.ts index 7ea22eb34..75f73db88 100644 --- a/src/transformation/visitors/call.ts +++ b/src/transformation/visitors/call.ts @@ -2,16 +2,16 @@ import * as ts from "typescript"; import * as lua from "../../LuaAST"; import { transformBuiltinCallExpression } from "../builtins"; import { FunctionVisitor, TransformationContext } from "../context"; -import { isInTupleReturnFunction, isTupleReturnCall, isVarargType } from "../utils/annotations"; +import { isInTupleReturnFunction, isTupleReturnCall } from "../utils/annotations"; import { validateAssignment } from "../utils/assignment-validation"; import { ContextType, getDeclarationContextType } from "../utils/function-context"; -import { createImmediatelyInvokedFunctionExpression, createUnpackCall, wrapInTable } from "../utils/lua-ast"; +import { createUnpackCall, wrapInTable } from "../utils/lua-ast"; import { LuaLibFeature, transformLuaLibFunction } from "../utils/lualib"; import { isValidLuaIdentifier } from "../utils/safe-names"; -import { isArrayType, isExpressionWithEvaluationEffect, isInDestructingAssignment } from "../utils/typescript"; +import { isExpressionWithEvaluationEffect, isInDestructingAssignment } from "../utils/typescript"; import { transformElementAccessArgument } from "./access"; import { transformLuaTableCallExpression } from "./lua-table"; -import { shouldMultiReturnCallBeWrapped, returnsMultiType } from "./language-extensions/multi"; +import { shouldMultiReturnCallBeWrapped } from "./language-extensions/multi"; import { isOperatorMapping, transformOperatorMappingExpression } from "./language-extensions/operators"; import { isTableGetCall, @@ -20,6 +20,10 @@ import { transformTableSetExpression, } from "./language-extensions/table"; import { invalidTableSetExpression } from "../utils/diagnostics"; +import { + ImmediatelyInvokedFunctionParameters, + transformToImmediatelyInvokedFunctionExpression, +} from "../utils/transform"; export type PropertyCallExpression = ts.CallExpression & { expression: ts.PropertyAccessExpression }; @@ -113,10 +117,32 @@ export function transformArguments( return parameters; } +function transformElementAccessCall( + context: TransformationContext, + left: ts.PropertyAccessExpression | ts.ElementAccessExpression, + args: ts.Expression[] | ts.NodeArray, + signature?: ts.Signature +): ImmediatelyInvokedFunctionParameters { + const transformedArguments = transformArguments(context, args, signature, ts.factory.createIdentifier("____self")); + + // Cache left-side if it has effects + // (function() local ____self = context; return ____self[argument](parameters); end)() + const argument = ts.isElementAccessExpression(left) + ? transformElementAccessArgument(context, left) + : lua.createStringLiteral(left.name.text); + const selfIdentifier = lua.createIdentifier("____self"); + const callContext = context.transformExpression(left.expression); + const selfAssignment = lua.createVariableDeclarationStatement(selfIdentifier, callContext); + const index = lua.createTableIndexExpression(selfIdentifier, argument); + const callExpression = lua.createCallExpression(index, transformedArguments); + return { statements: selfAssignment, result: callExpression }; +} + export function transformContextualCallExpression( context: TransformationContext, node: ts.CallExpression | ts.TaggedTemplateExpression, - transformedArguments: lua.Expression[] + args: ts.Expression[] | ts.NodeArray, + signature?: ts.Signature ): lua.Expression { const left = ts.isCallExpression(node) ? node.expression : node.tag; if (ts.isPropertyAccessExpression(left) && ts.isIdentifier(left.name) && isValidLuaIdentifier(left.name.text)) { @@ -126,32 +152,25 @@ export function transformContextualCallExpression( return lua.createMethodCallExpression( table, lua.createIdentifier(left.name.text, left.name), - transformedArguments, + transformArguments(context, args, signature), node ); } else if (ts.isElementAccessExpression(left) || ts.isPropertyAccessExpression(left)) { - const callContext = context.transformExpression(left.expression); if (isExpressionWithEvaluationEffect(left.expression)) { - // Inject context parameter - transformedArguments.unshift(lua.createIdentifier("____self")); - - // Cache left-side if it has effects - // (function() local ____self = context; return ____self[argument](parameters); end)() - const argument = ts.isElementAccessExpression(left) - ? transformElementAccessArgument(context, left) - : lua.createStringLiteral(left.name.text); - const selfIdentifier = lua.createIdentifier("____self"); - const selfAssignment = lua.createVariableDeclarationStatement(selfIdentifier, callContext); - const index = lua.createTableIndexExpression(selfIdentifier, argument); - const callExpression = lua.createCallExpression(index, transformedArguments); - return createImmediatelyInvokedFunctionExpression([selfAssignment], callExpression, node); + return transformToImmediatelyInvokedFunctionExpression( + context, + () => transformElementAccessCall(context, left, args, signature), + node + ); } else { + const callContext = context.transformExpression(left.expression); const expression = context.transformExpression(left); + const transformedArguments = transformArguments(context, args, signature); return lua.createCallExpression(expression, [callContext, ...transformedArguments]); } } else if (ts.isIdentifier(left)) { - const callContext = context.isStrict ? lua.createNilLiteral() : lua.createIdentifier("_G"); - transformedArguments.unshift(callContext); + const callContext = context.isStrict ? ts.factory.createNull() : ts.factory.createIdentifier("_G"); + const transformedArguments = transformArguments(context, args, signature, callContext); const expression = context.transformExpression(left); return lua.createCallExpression(expression, transformedArguments, node); } else { @@ -168,17 +187,17 @@ function transformPropertyCall(context: TransformationContext, node: PropertyCal return lua.createCallExpression(context.transformExpression(node.expression), parameters); } - const parameters = transformArguments(context, node.arguments, signature); const signatureDeclaration = signature?.getDeclaration(); if (!signatureDeclaration || getDeclarationContextType(context, signatureDeclaration) !== ContextType.Void) { // table:name() - return transformContextualCallExpression(context, node, parameters); + return transformContextualCallExpression(context, node, node.arguments, signature); } else { const table = context.transformExpression(node.expression.expression); // table.name() const name = node.expression.name.text; const callPath = lua.createTableIndexExpression(table, lua.createStringLiteral(name), node.expression); + const parameters = transformArguments(context, node.arguments, signature); return lua.createCallExpression(callPath, parameters, node); } } @@ -186,13 +205,13 @@ function transformPropertyCall(context: TransformationContext, node: PropertyCal function transformElementCall(context: TransformationContext, node: ts.CallExpression): lua.Expression { const signature = context.checker.getResolvedSignature(node); const signatureDeclaration = signature?.getDeclaration(); - const parameters = transformArguments(context, node.arguments, signature); if (!signatureDeclaration || getDeclarationContextType(context, signatureDeclaration) !== ContextType.Void) { // A contextual parameter must be given to this call expression - return transformContextualCallExpression(context, node, parameters); + return transformContextualCallExpression(context, node, node.arguments, signature); } else { // No context const expression = context.transformExpression(node.expression); + const parameters = transformArguments(context, node.arguments, signature); return lua.createCallExpression(expression, parameters); } } @@ -227,9 +246,10 @@ export const transformCallExpression: FunctionVisitor = (node if (isTableSetCall(context, node)) { context.diagnostics.push(invalidTableSetExpression(node)); - return createImmediatelyInvokedFunctionExpression( - [transformTableSetExpression(context, node)], - lua.createNilLiteral() + return transformToImmediatelyInvokedFunctionExpression( + context, + () => ({ statements: transformTableSetExpression(context, node), result: lua.createNilLiteral() }), + node ); } @@ -272,21 +292,3 @@ export const transformCallExpression: FunctionVisitor = (node const callExpression = lua.createCallExpression(callPath, parameters, node); return wrapResult ? wrapInTable(callExpression) : callExpression; }; - -// TODO: Currently it's also used as an array member -export const transformSpreadElement: FunctionVisitor = (node, context) => { - const innerExpression = context.transformExpression(node.expression); - if (isTupleReturnCall(context, node.expression)) return innerExpression; - if (ts.isCallExpression(node.expression) && returnsMultiType(context, node.expression)) return innerExpression; - - if (ts.isIdentifier(node.expression) && isVarargType(context, node.expression)) { - return lua.createDotsLiteral(node); - } - - const type = context.checker.getTypeAtLocation(node.expression); - if (isArrayType(context, type)) { - return createUnpackCall(context, innerExpression, node); - } - - return transformLuaLibFunction(context, LuaLibFeature.Spread, node, innerExpression); -}; diff --git a/src/transformation/visitors/class/index.ts b/src/transformation/visitors/class/index.ts index 8b40eadcd..9963c09f4 100644 --- a/src/transformation/visitors/class/index.ts +++ b/src/transformation/visitors/class/index.ts @@ -10,13 +10,9 @@ import { hasDefaultExportModifier, isSymbolExported, } from "../../utils/export"; -import { - createImmediatelyInvokedFunctionExpression, - createSelfIdentifier, - unwrapVisitorResult, -} from "../../utils/lua-ast"; +import { createSelfIdentifier, unwrapVisitorResult } from "../../utils/lua-ast"; import { createSafeName, isUnsafeName } from "../../utils/safe-names"; -import { popScope, pushScope, ScopeType } from "../../utils/scope"; +import { transformToImmediatelyInvokedFunctionExpression } from "../../utils/transform"; import { isAmbientNode } from "../../utils/typescript"; import { transformIdentifier } from "../identifier"; import { createDecoratingExpression, transformDecoratorExpression } from "./decorators"; @@ -50,11 +46,14 @@ export function transformClassAsExpression( expression: ts.ClassLikeDeclaration, context: TransformationContext ): lua.Expression { - pushScope(context, ScopeType.Function); - const { statements, name } = transformClassLikeDeclaration(expression, context); - popScope(context); - - return createImmediatelyInvokedFunctionExpression(unwrapVisitorResult(statements), name, expression); + return transformToImmediatelyInvokedFunctionExpression( + context, + () => { + const { statements, name } = transformClassLikeDeclaration(expression, context); + return { statements: unwrapVisitorResult(statements), result: name }; + }, + expression + ); } const classSuperInfos = new WeakMap(); diff --git a/src/transformation/visitors/function.ts b/src/transformation/visitors/function.ts index d3fb88a5a..cb1d29bfd 100644 --- a/src/transformation/visitors/function.ts +++ b/src/transformation/visitors/function.ts @@ -2,7 +2,6 @@ import * as ts from "typescript"; import * as lua from "../../LuaAST"; import { assert } from "../../utils"; import { FunctionVisitor, TransformationContext } from "../context"; -import { isVarargType } from "../utils/annotations"; import { createDefaultExportStringLiteral, hasDefaultExportModifier } from "../utils/export"; import { ContextType, getFunctionContextType } from "../utils/function-context"; import { @@ -38,7 +37,7 @@ function transformParameterDefaultValueDeclaration( return lua.createIfStatement(nilCondition, ifBlock, undefined, tsOriginal); } -function isRestParameterReferenced(context: TransformationContext, identifier: lua.Identifier, scope: Scope): boolean { +function isRestParameterReferenced(identifier: lua.Identifier, scope: Scope): boolean { if (!identifier.symbolId) { return true; } @@ -46,11 +45,7 @@ function isRestParameterReferenced(context: TransformationContext, identifier: l return false; } const references = scope.referencedSymbols.get(identifier.symbolId); - if (!references) { - return false; - } - // Ignore references to @vararg types in spread elements - return references.some(r => !r.parent || !ts.isSpreadElement(r.parent) || !isVarargType(context, r)); + return references !== undefined && references.length > 0; } export function transformFunctionBodyContent(context: TransformationContext, body: ts.ConciseBody): lua.Statement[] { @@ -99,7 +94,7 @@ export function transformFunctionBodyHeader( } // Push spread operator here - if (spreadIdentifier && isRestParameterReferenced(context, spreadIdentifier, bodyScope)) { + if (spreadIdentifier && isRestParameterReferenced(spreadIdentifier, bodyScope)) { const spreadTable = wrapInTable(lua.createDotsLiteral()); headerStatements.push(lua.createVariableDeclarationStatement(spreadIdentifier, spreadTable)); } @@ -114,9 +109,11 @@ export function transformFunctionBody( context: TransformationContext, parameters: ts.NodeArray, body: ts.ConciseBody, - spreadIdentifier?: lua.Identifier + spreadIdentifier?: lua.Identifier, + node?: ts.FunctionLikeDeclaration ): [lua.Statement[], Scope] { const scope = pushScope(context, ScopeType.Function); + scope.node = node; const bodyStatements = transformFunctionBodyContent(context, body); const headerStatements = transformFunctionBodyHeader(context, scope, parameters, spreadIdentifier); popScope(context); @@ -195,7 +192,8 @@ export function transformFunctionToExpression( context, node.parameters, node.body, - spreadIdentifier + spreadIdentifier, + node ); const functionExpression = lua.createFunctionExpression( lua.createBlock(transformedBody), @@ -236,6 +234,10 @@ export function transformFunctionLikeDeclaration( // Only wrap if the name is actually referenced inside the function if (isReferenced) { const nameIdentifier = transformIdentifier(context, node.name); + // We cannot use transformToImmediatelyInvokedFunctionExpression() here because we need to transpile + // the function first to determine if it's self-referencing. Fortunately, this does not cause issues + // with var-arg optimization because the IIFE is just wrapping another function which will already push + // another scope. return createImmediatelyInvokedFunctionExpression( [lua.createVariableDeclarationStatement(nameIdentifier, functionExpression)], lua.cloneIdentifier(nameIdentifier) diff --git a/src/transformation/visitors/identifier.ts b/src/transformation/visitors/identifier.ts index 5905f1485..191319160 100644 --- a/src/transformation/visitors/identifier.ts +++ b/src/transformation/visitors/identifier.ts @@ -8,6 +8,7 @@ import { invalidMultiFunctionUse, invalidOperatorMappingUse, invalidRangeUse, + invalidVarargUse, invalidTableExtensionUse, } from "../utils/diagnostics"; import { createExportedIdentifier, getSymbolExportScope } from "../utils/export"; @@ -18,6 +19,7 @@ import { isMultiFunctionNode } from "./language-extensions/multi"; import { isOperatorMapping } from "./language-extensions/operators"; import { isRangeFunctionNode } from "./language-extensions/range"; import { isTableExtensionIdentifier } from "./language-extensions/table"; +import { isVarargConstantNode } from "./language-extensions/vararg"; export function transformIdentifier(context: TransformationContext, identifier: ts.Identifier): lua.Identifier { if (isMultiFunctionNode(context, identifier)) { @@ -38,6 +40,11 @@ export function transformIdentifier(context: TransformationContext, identifier: return lua.createAnonymousIdentifier(identifier); } + if (isVarargConstantNode(context, identifier)) { + context.diagnostics.push(invalidVarargUse(identifier)); + return lua.createAnonymousIdentifier(identifier); + } + if (isForRangeType(context, identifier)) { const callExpression = findFirstNodeAbove(identifier, ts.isCallExpression); if (!callExpression || !callExpression.parent || !ts.isForOfStatement(callExpression.parent)) { diff --git a/src/transformation/visitors/index.ts b/src/transformation/visitors/index.ts index 48469a84b..b9983ae31 100644 --- a/src/transformation/visitors/index.ts +++ b/src/transformation/visitors/index.ts @@ -4,7 +4,8 @@ import { transformElementAccessExpression, transformPropertyAccessExpression, tr import { transformBinaryExpression } from "./binary-expression"; import { transformBlock } from "./block"; import { transformBreakStatement, transformContinueStatement } from "./break-continue"; -import { transformCallExpression, transformSpreadElement } from "./call"; +import { transformCallExpression } from "./call"; +import { transformSpreadElement } from "./spread"; import { transformClassAsExpression, transformClassDeclaration, diff --git a/src/transformation/visitors/language-extensions/multi.ts b/src/transformation/visitors/language-extensions/multi.ts index b292d9aa1..85b3c6a29 100644 --- a/src/transformation/visitors/language-extensions/multi.ts +++ b/src/transformation/visitors/language-extensions/multi.ts @@ -25,7 +25,7 @@ export function isMultiReturnCall(context: TransformationContext, expression: ts export function isMultiFunctionNode(context: TransformationContext, node: ts.Node): boolean { const symbol = context.checker.getSymbolAtLocation(node); - return symbol ? extensions.isExtensionFunction(context, symbol, extensions.ExtensionKind.MultiFunction) : false; + return symbol ? extensions.isExtensionValue(context, symbol, extensions.ExtensionKind.MultiFunction) : false; } export function isInMultiReturnFunction(context: TransformationContext, node: ts.Node) { @@ -98,7 +98,7 @@ export function findMultiAssignmentViolations( if (!ts.isShorthandPropertyAssignment(element)) continue; const valueSymbol = context.checker.getShorthandAssignmentValueSymbol(element); if (valueSymbol) { - if (extensions.isExtensionFunction(context, valueSymbol, extensions.ExtensionKind.MultiFunction)) { + if (extensions.isExtensionValue(context, valueSymbol, extensions.ExtensionKind.MultiFunction)) { context.diagnostics.push(invalidMultiFunctionUse(element)); result.push(element); } diff --git a/src/transformation/visitors/language-extensions/range.ts b/src/transformation/visitors/language-extensions/range.ts index 00a3c249e..43f297a9a 100644 --- a/src/transformation/visitors/language-extensions/range.ts +++ b/src/transformation/visitors/language-extensions/range.ts @@ -14,7 +14,7 @@ export function isRangeFunction(context: TransformationContext, expression: ts.C export function isRangeFunctionNode(context: TransformationContext, node: ts.Node): boolean { const symbol = context.checker.getSymbolAtLocation(node); - return symbol ? extensions.isExtensionFunction(context, symbol, extensions.ExtensionKind.RangeFunction) : false; + return symbol ? extensions.isExtensionValue(context, symbol, extensions.ExtensionKind.RangeFunction) : false; } function getControlVariable(context: TransformationContext, statement: ts.ForOfStatement) { diff --git a/src/transformation/visitors/language-extensions/vararg.ts b/src/transformation/visitors/language-extensions/vararg.ts new file mode 100644 index 000000000..99eea82a9 --- /dev/null +++ b/src/transformation/visitors/language-extensions/vararg.ts @@ -0,0 +1,16 @@ +import * as ts from "typescript"; +import { TransformationContext } from "../../context"; +import * as extensions from "../../utils/language-extensions"; +import { Scope, ScopeType } from "../../utils/scope"; + +export function isGlobalVarargConstant(context: TransformationContext, symbol: ts.Symbol, scope: Scope) { + return ( + scope.type === ScopeType.File && + extensions.isExtensionValue(context, symbol, extensions.ExtensionKind.VarargConstant) + ); +} + +export function isVarargConstantNode(context: TransformationContext, node: ts.Node): boolean { + const symbol = context.checker.getSymbolAtLocation(node); + return symbol ? extensions.isExtensionValue(context, symbol, extensions.ExtensionKind.VarargConstant) : false; +} diff --git a/src/transformation/visitors/spread.ts b/src/transformation/visitors/spread.ts new file mode 100644 index 000000000..ca5e2fb85 --- /dev/null +++ b/src/transformation/visitors/spread.ts @@ -0,0 +1,82 @@ +import * as ts from "typescript"; +import * as lua from "../../LuaAST"; +import { FunctionVisitor, TransformationContext } from "../context"; +import { AnnotationKind, isTupleReturnCall, isVarargType } from "../utils/annotations"; +import { createUnpackCall } from "../utils/lua-ast"; +import { LuaLibFeature, transformLuaLibFunction } from "../utils/lualib"; +import { + findScope, + hasReferencedSymbol, + hasReferencedUndefinedLocalFunction, + isFunctionScopeWithDefinition, + ScopeType, +} from "../utils/scope"; +import { isArrayType } from "../utils/typescript"; +import { returnsMultiType } from "./language-extensions/multi"; +import { annotationDeprecated } from "../utils/diagnostics"; +import { isGlobalVarargConstant } from "./language-extensions/vararg"; + +export function isOptimizedVarArgSpread(context: TransformationContext, symbol: ts.Symbol, identifier: ts.Identifier) { + if (!ts.isSpreadElement(identifier.parent)) { + return false; + } + + // Walk up, stopping at any scope types which could stop optimization + const scope = findScope(context, ScopeType.Function | ScopeType.Try | ScopeType.Catch | ScopeType.File); + if (!scope) { + return; + } + + // $vararg global constant + if (isGlobalVarargConstant(context, symbol, scope)) { + return true; + } + + // Scope must be a function scope associated with a real ts function + if (!isFunctionScopeWithDefinition(scope)) { + return false; + } + + // Identifier must be a vararg in the local function scope's parameters + const isSpreadParameter = (p: ts.ParameterDeclaration) => + p.dotDotDotToken && ts.isIdentifier(p.name) && context.checker.getSymbolAtLocation(p.name) === symbol; + if (!scope.node.parameters.some(isSpreadParameter)) { + return false; + } + + // De-optimize if already referenced outside of a spread, as the array may have been modified + if (hasReferencedSymbol(context, scope, symbol)) { + return false; + } + + // De-optimize if a function is being hoisted from below to above, as it may have modified the array + if (hasReferencedUndefinedLocalFunction(context, scope)) { + return false; + } + return true; +} + +// TODO: Currently it's also used as an array member +export const transformSpreadElement: FunctionVisitor = (node, context) => { + if (ts.isIdentifier(node.expression)) { + if (isVarargType(context, node.expression)) { + context.diagnostics.push(annotationDeprecated(node, AnnotationKind.Vararg)); + return lua.createDotsLiteral(node); + } + const symbol = context.checker.getSymbolAtLocation(node.expression); + if (symbol && isOptimizedVarArgSpread(context, symbol, node.expression)) { + return lua.createDotsLiteral(node); + } + } + + const innerExpression = context.transformExpression(node.expression); + if (isTupleReturnCall(context, node.expression)) return innerExpression; + if (ts.isCallExpression(node.expression) && returnsMultiType(context, node.expression)) return innerExpression; + + const type = context.checker.getTypeAtLocation(node.expression); + if (isArrayType(context, type)) { + return createUnpackCall(context, innerExpression, node); + } + + return transformLuaLibFunction(context, LuaLibFeature.Spread, node, innerExpression); +}; diff --git a/src/transformation/visitors/template.ts b/src/transformation/visitors/template.ts index 1b714d765..9f9a19b33 100644 --- a/src/transformation/visitors/template.ts +++ b/src/transformation/visitors/template.ts @@ -65,29 +65,35 @@ export const transformTaggedTemplateExpression: FunctionVisitor lua.createTableFieldExpression(lua.createStringLiteral(text))) + const rawStringsArray = ts.factory.createArrayLiteralExpression( + rawStrings.map(text => ts.factory.createStringLiteral(text)) ); - const stringTableLiteral = lua.createTableExpression([ - ...strings.map(partialString => lua.createTableFieldExpression(lua.createStringLiteral(partialString))), - lua.createTableFieldExpression(rawStringsTable, lua.createStringLiteral("raw")), + const stringObject = ts.factory.createObjectLiteralExpression([ + ...strings.map((partialString, i) => + ts.factory.createPropertyAssignment( + ts.factory.createNumericLiteral(i + 1), + ts.factory.createStringLiteral(partialString) + ) + ), + ts.factory.createPropertyAssignment("raw", rawStringsArray), ]); + expressions.unshift(stringObject); + // Evaluate if there is a self parameter to be used. const signature = context.checker.getResolvedSignature(expression); const signatureDeclaration = signature?.getDeclaration(); const useSelfParameter = signatureDeclaration && getDeclarationContextType(context, signatureDeclaration) !== ContextType.Void; - // Argument evaluation. - const callArguments = transformArguments(context, expressions, signature); - callArguments.unshift(stringTableLiteral); - if (useSelfParameter) { - return transformContextualCallExpression(context, expression, callArguments); + return transformContextualCallExpression(context, expression, expressions, signature); } + // Argument evaluation. + const callArguments = transformArguments(context, expressions, signature); + const leftHandSideExpression = context.transformExpression(expression.tag); return lua.createCallExpression(leftHandSideExpression, callArguments); }; diff --git a/test/unit/__snapshots__/spread.spec.ts.snap b/test/unit/__snapshots__/spread.spec.ts.snap new file mode 100644 index 000000000..f16bc746a --- /dev/null +++ b/test/unit/__snapshots__/spread.spec.ts.snap @@ -0,0 +1,127 @@ +// Jest Snapshot v1, https://goo.gl/fbAQLP + +exports[`vararg spread optimization $multi 1`] = ` +"local ____exports = {} +function ____exports.__main(self) + local function multi(self, ...) + return ... + end + local function test(self, ...) + return select( + 2, + multi(nil, ...) + ) + end + return test(nil, \\"a\\", \\"b\\", \\"c\\") +end +return ____exports" +`; + +exports[`vararg spread optimization basic use 1`] = ` +"local ____exports = {} +function ____exports.__main(self) + local function pick(self, ...) + local args = {...} + return args[2] + end + local function test(self, ...) + return pick(nil, ...) + end + return test(nil, \\"a\\", \\"b\\", \\"c\\") +end +return ____exports" +`; + +exports[`vararg spread optimization block statement 1`] = ` +"local ____exports = {} +function ____exports.__main(self) + local function pick(self, ...) + local args = {...} + return args[2] + end + local function test(self, ...) + local result + do + result = pick(nil, ...) + end + return result + end + return test(nil, \\"a\\", \\"b\\", \\"c\\") +end +return ____exports" +`; + +exports[`vararg spread optimization body-less arrow function 1`] = ` +"local ____exports = {} +function ____exports.__main(self) + local function pick(self, ...) + local args = {...} + return args[2] + end + local function test(____, ...) + return pick(nil, ...) + end + return test(nil, \\"a\\", \\"b\\", \\"c\\") +end +return ____exports" +`; + +exports[`vararg spread optimization finally clause 1`] = ` +"local ____exports = {} +function ____exports.__main(self) + local function pick(self, ...) + local args = {...} + return args[2] + end + local function test(self, ...) + do + pcall( + function() + error(\\"foobar\\", 0) + end + ) + do + return pick(nil, ...) + end + end + end + return test(nil, \\"a\\", \\"b\\", \\"c\\") +end +return ____exports" +`; + +exports[`vararg spread optimization if statement 1`] = ` +"local ____exports = {} +function ____exports.__main(self) + local function pick(self, ...) + local args = {...} + return args[2] + end + local function test(self, ...) + if true then + return pick(nil, ...) + end + end + return test(nil, \\"a\\", \\"b\\", \\"c\\") +end +return ____exports" +`; + +exports[`vararg spread optimization loop statement 1`] = ` +"local ____exports = {} +function ____exports.__main(self) + local function pick(self, ...) + local args = {...} + return args[2] + end + local function test(self, ...) + repeat + do + return pick(nil, ...) + end + until not false + end + return test(nil, \\"a\\", \\"b\\", \\"c\\") +end +return ____exports" +`; diff --git a/test/unit/annotations/__snapshots__/deprecated.spec.ts.snap b/test/unit/annotations/__snapshots__/deprecated.spec.ts.snap index 3b1d5e2af..750b0596f 100644 --- a/test/unit/annotations/__snapshots__/deprecated.spec.ts.snap +++ b/test/unit/annotations/__snapshots__/deprecated.spec.ts.snap @@ -43,3 +43,13 @@ __TS__ClassExtends(ClassB, ClassA)" `; exports[`pureAbstract removed: diagnostics 1`] = `"main.ts(4,22): error TSTL: '@pureAbstract' has been removed and will no longer have any effect.See https://typescripttolua.github.io/docs/advanced/compiler-annotations#pureabstract for more information."`; + +exports[`vararg deprecation: code 1`] = ` +"function foo(self, ...) +end +function vararg(self, ...) + foo(_G, ...) +end" +`; + +exports[`vararg deprecation: diagnostics 1`] = `"main.ts(6,17): warning TSTL: '@vararg' is deprecated and will be removed in a future update. Please update your code before upgrading to the next release, otherwise your project will no longer compile. See https://typescripttolua.github.io/docs/advanced/compiler-annotations#vararg for more information."`; diff --git a/test/unit/annotations/deprecated.spec.ts b/test/unit/annotations/deprecated.spec.ts index 072ff6372..9c81b2176 100644 --- a/test/unit/annotations/deprecated.spec.ts +++ b/test/unit/annotations/deprecated.spec.ts @@ -33,3 +33,14 @@ test("forRange deprecation", () => { for (const i of forRange(1, 10)) {} `.expectDiagnosticsToMatchSnapshot([annotationDeprecated.code]); }); + +test("vararg deprecation", () => { + util.testModule` + /** @vararg */ + type VarArg = T & { readonly __brand: unique symbol }; + function foo(...args: any[]) {} + function vararg(...args: VarArg) { + foo(...args); + } + `.expectDiagnosticsToMatchSnapshot([annotationDeprecated.code]); +}); diff --git a/test/unit/annotations/vararg.spec.ts b/test/unit/annotations/vararg.spec.ts index 1635ef13a..1b80e73db 100644 --- a/test/unit/annotations/vararg.spec.ts +++ b/test/unit/annotations/vararg.spec.ts @@ -1,3 +1,4 @@ +import { annotationDeprecated } from "../../../src/transformation/utils/diagnostics"; import * as util from "../../util"; const varargDeclaration = ` @@ -19,6 +20,7 @@ test("@vararg", () => { ` .tap(builder => expect(builder.getMainLuaCodeChunk()).not.toMatch("b = ")) .tap(builder => expect(builder.getMainLuaCodeChunk()).not.toMatch("unpack")) + .ignoreDiagnostics([annotationDeprecated.code]) .expectToMatchJsResult(); }); @@ -30,7 +32,9 @@ test("@vararg array access", () => { return c.join("") + b[0]; } return foo("A", "B", "C", "D"); - `.expectToMatchJsResult(); + ` + .ignoreDiagnostics([annotationDeprecated.code]) + .expectToMatchJsResult(); }); test("@vararg global", () => { @@ -41,5 +45,6 @@ test("@vararg global", () => { ` .setLuaFactory(code => `return (function(...) ${code} end)("A", "B", "C", "D")`) .tap(builder => expect(builder.getMainLuaCodeChunk()).not.toMatch("unpack")) + .ignoreDiagnostics([annotationDeprecated.code]) .expectToEqual({ result: "ABCD" }); }); diff --git a/test/unit/language-extensions/__snapshots__/multi.spec.ts.snap b/test/unit/language-extensions/__snapshots__/multi.spec.ts.snap index a280fe331..c5f3d5838 100644 --- a/test/unit/language-extensions/__snapshots__/multi.spec.ts.snap +++ b/test/unit/language-extensions/__snapshots__/multi.spec.ts.snap @@ -4,8 +4,7 @@ exports[`disallow LuaMultiReturn non-numeric access: code 1`] = ` "local ____exports = {} function ____exports.__main(self) local function multi(self, ...) - local args = {...} - return table.unpack(args) + return ... end return select( \\"forEach\\", @@ -19,8 +18,7 @@ exports[`disallow LuaMultiReturn non-numeric access: code 2`] = ` "local ____exports = {} function ____exports.__main(self) local function multi(self, ...) - local args = {...} - return table.unpack(args) + return ... end return ({ multi(nil) @@ -94,8 +92,7 @@ exports[`invalid $multi implicit cast: diagnostics 1`] = `"main.ts(3,20): error exports[`invalid direct $multi function use (const [a = 1] = $multi()): code 1`] = ` "local ____exports = {} local function multi(self, ...) - local args = {...} - return table.unpack(args) + return ... end local a = ____(nil) if a == nil then @@ -110,8 +107,7 @@ exports[`invalid direct $multi function use (const [a = 1] = $multi()): diagnost exports[`invalid direct $multi function use (const [a = 1] = $multi(2)): code 1`] = ` "local ____exports = {} local function multi(self, ...) - local args = {...} - return table.unpack(args) + return ... end local a = ____(nil, 2) if a == nil then @@ -126,8 +122,7 @@ exports[`invalid direct $multi function use (const [a = 1] = $multi(2)): diagnos exports[`invalid direct $multi function use (const [a] = $multi()): code 1`] = ` "local ____exports = {} local function multi(self, ...) - local args = {...} - return table.unpack(args) + return ... end local a = ____(nil) ____exports.a = a @@ -139,8 +134,7 @@ exports[`invalid direct $multi function use (const [a] = $multi()): diagnostics exports[`invalid direct $multi function use (const [a] = $multi(1)): code 1`] = ` "local ____exports = {} local function multi(self, ...) - local args = {...} - return table.unpack(args) + return ... end local a = ____(nil, 1) ____exports.a = a @@ -152,8 +146,7 @@ exports[`invalid direct $multi function use (const [a] = $multi(1)): diagnostics exports[`invalid direct $multi function use (const _ = null, [a] = $multi(1)): code 1`] = ` "local ____exports = {} local function multi(self, ...) - local args = {...} - return table.unpack(args) + return ... end local _ = nil local a = ____(nil, 1) @@ -166,8 +159,7 @@ exports[`invalid direct $multi function use (const _ = null, [a] = $multi(1)): d exports[`invalid direct $multi function use (const ar = [1]; const [a] = $multi(...ar)): code 1`] = ` "local ____exports = {} local function multi(self, ...) - local args = {...} - return table.unpack(args) + return ... end local ar = {1} local a = ____( @@ -183,8 +175,7 @@ exports[`invalid direct $multi function use (const ar = [1]; const [a] = $multi( exports[`invalid direct $multi function use (let a; [a] = $multi()): code 1`] = ` "local ____exports = {} local function multi(self, ...) - local args = {...} - return table.unpack(args) + return ... end local a local ____ = { @@ -201,8 +192,7 @@ exports[`invalid direct $multi function use (let a; [a] = $multi()): diagnostics exports[`invalid direct $multi function use (let a; for ([a] = $multi(1, 2); false; 1) {}): code 1`] = ` "local ____exports = {} local function multi(self, ...) - local args = {...} - return table.unpack(args) + return ... end local a do @@ -224,8 +214,7 @@ exports[`invalid direct $multi function use (let a; for ([a] = $multi(1, 2); fal exports[`invalid direct $multi function use (let a; for (const [a] = $multi(1, 2); false; 1) {}): code 1`] = ` "local ____exports = {} local function multi(self, ...) - local args = {...} - return table.unpack(args) + return ... end local a do @@ -243,8 +232,7 @@ exports[`invalid direct $multi function use (let a; for (const [a] = $multi(1, 2 exports[`invalid direct $multi function use (let a; if ([a] = $multi(1)) { ++a; }): code 1`] = ` "local ____exports = {} local function multi(self, ...) - local args = {...} - return table.unpack(args) + return ... end local a if (function() diff --git a/test/unit/language-extensions/__snapshots__/vararg.spec.ts.snap b/test/unit/language-extensions/__snapshots__/vararg.spec.ts.snap new file mode 100644 index 000000000..819cc27a0 --- /dev/null +++ b/test/unit/language-extensions/__snapshots__/vararg.spec.ts.snap @@ -0,0 +1,37 @@ +// Jest Snapshot v1, https://goo.gl/fbAQLP + +exports[`$vararg invalid use ("const l = $vararg.length"): code 1`] = `"l = #____"`; + +exports[`$vararg invalid use ("const l = $vararg.length"): diagnostics 1`] = `"main.ts(2,19): error TSTL: $vararg can only be used in a spread element ('...$vararg') in global scope."`; + +exports[`$vararg invalid use ("const x = $vararg;"): code 1`] = `"x = ____"`; + +exports[`$vararg invalid use ("const x = $vararg;"): diagnostics 1`] = `"main.ts(2,19): error TSTL: $vararg can only be used in a spread element ('...$vararg') in global scope."`; + +exports[`$vararg invalid use ("for (const x of $vararg) {}"): code 1`] = ` +"for ____, x in ipairs(____) do +end" +`; + +exports[`$vararg invalid use ("for (const x of $vararg) {}"): diagnostics 1`] = `"main.ts(2,25): error TSTL: $vararg can only be used in a spread element ('...$vararg') in global scope."`; + +exports[`$vararg invalid use ("function f(s: string[]) {} f($vararg);"): code 1`] = ` +"function f(self, s) +end +f(_G, ____)" +`; + +exports[`$vararg invalid use ("function f(s: string[]) {} f($vararg);"): diagnostics 1`] = `"main.ts(2,38): error TSTL: $vararg can only be used in a spread element ('...$vararg') in global scope."`; + +exports[`$vararg invalid use ("function foo(...args: string[]) {} function bar() { foo(...$vararg); }"): code 1`] = ` +"function foo(self, ...) +end +function bar(self) + foo( + _G, + table.unpack(____) + ) +end" +`; + +exports[`$vararg invalid use ("function foo(...args: string[]) {} function bar() { foo(...$vararg); }"): diagnostics 1`] = `"main.ts(2,68): error TSTL: $vararg can only be used in a spread element ('...$vararg') in global scope."`; diff --git a/test/unit/language-extensions/vararg.spec.ts b/test/unit/language-extensions/vararg.spec.ts new file mode 100644 index 000000000..f92cc0f96 --- /dev/null +++ b/test/unit/language-extensions/vararg.spec.ts @@ -0,0 +1,38 @@ +import * as path from "path"; +import * as util from "../../util"; +import * as tstl from "../../../src"; +import { invalidVarargUse } from "../../../src/transformation/utils/diagnostics"; + +const varargProjectOptions: tstl.CompilerOptions = { + types: [path.resolve(__dirname, "../../../language-extensions")], +}; + +test.each([ + 'const result = [...$vararg].join("")', + 'let result: string; { result = [...$vararg].join(""); }', + 'let result: string; if (true) { result = [...$vararg].join(""); }', + 'let result: string; do { result = [...$vararg].join(""); } while (false);', +])("$vararg valid use (%p)", statement => { + util.testModule` + ${statement} + export { result }; + ` + .setOptions(varargProjectOptions) + .setLuaFactory(code => `return (function(...) ${code} end)("A", "B", "C", "D")`) + .tap(builder => expect(builder.getMainLuaCodeChunk()).not.toMatch("unpack")) + .expectToEqual({ result: "ABCD" }); +}); + +test.each([ + "const x = $vararg;", + "for (const x of $vararg) {}", + "const l = $vararg.length", + "function f(s: string[]) {} f($vararg);", + "function foo(...args: string[]) {} function bar() { foo(...$vararg); }", +])("$vararg invalid use (%p)", statement => { + util.testModule` + ${statement} + ` + .setOptions(varargProjectOptions) + .expectDiagnosticsToMatchSnapshot([invalidVarargUse.code]); +}); diff --git a/test/unit/spread.spec.ts b/test/unit/spread.spec.ts index 3384b60fb..126ce89cd 100644 --- a/test/unit/spread.spec.ts +++ b/test/unit/spread.spec.ts @@ -1,3 +1,4 @@ +import * as path from "path"; import * as tstl from "../../src"; import * as util from "../util"; import { formatCode } from "../util"; @@ -128,3 +129,303 @@ describe("in object literal", () => { `.expectToMatchJsResult(); }); }); + +describe("vararg spread optimization", () => { + test("basic use", () => { + util.testFunction` + function pick(...args: string[]) { return args[1]; } + function test(...args: string[]) { + return pick(...args); + } + return test("a", "b", "c"); + ` + .expectLuaToMatchSnapshot() + .expectToMatchJsResult(); + }); + + test("body-less arrow function", () => { + util.testFunction` + function pick(...args: string[]) { return args[1]; } + const test = (...args: string[]) => pick(...args); + return test("a", "b", "c"); + ` + .expectLuaToMatchSnapshot() + .expectToMatchJsResult(); + }); + + test("if statement", () => { + util.testFunction` + function pick(...args: string[]) { return args[1]; } + function test(...args: string[]) { + if (true) { + return pick(...args); + } + } + return test("a", "b", "c"); + ` + .expectLuaToMatchSnapshot() + .expectToMatchJsResult(); + }); + + test("loop statement", () => { + util.testFunction` + function pick(...args: string[]) { return args[1]; } + function test(...args: string[]) { + do { + return pick(...args); + } while (false); + } + return test("a", "b", "c"); + ` + .expectLuaToMatchSnapshot() + .expectToMatchJsResult(); + }); + + test("block statement", () => { + util.testFunction` + function pick(...args: string[]) { return args[1]; } + function test(...args: string[]) { + let result: string; + { + result = pick(...args); + } + return result; + } + return test("a", "b", "c"); + ` + .expectLuaToMatchSnapshot() + .expectToMatchJsResult(); + }); + + test("finally clause", () => { + util.testFunction` + function pick(...args: string[]) { return args[1]; } + function test(...args: string[]) { + try { + throw "foobar"; + } catch { + } finally { + return pick(...args); + } + } + return test("a" ,"b", "c"); + ` + .expectLuaToMatchSnapshot() + .expectToMatchJsResult(); + }); + + test("$multi", () => { + util.testFunction` + function multi(...args: string[]) { + return $multi(...args); + } + function test(...args: string[]) { + return multi(...args)[1]; + } + return test("a" ,"b", "c"); + ` + .setOptions({ types: [path.resolve(__dirname, "../../language-extensions")] }) + .expectLuaToMatchSnapshot() + .expectToEqual("b"); + }); +}); + +describe("vararg spread de-optimization", () => { + test("array modification", () => { + util.testFunction` + function pick(...args: string[]) { return args[1]; } + function test(...args: string[]) { + args[1] = "foobar"; + return pick(...args); + } + return test("a", "b", "c"); + `.expectToMatchJsResult(); + }); + + test("array modification in hoisted function", () => { + util.testFunction` + function pick(...args: string[]) { return args[1]; } + function test(...args: string[]) { + hoisted(); + const result = pick(...args); + function hoisted() { args[1] = "foobar"; } + return result; + } + return test("a", "b", "c"); + `.expectToMatchJsResult(); + }); + + test("array modification in secondary hoisted function", () => { + util.testFunction` + function pick(...args: string[]) { return args[1]; } + function test(...args: string[]) { + function triggersHoisted() { hoisted(); } + triggersHoisted(); + const result = pick(...args); + function hoisted() { args[1] = "foobar"; } + return result; + } + return test("a", "b", "c"); + `.expectToMatchJsResult(); + }); +}); + +describe("vararg spread in IIFE", () => { + test("comma operator", () => { + util.testFunction` + function dummy() { return "foobar"; } + function pick(...args: string[]) { return args[1]; } + function test(...args: string[]) { + return (dummy(), pick(...args)); + } + return test("a", "b", "c"); + `.expectToMatchJsResult(); + }); + + test("assignment expression", () => { + util.testFunction` + function pick(...args: string[]) { return args[1]; } + function test(...args: string[]) { + let x: string; + return (x = pick(...args)); + } + return test("a", "b", "c"); + `.expectToMatchJsResult(); + }); + + test("destructured assignment expression", () => { + util.testFunction` + function pick(...args: string[]) { return [args[1]]; } + function test(...args: string[]) { + let x: string; + return ([x] = pick(...args)); + } + return test("a", "b", "c"); + `.expectToMatchJsResult(); + }); + + test("property-access assignment expression", () => { + util.testFunction` + function pick(...args: string[]) { return args[1]; } + function test(...args: string[]) { + let x: {val?: string} = {}; + return (x.val = pick(...args)); + } + return test("a", "b", "c"); + `.expectToMatchJsResult(); + }); + + test("binary compound assignment", () => { + util.testFunction` + function pick(...args: number[]) { return args[1]; } + function test(...args: number[]) { + let x = 1; + return x += pick(...args); + } + return test(1, 2, 3); + `.expectToMatchJsResult(); + }); + + test("postfix unary compound assignment", () => { + util.testFunction` + function pick(...args: number[]) { return args[1]; } + function test(...args: number[]) { + let x = [7, 8, 9]; + return x[pick(...args)]++; + } + return test(1, 2, 3); + `.expectToMatchJsResult(); + }); + + test("prefix unary compound assignment", () => { + util.testFunction` + function pick(...args: number[]) { return args[1]; } + function test(...args: number[]) { + let x = [7, 8, 9]; + return ++x[pick(...args)]; + } + return test(1, 2, 3); + `.expectToMatchJsResult(); + }); + + test("try clause", () => { + util.testFunction` + function pick(...args: string[]) { return args[1]; } + function test(...args: string[]) { + try { + return pick(...args) + } catch {} + } + return test("a" ,"b", "c"); + `.expectToMatchJsResult(); + }); + + test("catch clause", () => { + util.testFunction` + function pick(...args: string[]) { return args[1]; } + function test(...args: string[]) { + try { + throw "foobar"; + } catch { + return pick(...args) + } + } + return test("a" ,"b", "c"); + `.expectToMatchJsResult(); + }); + + test("class expression", () => { + util.testFunction` + function pick(...args: string[]) { return args[1]; } + function test(...args: string[]) { + const fooClass = class Foo { foo = pick(...args); }; + return new fooClass().foo; + } + return test("a" ,"b", "c"); + `.expectToMatchJsResult(); + }); + + test("self-referencing function expression", () => { + util.testFunction` + function pick(...args: string[]) { return args[1]; } + const test = function testName(...args: string[]) { + return \`\${typeof testName}:\${pick(...args)}\`; + } + return test("a" ,"b", "c"); + `.expectToMatchJsResult(); + }); + + test("method indirect access (access args)", () => { + util.testFunction` + const obj = { $method: () => obj.arg, arg: "foobar" }; + function getObj(...args: string[]) { obj.arg = args[1]; return obj; } + function test(...args: string[]) { + return getObj(...args).$method(); + } + return test("a" ,"b", "c"); + `.expectToMatchJsResult(); + }); + + test("method indirect access (method args)", () => { + util.testFunction` + const obj = { $pick: (...args: string[]) => args[1] }; + function getObj() { return obj; } + function test(...args: string[]) { + return getObj().$pick(...args); + } + return test("a" ,"b", "c"); + `.expectToMatchJsResult(); + }); + + test("tagged template method indirect access", () => { + util.testFunction` + const obj = { $tag: (t: TemplateStringsArray, ...args: string[]) => args[1] }; + function getObj() { return obj; } + function pick(...args: string[]): string { return args[1]; } + function test(...args: string[]) { + return getObj().$tag\`FOO\${pick(...args)}BAR\`; + } + return test("a" ,"b", "c"); + `.expectToMatchJsResult(); + }); +});