diff --git a/src/compiler/diagnosticMessages.json b/src/compiler/diagnosticMessages.json index 0bd4693875e..cdfacf6a9ee 100644 --- a/src/compiler/diagnosticMessages.json +++ b/src/compiler/diagnosticMessages.json @@ -5116,6 +5116,10 @@ "category": "Message", "code": 95088 }, + "Add 'await' to initializers": { + "category": "Message", + "code": 95089 + }, "No value exists in scope for the shorthand property '{0}'. Either declare one or provide an initializer.": { "category": "Error", diff --git a/src/services/codefixes/addMissingAwait.ts b/src/services/codefixes/addMissingAwait.ts index 3545c65faf9..57cb3e9ad57 100644 --- a/src/services/codefixes/addMissingAwait.ts +++ b/src/services/codefixes/addMissingAwait.ts @@ -31,7 +31,7 @@ namespace ts.codefix { errorCodes, getCodeActions: context => { const { sourceFile, errorCode, span, cancellationToken, program } = context; - const expression = getAwaitableExpression(sourceFile, errorCode, span, cancellationToken, program); + const expression = getFixableErrorSpanExpression(sourceFile, errorCode, span, cancellationToken, program); if (!expression) { return; } @@ -45,32 +45,40 @@ namespace ts.codefix { getAllCodeActions: context => { const { sourceFile, program, cancellationToken } = context; const checker = context.program.getTypeChecker(); + const fixedDeclarations = createMap(); return codeFixAll(context, errorCodes, (t, diagnostic) => { - const expression = getAwaitableExpression(sourceFile, diagnostic.code, diagnostic, cancellationToken, program); + const expression = getFixableErrorSpanExpression(sourceFile, diagnostic.code, diagnostic, cancellationToken, program); if (!expression) { return; } const trackChanges: ContextualTrackChangesFunction = cb => (cb(t), []); - return getDeclarationSiteFix(context, expression, diagnostic.code, checker, trackChanges) - || getUseSiteFix(context, expression, diagnostic.code, checker, trackChanges); + return getDeclarationSiteFix(context, expression, diagnostic.code, checker, trackChanges, fixedDeclarations) + || getUseSiteFix(context, expression, diagnostic.code, checker, trackChanges, fixedDeclarations); }); }, }); - function getDeclarationSiteFix(context: CodeFixContext | CodeFixAllContext, expression: Expression, errorCode: number, checker: TypeChecker, trackChanges: ContextualTrackChangesFunction) { - const { sourceFile } = context; - const awaitableInitializer = findAwaitableInitializer(expression, sourceFile, checker); - if (awaitableInitializer) { - const initializerChanges = trackChanges(t => makeChange(t, errorCode, sourceFile, checker, awaitableInitializer)); + function getDeclarationSiteFix(context: CodeFixContext | CodeFixAllContext, expression: Expression, errorCode: number, checker: TypeChecker, trackChanges: ContextualTrackChangesFunction, fixedDeclarations?: Map) { + const { sourceFile, program, cancellationToken } = context; + const awaitableInitializers = findAwaitableInitializers(expression, sourceFile, cancellationToken, program, checker); + if (awaitableInitializers) { + const initializerChanges = trackChanges(t => { + forEach(awaitableInitializers.initializers, ({ expression }) => makeChange(t, errorCode, sourceFile, checker, expression, fixedDeclarations)); + if (fixedDeclarations && awaitableInitializers.needsSecondPassForFixAll) { + makeChange(t, errorCode, sourceFile, checker, expression, fixedDeclarations); + } + }); return createCodeFixActionNoFixId( "addMissingAwaitToInitializer", initializerChanges, - [Diagnostics.Add_await_to_initializer_for_0, expression.getText(sourceFile)]); + awaitableInitializers.initializers.length === 1 + ? [Diagnostics.Add_await_to_initializer_for_0, awaitableInitializers.initializers[0].declarationSymbol.name] + : Diagnostics.Add_await_to_initializers); } } - function getUseSiteFix(context: CodeFixContext | CodeFixAllContext, expression: Expression, errorCode: number, checker: TypeChecker, trackChanges: ContextualTrackChangesFunction) { - const changes = trackChanges(t => makeChange(t, errorCode, context.sourceFile, checker, expression)); + function getUseSiteFix(context: CodeFixContext | CodeFixAllContext, expression: Expression, errorCode: number, checker: TypeChecker, trackChanges: ContextualTrackChangesFunction, fixedDeclarations?: Map) { + const changes = trackChanges(t => makeChange(t, errorCode, context.sourceFile, checker, expression, fixedDeclarations)); return createCodeFixAction(fixId, changes, Diagnostics.Add_await, fixId, Diagnostics.Fix_all_expressions_possibly_missing_await); } @@ -84,7 +92,7 @@ namespace ts.codefix { some(relatedInformation, related => related.code === Diagnostics.Did_you_forget_to_use_await.code)); } - function getAwaitableExpression(sourceFile: SourceFile, errorCode: number, span: TextSpan, cancellationToken: CancellationToken, program: Program): Expression | undefined { + function getFixableErrorSpanExpression(sourceFile: SourceFile, errorCode: number, span: TextSpan, cancellationToken: CancellationToken, program: Program): Expression | undefined { const token = getTokenAtPosition(sourceFile, span.start); // Checker has already done work to determine that await might be possible, and has attached // related info to the node, so start by finding the expression that exactly matches up @@ -103,38 +111,117 @@ namespace ts.codefix { : undefined; } - function findAwaitableInitializer(expression: Node, sourceFile: SourceFile, checker: TypeChecker): Expression | undefined { - if (!isIdentifier(expression)) { + interface AwaitableInitializer { + expression: Expression; + declarationSymbol: Symbol; + } + + interface AwaitableInitializers { + initializers: readonly AwaitableInitializer[]; + needsSecondPassForFixAll: boolean; + } + + function findAwaitableInitializers( + expression: Node, + sourceFile: SourceFile, + cancellationToken: CancellationToken, + program: Program, + checker: TypeChecker, + ): AwaitableInitializers | undefined { + const identifiers = getIdentifiersFromErrorSpanExpression(expression, checker); + if (!identifiers) { return; } - const symbol = checker.getSymbolAtLocation(expression); - if (!symbol) { - return; + let isCompleteFix = identifiers.isCompleteFix; + let initializers: AwaitableInitializer[] | undefined; + for (const identifier of identifiers.identifiers) { + const symbol = checker.getSymbolAtLocation(identifier); + if (!symbol) { + continue; + } + + const declaration = tryCast(symbol.valueDeclaration, isVariableDeclaration); + const variableName = declaration && tryCast(declaration.name, isIdentifier); + const variableStatement = getAncestor(declaration, SyntaxKind.VariableStatement); + if (!declaration || !variableStatement || + declaration.type || + !declaration.initializer || + variableStatement.getSourceFile() !== sourceFile || + hasModifier(variableStatement, ModifierFlags.Export) || + !variableName || + !isInsideAwaitableBody(declaration.initializer)) { + isCompleteFix = false; + continue; + } + + const diagnostics = program.getSemanticDiagnostics(sourceFile, cancellationToken); + const isUsedElsewhere = FindAllReferences.Core.eachSymbolReferenceInFile(variableName, checker, sourceFile, reference => { + return identifier !== reference && !symbolReferenceIsAlsoMissingAwait(reference, diagnostics, sourceFile, checker); + }); + + if (isUsedElsewhere) { + isCompleteFix = false; + continue; + } + + (initializers || (initializers = [])).push({ + expression: declaration.initializer, + declarationSymbol: symbol, + }); } + return initializers && { + initializers, + needsSecondPassForFixAll: !isCompleteFix, + }; + } - const declaration = tryCast(symbol.valueDeclaration, isVariableDeclaration); - const variableName = tryCast(declaration && declaration.name, isIdentifier); - const variableStatement = getAncestor(declaration, SyntaxKind.VariableStatement); - if (!declaration || !variableStatement || - declaration.type || - !declaration.initializer || - variableStatement.getSourceFile() !== sourceFile || - hasModifier(variableStatement, ModifierFlags.Export) || - !variableName || - !isInsideAwaitableBody(declaration.initializer)) { - return; + interface Identifiers { + identifiers: readonly Identifier[]; + isCompleteFix: boolean; + } + + function getIdentifiersFromErrorSpanExpression(expression: Node, checker: TypeChecker): Identifiers | undefined { + if (isPropertyAccessExpression(expression.parent) && isIdentifier(expression.parent.expression)) { + return { identifiers: [expression.parent.expression], isCompleteFix: true }; } - - const isUsedElsewhere = FindAllReferences.Core.eachSymbolReferenceInFile(variableName, checker, sourceFile, identifier => { - return identifier !== expression; - }); - - if (isUsedElsewhere) { - return; + if (isIdentifier(expression)) { + return { identifiers: [expression], isCompleteFix: true }; } + if (isBinaryExpression(expression)) { + let sides: Identifier[] | undefined; + let isCompleteFix = true; + for (const side of [expression.left, expression.right]) { + const type = checker.getTypeAtLocation(side); + if (checker.getPromisedTypeOfPromise(type)) { + if (!isIdentifier(side)) { + isCompleteFix = false; + continue; + } + (sides || (sides = [])).push(side); + } + } + return sides && { identifiers: sides, isCompleteFix }; + } + } - return declaration.initializer; + function symbolReferenceIsAlsoMissingAwait(reference: Identifier, diagnostics: readonly Diagnostic[], sourceFile: SourceFile, checker: TypeChecker) { + const errorNode = isPropertyAccessExpression(reference.parent) ? reference.parent.name : + isBinaryExpression(reference.parent) ? reference.parent : + reference; + const diagnostic = find(diagnostics, diagnostic => + diagnostic.start === errorNode.getStart(sourceFile) && + diagnostic.start + diagnostic.length! === errorNode.getEnd()); + + return diagnostic && contains(errorCodes, diagnostic.code) || + // A Promise is usually not correct in a binary expression (it’s not valid + // in an arithmetic expression and an equality comparison seems unusual), + // but if the other side of the binary expression has an error, the side + // is typed `any` which will squash the error that would identify this + // Promise as an invalid operand. So if the whole binary expression is + // typed `any` as a result, there is a strong likelihood that this Promise + // is accidentally missing `await`. + checker.getTypeAtLocation(errorNode).flags & TypeFlags.Any; } function isInsideAwaitableBody(node: Node) { @@ -147,26 +234,48 @@ namespace ts.codefix { ancestor.parent.kind === SyntaxKind.MethodDeclaration)); } - function makeChange(changeTracker: textChanges.ChangeTracker, errorCode: number, sourceFile: SourceFile, checker: TypeChecker, insertionSite: Expression) { + function makeChange(changeTracker: textChanges.ChangeTracker, errorCode: number, sourceFile: SourceFile, checker: TypeChecker, insertionSite: Expression, fixedDeclarations?: Map) { if (isBinaryExpression(insertionSite)) { - const { left, right } = insertionSite; - const leftType = checker.getTypeAtLocation(left); - const rightType = checker.getTypeAtLocation(right); - const newLeft = checker.getPromisedTypeOfPromise(leftType) ? createAwait(left) : left; - const newRight = checker.getPromisedTypeOfPromise(rightType) ? createAwait(right) : right; - changeTracker.replaceNode(sourceFile, left, newLeft); - changeTracker.replaceNode(sourceFile, right, newRight); + for (const side of [insertionSite.left, insertionSite.right]) { + if (fixedDeclarations && isIdentifier(side)) { + const symbol = checker.getSymbolAtLocation(side); + if (symbol && fixedDeclarations.has(getSymbolId(symbol).toString())) { + continue; + } + } + const type = checker.getTypeAtLocation(side); + const newNode = checker.getPromisedTypeOfPromise(type) ? createAwait(side) : side; + changeTracker.replaceNode(sourceFile, side, newNode); + } } else if (errorCode === propertyAccessCode && isPropertyAccessExpression(insertionSite.parent)) { + if (fixedDeclarations && isIdentifier(insertionSite.parent.expression)) { + const symbol = checker.getSymbolAtLocation(insertionSite.parent.expression); + if (symbol && fixedDeclarations.has(getSymbolId(symbol).toString())) { + return; + } + } changeTracker.replaceNode( sourceFile, insertionSite.parent.expression, createParen(createAwait(insertionSite.parent.expression))); } else if (contains(callableConstructableErrorCodes, errorCode) && isCallOrNewExpression(insertionSite.parent)) { + if (fixedDeclarations && isIdentifier(insertionSite)) { + const symbol = checker.getSymbolAtLocation(insertionSite); + if (symbol && fixedDeclarations.has(getSymbolId(symbol).toString())) { + return; + } + } changeTracker.replaceNode(sourceFile, insertionSite, createParen(createAwait(insertionSite))); } else { + if (fixedDeclarations && isVariableDeclaration(insertionSite.parent) && isIdentifier(insertionSite.parent.name)) { + const symbol = checker.getSymbolAtLocation(insertionSite.parent.name); + if (symbol && !addToSeen(fixedDeclarations, getSymbolId(symbol))) { + return; + } + } changeTracker.replaceNode(sourceFile, insertionSite, createAwait(insertionSite)); } } diff --git a/tests/cases/fourslash/codeFixAddMissingAwait_initializer.ts b/tests/cases/fourslash/codeFixAddMissingAwait_initializer1.ts similarity index 100% rename from tests/cases/fourslash/codeFixAddMissingAwait_initializer.ts rename to tests/cases/fourslash/codeFixAddMissingAwait_initializer1.ts diff --git a/tests/cases/fourslash/codeFixAddMissingAwait_initializer2.ts b/tests/cases/fourslash/codeFixAddMissingAwait_initializer2.ts new file mode 100644 index 00000000000..4f5d604d448 --- /dev/null +++ b/tests/cases/fourslash/codeFixAddMissingAwait_initializer2.ts @@ -0,0 +1,15 @@ +/// +////async function fn(a: Promise) { +//// const x = a; +//// x.toLowerCase(); +////} + +verify.codeFix({ + description: "Add 'await' to initializer for 'x'", + index: 0, + newFileContent: +`async function fn(a: Promise) { + const x = await a; + x.toLowerCase(); +}` +}); diff --git a/tests/cases/fourslash/codeFixAddMissingAwait_initializer3.ts b/tests/cases/fourslash/codeFixAddMissingAwait_initializer3.ts new file mode 100644 index 00000000000..9fd6cad6420 --- /dev/null +++ b/tests/cases/fourslash/codeFixAddMissingAwait_initializer3.ts @@ -0,0 +1,31 @@ +/// +////async function fn(a: number, b: Promise) { +//// const x = b; +//// const y = b; +//// fn(x, b); +//// fn(y, b); +//// x.toFixed(); +//// y.then; +//// +//// b + b; +//// x + b; +//// x + x.toFixed(); +////} + +verify.codeFixAll({ + fixAllDescription: ts.Diagnostics.Fix_all_expressions_possibly_missing_await.message, + fixId: "addMissingAwait", + newFileContent: +`async function fn(a: number, b: Promise) { + const x = await b; + const y = b; + fn(x, b); + fn(await y, b); + x.toFixed(); + y.then; + + await b + await b; + x + await b; + x + x.toFixed(); +}` +}); diff --git a/tests/cases/fourslash/codeFixAddMissingAwait_initializer4.ts b/tests/cases/fourslash/codeFixAddMissingAwait_initializer4.ts new file mode 100644 index 00000000000..3106eaeccb6 --- /dev/null +++ b/tests/cases/fourslash/codeFixAddMissingAwait_initializer4.ts @@ -0,0 +1,28 @@ +/// +////async function fn(a: string, b: Promise) { +//// const x = b; +//// const y = b; +//// x + y; +////} + +verify.codeFix({ + description: "Add 'await' to initializers", + index: 0, + newFileContent: +`async function fn(a: string, b: Promise) { + const x = await b; + const y = await b; + x + y; +}` +}); + +verify.codeFixAll({ + fixAllDescription: ts.Diagnostics.Fix_all_expressions_possibly_missing_await.message, + fixId: "addMissingAwait", + newFileContent: +`async function fn(a: string, b: Promise) { + const x = await b; + const y = await b; + x + y; +}` +}); diff --git a/tests/cases/user/prettier/prettier b/tests/cases/user/prettier/prettier index 23146404850..1e471a00796 160000 --- a/tests/cases/user/prettier/prettier +++ b/tests/cases/user/prettier/prettier @@ -1 +1 @@ -Subproject commit 23146404850011972f695fb6bc2b8113c3cffbfc +Subproject commit 1e471a007968b7490563b91ed6909ae6046f3fe8