diff --git a/language-extensions/index.d.ts b/language-extensions/index.d.ts new file mode 100644 index 000000000..d7270be32 --- /dev/null +++ b/language-extensions/index.d.ts @@ -0,0 +1,2 @@ +declare function $multi(...values: T): MultiReturn; +declare type MultiReturn = T & { readonly " __multiBrand": unique symbol }; diff --git a/package.json b/package.json index 16e3e20a1..017169255 100644 --- a/package.json +++ b/package.json @@ -13,7 +13,8 @@ "files": [ "dist/**/*.js", "dist/**/*.lua", - "dist/**/*.ts" + "dist/**/*.ts", + "language-extensions/**/*.ts" ], "main": "dist/index.js", "types": "dist/index.d.ts", diff --git a/src/lualib/tsconfig.json b/src/lualib/tsconfig.json index b9d6de8ce..f16968dc1 100644 --- a/src/lualib/tsconfig.json +++ b/src/lualib/tsconfig.json @@ -3,7 +3,7 @@ "outDir": "../../dist/lualib", "target": "esnext", "lib": ["esnext"], - "types": [], + "types": ["../../language-extensions"], "skipLibCheck": true, "noUnusedLocals": true, diff --git a/src/transformation/utils/diagnostics.ts b/src/transformation/utils/diagnostics.ts index 0197d90c2..8f69eadd9 100644 --- a/src/transformation/utils/diagnostics.ts +++ b/src/transformation/utils/diagnostics.ts @@ -142,6 +142,32 @@ export const unsupportedVarDeclaration = createErrorDiagnosticFactory( "`var` declarations are not supported. Use `let` or `const` instead." ); +export const invalidMultiFunctionUse = createErrorDiagnosticFactory( + "The $multi function must be called in an expression that is returned." +); + +export const invalidMultiTypeToNonArrayBindingPattern = createErrorDiagnosticFactory( + "Expected an array destructuring pattern." +); + +export const invalidMultiTypeToNonArrayLiteral = createErrorDiagnosticFactory("Expected an array literal."); + +export const invalidMultiTypeToEmptyPatternOrArrayLiteral = createErrorDiagnosticFactory( + "There must be one or more elements specified here." +); + +export const invalidMultiTypeArrayBindingPatternElementInitializer = createErrorDiagnosticFactory( + "This array binding pattern cannot have initializers." +); + +export const invalidMultiTypeArrayLiteralElementInitializer = createErrorDiagnosticFactory( + "This array literal pattern cannot have initializers." +); + +export const unsupportedMultiFunctionAssignment = createErrorDiagnosticFactory( + "Omitted expressions and BindingElements are expected here." +); + export const annotationDeprecated = createWarningDiagnosticFactory( (kind: AnnotationKind) => `'@${kind}' 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. ` + diff --git a/src/transformation/utils/language-extensions.ts b/src/transformation/utils/language-extensions.ts new file mode 100644 index 000000000..8f3b07672 --- /dev/null +++ b/src/transformation/utils/language-extensions.ts @@ -0,0 +1,28 @@ +import * as ts from "typescript"; +import * as path from "path"; + +export enum ExtensionKind { + MultiFunction = "MultiFunction", + MultiType = "MultiType", +} + +function isSourceFileFromLanguageExtensions(sourceFile: ts.SourceFile): boolean { + const extensionDirectory = path.resolve(__dirname, "../../../language-extensions"); + const sourceFileDirectory = path.dirname(path.normalize(sourceFile.fileName)); + return extensionDirectory === sourceFileDirectory; +} + +export function getExtensionKind(declaration: ts.Declaration): ExtensionKind | undefined { + const sourceFile = declaration.getSourceFile(); + if (isSourceFileFromLanguageExtensions(sourceFile)) { + if (ts.isFunctionDeclaration(declaration) && declaration?.name?.text === "$multi") { + return ExtensionKind.MultiFunction; + } + + if (ts.isTypeAliasDeclaration(declaration) && declaration.name.text === "MultiReturn") { + return ExtensionKind.MultiType; + } + + throw new Error("Unknown extension kind"); + } +} diff --git a/src/transformation/visitors/call.ts b/src/transformation/visitors/call.ts index 366ff8041..90959a580 100644 --- a/src/transformation/visitors/call.ts +++ b/src/transformation/visitors/call.ts @@ -11,6 +11,7 @@ import { isValidLuaIdentifier } from "../utils/safe-names"; import { isArrayType, isExpressionWithEvaluationEffect, isInDestructingAssignment } from "../utils/typescript"; import { transformElementAccessArgument } from "./access"; import { transformLuaTableCallExpression } from "./lua-table"; +import { returnsMultiType } from "./language-extensions/multi"; export type PropertyCallExpression = ts.CallExpression & { expression: ts.PropertyAccessExpression }; @@ -250,9 +251,8 @@ export const transformCallExpression: FunctionVisitor = (node // 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 (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); diff --git a/src/transformation/visitors/expression-statement.ts b/src/transformation/visitors/expression-statement.ts index 890695be0..4901a3210 100644 --- a/src/transformation/visitors/expression-statement.ts +++ b/src/transformation/visitors/expression-statement.ts @@ -4,8 +4,18 @@ import { FunctionVisitor } from "../context"; import { transformBinaryExpressionStatement } from "./binary-expression"; import { transformLuaTableExpressionStatement } from "./lua-table"; import { transformUnaryExpressionStatement } from "./unary-expression"; +import { returnsMultiType, transformMultiDestructuringAssignmentStatement } from "./language-extensions/multi"; export const transformExpressionStatement: FunctionVisitor = (node, context) => { + if ( + ts.isBinaryExpression(node.expression) && + node.expression.operatorToken.kind === ts.SyntaxKind.EqualsToken && + ts.isCallExpression(node.expression.right) && + returnsMultiType(context, node.expression.right) + ) { + return transformMultiDestructuringAssignmentStatement(context, node); + } + const luaTableResult = transformLuaTableExpressionStatement(context, node); if (luaTableResult) { return luaTableResult; diff --git a/src/transformation/visitors/function.ts b/src/transformation/visitors/function.ts index d3fb88a5a..424e75efe 100644 --- a/src/transformation/visitors/function.ts +++ b/src/transformation/visitors/function.ts @@ -15,6 +15,7 @@ import { import { LuaLibFeature, transformLuaLibFunction } from "../utils/lualib"; import { peekScope, performHoisting, popScope, pushScope, Scope, ScopeType } from "../utils/scope"; import { transformIdentifier } from "./identifier"; +import { isMultiFunction, transformMultiCallExpressionToReturnStatement } from "./language-extensions/multi"; import { transformExpressionBodyToReturnStatement } from "./return"; import { transformBindingPattern } from "./variable-declaration"; @@ -55,6 +56,10 @@ function isRestParameterReferenced(context: TransformationContext, identifier: l export function transformFunctionBodyContent(context: TransformationContext, body: ts.ConciseBody): lua.Statement[] { if (!ts.isBlock(body)) { + if (ts.isCallExpression(body) && isMultiFunction(context, body)) { + return [transformMultiCallExpressionToReturnStatement(context, body)]; + } + const returnStatement = transformExpressionBodyToReturnStatement(context, body); return [returnStatement]; } diff --git a/src/transformation/visitors/identifier.ts b/src/transformation/visitors/identifier.ts index 54ed2a2f0..6e95411a5 100644 --- a/src/transformation/visitors/identifier.ts +++ b/src/transformation/visitors/identifier.ts @@ -3,13 +3,19 @@ import * as lua from "../../LuaAST"; import { transformBuiltinIdentifierExpression } from "../builtins"; import { FunctionVisitor, TransformationContext } from "../context"; import { isForRangeType } from "../utils/annotations"; -import { invalidForRangeCall } from "../utils/diagnostics"; +import { invalidForRangeCall, invalidMultiFunctionUse } from "../utils/diagnostics"; import { createExportedIdentifier, getSymbolExportScope } from "../utils/export"; import { createSafeName, hasUnsafeIdentifierName } from "../utils/safe-names"; import { getIdentifierSymbolId } from "../utils/symbols"; import { findFirstNodeAbove } from "../utils/typescript"; +import { isMultiFunctionNode } from "./language-extensions/multi"; export function transformIdentifier(context: TransformationContext, identifier: ts.Identifier): lua.Identifier { + if (isMultiFunctionNode(context, identifier)) { + context.diagnostics.push(invalidMultiFunctionUse(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/language-extensions/multi.ts b/src/transformation/visitors/language-extensions/multi.ts new file mode 100644 index 000000000..a6cb2a755 --- /dev/null +++ b/src/transformation/visitors/language-extensions/multi.ts @@ -0,0 +1,188 @@ +import * as ts from "typescript"; +import * as lua from "../../../LuaAST"; +import * as extensions from "../../utils/language-extensions"; +import { TransformationContext } from "../../context"; +import { transformAssignmentLeftHandSideExpression } from "../binary-expression/assignments"; +import { transformIdentifier } from "../identifier"; +import { transformArguments } from "../call"; +import { getDependenciesOfSymbol, createExportedIdentifier } from "../../utils/export"; +import { createLocalOrExportedOrGlobalDeclaration } from "../../utils/lua-ast"; +import { + invalidMultiTypeArrayBindingPatternElementInitializer, + invalidMultiTypeArrayLiteralElementInitializer, + invalidMultiTypeToEmptyPatternOrArrayLiteral, + invalidMultiTypeToNonArrayBindingPattern, + invalidMultiTypeToNonArrayLiteral, + unsupportedMultiFunctionAssignment, + invalidMultiFunctionUse, +} from "../../utils/diagnostics"; +import { assert } from "../../../utils"; + +const isMultiFunctionDeclaration = (declaration: ts.Declaration): boolean => + extensions.getExtensionKind(declaration) === extensions.ExtensionKind.MultiFunction; + +const isMultiTypeDeclaration = (declaration: ts.Declaration): boolean => + extensions.getExtensionKind(declaration) === extensions.ExtensionKind.MultiType; + +export function isMultiFunction(context: TransformationContext, expression: ts.CallExpression): boolean { + const type = context.checker.getTypeAtLocation(expression.expression); + return type.symbol?.declarations?.some(isMultiFunctionDeclaration) ?? false; +} + +export function returnsMultiType(context: TransformationContext, node: ts.CallExpression): boolean { + const signature = context.checker.getResolvedSignature(node); + return signature?.getReturnType().aliasSymbol?.declarations?.some(isMultiTypeDeclaration) ?? false; +} + +export function isMultiFunctionNode(context: TransformationContext, node: ts.Node): boolean { + const type = context.checker.getTypeAtLocation(node); + return type.symbol?.declarations?.some(isMultiFunctionDeclaration) ?? false; +} + +export function transformMultiCallExpressionToReturnStatement( + context: TransformationContext, + expression: ts.Expression +): lua.Statement { + assert(ts.isCallExpression(expression)); + + const expressions = transformArguments(context, expression.arguments); + return lua.createReturnStatement(expressions, expression); +} + +export function transformMultiReturnStatement( + context: TransformationContext, + statement: ts.ReturnStatement +): lua.Statement { + assert(statement.expression); + + return transformMultiCallExpressionToReturnStatement(context, statement.expression); +} + +function transformMultiFunctionArguments( + context: TransformationContext, + expression: ts.CallExpression +): lua.Expression[] | lua.Expression { + if (!isMultiFunction(context, expression)) { + return context.transformExpression(expression); + } + + if (expression.arguments.length === 0) { + return lua.createNilLiteral(expression); + } + + return expression.arguments.map(e => context.transformExpression(e)); +} + +export function transformMultiVariableDeclaration( + context: TransformationContext, + declaration: ts.VariableDeclaration +): lua.Statement[] { + assert(declaration.initializer); + assert(ts.isCallExpression(declaration.initializer)); + + if (!ts.isArrayBindingPattern(declaration.name)) { + context.diagnostics.push(invalidMultiTypeToNonArrayBindingPattern(declaration.name)); + return []; + } + + if (declaration.name.elements.length < 1) { + context.diagnostics.push(invalidMultiTypeToEmptyPatternOrArrayLiteral(declaration.name)); + return []; + } + + if (declaration.name.elements.some(e => ts.isBindingElement(e) && e.initializer)) { + context.diagnostics.push(invalidMultiTypeArrayBindingPatternElementInitializer(declaration.name)); + return []; + } + + if (isMultiFunction(context, declaration.initializer)) { + context.diagnostics.push(invalidMultiFunctionUse(declaration.initializer)); + return []; + } + + const leftIdentifiers: lua.Identifier[] = []; + + for (const element of declaration.name.elements) { + if (ts.isBindingElement(element)) { + if (ts.isIdentifier(element.name)) { + leftIdentifiers.push(transformIdentifier(context, element.name)); + } else { + context.diagnostics.push(unsupportedMultiFunctionAssignment(element)); + } + } else if (ts.isOmittedExpression(element)) { + leftIdentifiers.push(lua.createAnonymousIdentifier(element)); + } + } + + const rightExpressions = transformMultiFunctionArguments(context, declaration.initializer); + return createLocalOrExportedOrGlobalDeclaration(context, leftIdentifiers, rightExpressions, declaration); +} + +export function transformMultiDestructuringAssignmentStatement( + context: TransformationContext, + statement: ts.ExpressionStatement +): lua.Statement[] | undefined { + assert(ts.isBinaryExpression(statement.expression)); + assert(ts.isCallExpression(statement.expression.right)); + + if (!ts.isArrayLiteralExpression(statement.expression.left)) { + context.diagnostics.push(invalidMultiTypeToNonArrayLiteral(statement.expression.left)); + return []; + } + + if (statement.expression.left.elements.some(ts.isBinaryExpression)) { + context.diagnostics.push(invalidMultiTypeArrayLiteralElementInitializer(statement.expression.left)); + return []; + } + + if (statement.expression.left.elements.length < 1) { + context.diagnostics.push(invalidMultiTypeToEmptyPatternOrArrayLiteral(statement.expression.left)); + return []; + } + + if (isMultiFunction(context, statement.expression.right)) { + context.diagnostics.push(invalidMultiFunctionUse(statement.expression.right)); + return []; + } + + const transformLeft = (expression: ts.Expression): lua.AssignmentLeftHandSideExpression => + ts.isOmittedExpression(expression) + ? lua.createAnonymousIdentifier(expression) + : transformAssignmentLeftHandSideExpression(context, expression); + + const leftIdentifiers = statement.expression.left.elements.map(transformLeft); + + const rightExpressions = transformMultiFunctionArguments(context, statement.expression.right); + + const trailingStatements = statement.expression.left.elements.flatMap(expression => { + const symbol = context.checker.getSymbolAtLocation(expression); + const dependentSymbols = symbol ? getDependenciesOfSymbol(context, symbol) : []; + return dependentSymbols.map(symbol => { + const identifierToAssign = createExportedIdentifier(context, lua.createIdentifier(symbol.name)); + return lua.createAssignmentStatement(identifierToAssign, transformLeft(expression)); + }); + }); + + return [lua.createAssignmentStatement(leftIdentifiers, rightExpressions, statement), ...trailingStatements]; +} + +export function findMultiAssignmentViolations( + context: TransformationContext, + node: ts.ObjectLiteralExpression +): ts.Node[] { + const result: ts.Node[] = []; + + for (const element of node.properties) { + if (!ts.isShorthandPropertyAssignment(element)) continue; + const valueSymbol = context.checker.getShorthandAssignmentValueSymbol(element); + if (valueSymbol) { + const declaration = valueSymbol.valueDeclaration; + if (declaration && isMultiFunctionDeclaration(declaration)) { + context.diagnostics.push(invalidMultiFunctionUse(element)); + result.push(element); + } + } + } + + return result; +} diff --git a/src/transformation/visitors/literal.ts b/src/transformation/visitors/literal.ts index 6544367a2..36c7e48c6 100644 --- a/src/transformation/visitors/literal.ts +++ b/src/transformation/visitors/literal.ts @@ -2,7 +2,7 @@ import * as ts from "typescript"; import * as lua from "../../LuaAST"; import { assertNever } from "../../utils"; import { FunctionVisitor, TransformationContext, Visitors } from "../context"; -import { unsupportedAccessorInObjectLiteral } from "../utils/diagnostics"; +import { unsupportedAccessorInObjectLiteral, invalidMultiFunctionUse } from "../utils/diagnostics"; import { createExportedIdentifier, getSymbolExportScope } from "../utils/export"; import { LuaLibFeature, transformLuaLibFunction } from "../utils/lualib"; import { createSafeName, hasUnsafeIdentifierName, hasUnsafeSymbolName } from "../utils/safe-names"; @@ -10,6 +10,7 @@ import { getSymbolIdOfSymbol, trackSymbolReference } from "../utils/symbols"; import { isArrayType } from "../utils/typescript"; import { transformFunctionLikeDeclaration } from "./function"; import { flattenSpreadExpressions } from "./call"; +import { findMultiAssignmentViolations } from "./language-extensions/multi"; // TODO: Move to object-literal.ts? export function transformPropertyName(context: TransformationContext, node: ts.PropertyName): lua.Expression { @@ -62,6 +63,12 @@ const transformNumericLiteralExpression: FunctionVisitor = ex }; const transformObjectLiteralExpression: FunctionVisitor = (expression, context) => { + const violations = findMultiAssignmentViolations(context, expression); + if (violations.length > 0) { + context.diagnostics.push(...violations.map(e => invalidMultiFunctionUse(e))); + return lua.createNilLiteral(expression); + } + let properties: lua.TableFieldExpression[] = []; const tableExpressions: lua.Expression[] = []; diff --git a/src/transformation/visitors/return.ts b/src/transformation/visitors/return.ts index c5ec3c740..c13d571ff 100644 --- a/src/transformation/visitors/return.ts +++ b/src/transformation/visitors/return.ts @@ -6,6 +6,7 @@ import { validateAssignment } from "../utils/assignment-validation"; import { createUnpackCall, wrapInTable } from "../utils/lua-ast"; import { ScopeType, walkScopesUp } from "../utils/scope"; import { isArrayType } from "../utils/typescript"; +import { isMultiFunction, transformMultiReturnStatement } from "./language-extensions/multi"; function transformExpressionsInReturn( context: TransformationContext, @@ -59,6 +60,14 @@ export const transformReturnStatement: FunctionVisitor = (st insideTryCatch = insideTryCatch || scope.type === ScopeType.Try || scope.type === ScopeType.Catch; } + if ( + statement.expression && + ts.isCallExpression(statement.expression) && + isMultiFunction(context, statement.expression) + ) { + return transformMultiReturnStatement(context, statement); + } + let results: lua.Expression[]; if (statement.expression) { diff --git a/src/transformation/visitors/variable-declaration.ts b/src/transformation/visitors/variable-declaration.ts index bec7548e5..6ddd4acdc 100644 --- a/src/transformation/visitors/variable-declaration.ts +++ b/src/transformation/visitors/variable-declaration.ts @@ -10,6 +10,7 @@ import { createLocalOrExportedOrGlobalDeclaration, createUnpackCall } from "../u import { LuaLibFeature, transformLuaLibFunction } from "../utils/lualib"; import { transformIdentifier } from "./identifier"; import { transformPropertyName } from "./literal"; +import { returnsMultiType, transformMultiVariableDeclaration } from "./language-extensions/multi"; export function transformArrayBindingElement( context: TransformationContext, @@ -229,6 +230,14 @@ export function transformVariableDeclaration( context: TransformationContext, statement: ts.VariableDeclaration ): lua.Statement[] { + if ( + statement.initializer && + ts.isCallExpression(statement.initializer) && + returnsMultiType(context, statement.initializer) + ) { + return transformMultiVariableDeclaration(context, statement); + } + if (statement.initializer && statement.type) { const initializerType = context.checker.getTypeAtLocation(statement.initializer); const varType = context.checker.getTypeFromTypeNode(statement.type); diff --git a/src/transpilation/index.ts b/src/transpilation/index.ts index 628246ec9..608b42818 100644 --- a/src/transpilation/index.ts +++ b/src/transpilation/index.ts @@ -59,14 +59,22 @@ export function createVirtualProgram(input: Record, options: Com return ts.createSourceFile(fileName, input[fileName], ts.ScriptTarget.Latest, false); } + let filePath: string | undefined; + if (fileName.startsWith("lib.")) { - if (libCache[fileName]) return libCache[fileName]; const typeScriptDir = path.dirname(require.resolve("typescript")); - const filePath = path.join(typeScriptDir, fileName); - const content = fs.readFileSync(filePath, "utf8"); + filePath = path.join(typeScriptDir, fileName); + } - libCache[fileName] = ts.createSourceFile(fileName, content, ts.ScriptTarget.Latest, false); + if (fileName.includes("language-extensions")) { + const dtsName = fileName.replace(/(\.d)?(\.ts)$/, ".d.ts"); + filePath = path.resolve(dtsName); + } + if (filePath !== undefined) { + if (libCache[fileName]) return libCache[fileName]; + const content = fs.readFileSync(filePath, "utf8"); + libCache[fileName] = ts.createSourceFile(filePath, content, ts.ScriptTarget.Latest, false); return libCache[fileName]; } }, diff --git a/test/tsconfig.json b/test/tsconfig.json index de0b21bdd..2c06244c7 100644 --- a/test/tsconfig.json +++ b/test/tsconfig.json @@ -13,6 +13,7 @@ "cli/watch", "transpile/directories", "transpile/outFile", - "../src/lualib" + "../src/lualib", + "../language-extensions" ] } diff --git a/test/unit/language-extensions/__snapshots__/multi.spec.ts.snap b/test/unit/language-extensions/__snapshots__/multi.spec.ts.snap new file mode 100644 index 000000000..587cdfd90 --- /dev/null +++ b/test/unit/language-extensions/__snapshots__/multi.spec.ts.snap @@ -0,0 +1,135 @@ +// Jest Snapshot v1, https://goo.gl/fbAQLP + +exports[`invalid $multi call ($multi()): code 1`] = `"____(_G)"`; + +exports[`invalid $multi call ($multi()): diagnostics 1`] = `"main.ts(2,9): error TSTL: The $multi function must be called in an expression that is returned."`; + +exports[`invalid $multi call ($multi): code 1`] = `"local ____ = ____"`; + +exports[`invalid $multi call ($multi): diagnostics 1`] = `"main.ts(2,9): error TSTL: The $multi function must be called in an expression that is returned."`; + +exports[`invalid $multi call (([a] = $multi(1)) => {}): code 1`] = ` +"local function ____(____, ____bindingPattern0) + if ____bindingPattern0 == nil then + ____bindingPattern0 = ____(_G, 1) + end + local a = ____bindingPattern0[1] +end" +`; + +exports[`invalid $multi call (([a] = $multi(1)) => {}): diagnostics 1`] = `"main.ts(2,16): error TSTL: The $multi function must be called in an expression that is returned."`; + +exports[`invalid $multi call (({ $multi });): code 1`] = `"local ____ = nil"`; + +exports[`invalid $multi call (({ $multi });): diagnostics 1`] = `"main.ts(2,12): error TSTL: The $multi function must be called in an expression that is returned."`; + +exports[`invalid $multi call (const [a = 0] = $multi()): code 1`] = `""`; + +exports[`invalid $multi call (const [a = 0] = $multi()): diagnostics 1`] = `"main.ts(2,15): error TSTL: This array binding pattern cannot have initializers."`; + +exports[`invalid $multi call (const {} = $multi();): code 1`] = `""`; + +exports[`invalid $multi call (const {} = $multi();): diagnostics 1`] = `"main.ts(2,15): error TSTL: Expected an array destructuring pattern."`; + +exports[`invalid $multi call (const a = $multi();): code 1`] = `""`; + +exports[`invalid $multi call (const a = $multi();): diagnostics 1`] = `"main.ts(2,15): error TSTL: Expected an array destructuring pattern."`; + +exports[`invalid direct $multi function use (const [a] = $multi()): code 1`] = ` +"local ____exports = {} +local function multi(self, ...) + local args = {...} + return table.unpack(args) +end +____exports.a = a +return ____exports" +`; + +exports[`invalid direct $multi function use (const [a] = $multi()): diagnostics 1`] = `"main.ts(7,21): error TSTL: The $multi function must be called in an expression that is returned."`; + +exports[`invalid direct $multi function use (const [a] = $multi(1)): code 1`] = ` +"local ____exports = {} +local function multi(self, ...) + local args = {...} + return table.unpack(args) +end +____exports.a = a +return ____exports" +`; + +exports[`invalid direct $multi function use (const [a] = $multi(1)): diagnostics 1`] = `"main.ts(7,21): error TSTL: The $multi function must be called in an expression that is returned."`; + +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) +end +local _ = nil +____exports.a = a +return ____exports" +`; + +exports[`invalid direct $multi function use (const _ = null, [a] = $multi(1)): diagnostics 1`] = `"main.ts(7,31): error TSTL: The $multi function must be called in an expression that is returned."`; + +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) +end +local ar = {1} +____exports.a = a +return ____exports" +`; + +exports[`invalid direct $multi function use (const ar = [1]; const [a] = $multi(...ar)): diagnostics 1`] = `"main.ts(7,37): error TSTL: The $multi function must be called in an expression that is returned."`; + +exports[`invalid direct $multi function use (let a; [a] = $multi()): code 1`] = ` +"local ____exports = {} +local function multi(self, ...) + local args = {...} + return table.unpack(args) +end +local a +____exports.a = a +return ____exports" +`; + +exports[`invalid direct $multi function use (let a; [a] = $multi()): diagnostics 1`] = `"main.ts(7,22): error TSTL: The $multi function must be called in an expression that is returned."`; + +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) +end +local a +do + while false do + local ____ = 1 + end +end +____exports.a = a +return ____exports" +`; + +exports[`invalid direct $multi function use (let a; for ([a] = $multi(1, 2); false; 1) {}): diagnostics 1`] = `"main.ts(7,27): error TSTL: The $multi function must be called in an expression that is returned."`; + +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) +end +local a +do + while false do + local ____ = 1 + end +end +____exports.a = a +return ____exports" +`; + +exports[`invalid direct $multi function use (let a; for (const [a] = $multi(1, 2); false; 1) {}): diagnostics 1`] = `"main.ts(7,33): error TSTL: The $multi function must be called in an expression that is returned."`; diff --git a/test/unit/language-extensions/multi.spec.ts b/test/unit/language-extensions/multi.spec.ts new file mode 100644 index 000000000..ed5353bbc --- /dev/null +++ b/test/unit/language-extensions/multi.spec.ts @@ -0,0 +1,128 @@ +import * as path from "path"; +import * as util from "../../util"; +import * as tstl from "../../../src"; +import { + invalidMultiFunctionUse, + invalidMultiTypeToNonArrayBindingPattern, + invalidMultiTypeArrayBindingPatternElementInitializer, +} from "../../../src/transformation/utils/diagnostics"; + +const multiProjectOptions: tstl.CompilerOptions = { + types: [path.resolve(__dirname, "../../../language-extensions")], +}; + +test("multi example use case", () => { + util.testModule` + function multiReturn(): MultiReturn<[string, number]> { + return $multi("foo", 5); + } + + const [a, b] = multiReturn(); + export { a, b }; + ` + .setOptions(multiProjectOptions) + .expectToEqual({ a: "foo", b: 5 }); +}); + +test.each<[string, any]>([ + ["$multi()", undefined], + ["$multi(true)", true], + ["$multi(1, 2)", 1], +])("$multi call on return statement (%s)", (expression, result) => { + util.testFunction` + return ${expression}; + ` + .setOptions(multiProjectOptions) + .expectToEqual(result); +}); + +const multiFunction = ` +function multi(...args) { + return $multi(...args); +} +`; + +const createCasesThatCall = (name: string): Array<[string, any]> => [ + [`let a; [a] = ${name}()`, undefined], + [`const [a] = ${name}()`, undefined], + [`const [a] = ${name}(1)`, 1], + [`const ar = [1]; const [a] = ${name}(...ar)`, 1], + [`const _ = null, [a] = ${name}(1)`, 1], + [`let a; for (const [a] = ${name}(1, 2); false; 1) {}`, undefined], + [`let a; for ([a] = ${name}(1, 2); false; 1) {}`, 1], +]; + +test.each<[string, any]>(createCasesThatCall("$multi"))("invalid direct $multi function use (%s)", statement => { + util.testModule` + ${multiFunction} + ${statement} + export { a }; + ` + .setOptions(multiProjectOptions) + .setReturnExport("a") + .expectDiagnosticsToMatchSnapshot([invalidMultiFunctionUse.code]); +}); + +test.each<[string, any]>(createCasesThatCall("multi"))( + "valid indirect $multi function use (%s)", + (statement, result) => { + util.testModule` + ${multiFunction} + ${statement} + export { a }; + ` + .setOptions(multiProjectOptions) + .setReturnExport("a") + .expectToEqual(result); + } +); + +test.each<[string, number[]]>([ + ["$multi", [invalidMultiFunctionUse.code]], + ["$multi()", [invalidMultiFunctionUse.code]], + ["({ $multi });", [invalidMultiFunctionUse.code]], + ["const a = $multi();", [invalidMultiTypeToNonArrayBindingPattern.code]], + ["const {} = $multi();", [invalidMultiTypeToNonArrayBindingPattern.code]], + ["([a] = $multi(1)) => {}", [invalidMultiFunctionUse.code]], + ["const [a = 0] = $multi()", [invalidMultiTypeArrayBindingPatternElementInitializer.code]], +])("invalid $multi call (%s)", (statement, diagnostics) => { + util.testModule` + ${statement} + ` + .setOptions(multiProjectOptions) + .expectDiagnosticsToMatchSnapshot(diagnostics); +}); + +test("function to spread multi type result from multi type function", () => { + util.testFunction` + ${multiFunction} + function m() { + return $multi(...multi(true)); + } + return m(); + ` + .setOptions(multiProjectOptions) + .expectToEqual(true); +}); + +test("$multi call with destructuring assignment side effects", () => { + util.testModule` + ${multiFunction} + let a; + export { a }; + [a] = multi(1); + ` + .setOptions(multiProjectOptions) + .setReturnExport("a") + .expectToEqual(1); +}); + +test("allow $multi call in ArrowFunction body", () => { + util.testFunction` + const call = () => $multi(1); + const [result] = call(); + return result; + ` + .setOptions(multiProjectOptions) + .expectToEqual(1); +});