From ce42aa43a8b11a1fdae4313225d9e6012b3d70e7 Mon Sep 17 00:00:00 2001 From: Gabriela Araujo Britto Date: Fri, 22 Feb 2019 16:31:40 -0800 Subject: [PATCH] check usages of class if refactoring a constructor --- .../refactors/convertToNamedParameters.ts | 251 ++++++++++++------ 1 file changed, 166 insertions(+), 85 deletions(-) diff --git a/src/services/refactors/convertToNamedParameters.ts b/src/services/refactors/convertToNamedParameters.ts index 8d4f3d6e3c9..9bcea0e5983 100644 --- a/src/services/refactors/convertToNamedParameters.ts +++ b/src/services/refactors/convertToNamedParameters.ts @@ -31,9 +31,8 @@ namespace ts.refactor.convertToNamedParameters { const functionDeclaration = getFunctionDeclarationAtPosition(file, startPosition, program.getTypeChecker()); if (!functionDeclaration || !cancellationToken) return undefined; - const functionNames = getFunctionDeclarationNames(functionDeclaration); - const groupedReferences = getGroupedReferences(functionNames, program, cancellationToken); - if (checkReferences(functionNames, groupedReferences)) { + const groupedReferences = getGroupedReferences(functionDeclaration, program, cancellationToken); + if (groupedReferences.valid) { const edits = textChanges.ChangeTracker.with(context, t => doChange(file, program, host, t, functionDeclaration, groupedReferences)); return { renameFilename: undefined, renameLocation: undefined, edits }; } @@ -41,7 +40,6 @@ namespace ts.refactor.convertToNamedParameters { return { edits: [] }; } - function doChange(sourceFile: SourceFile, program: Program, host: LanguageServiceHost, changes: textChanges.ChangeTracker, functionDeclaration: ValidFunctionDeclaration, groupedReferences: GroupedReferences): void { const newParamDeclaration = map(createNewParameters(functionDeclaration, program, host), param => getSynthesizedDeepClone(param)); changes.replaceNodeRangeWithNodes( @@ -57,7 +55,7 @@ namespace ts.refactor.convertToNamedParameters { }); - const functionCalls = groupedReferences.calls; + const functionCalls = deduplicate(groupedReferences.functionCalls, (a, b) => a === b); forEach(functionCalls, call => { if (call.arguments && call.arguments.length) { const newArgument = getSynthesizedDeepClone(createNewArgument(functionDeclaration, call.arguments), /*includeTrivia*/ true); @@ -70,99 +68,172 @@ namespace ts.refactor.convertToNamedParameters { }}); } - function getGroupedReferences(functionNames: Node[], program: Program, cancellationToken: CancellationToken): GroupedReferences { - const functionReferences = flatMap(functionNames, name => FindAllReferences.getReferenceEntriesForNode(-1, name, program, program.getSourceFiles(), cancellationToken)); - const groupedReferences = groupReferences(functionReferences); + function getGroupedReferences(functionDeclaration: ValidFunctionDeclaration, program: Program, cancellationToken: CancellationToken): GroupedReferences { + const names = getDeclarationNames(functionDeclaration); + const references = flatMap(names, name => FindAllReferences.getReferenceEntriesForNode(-1, name, program, program.getSourceFiles(), cancellationToken)); + let groupedReferences = groupReferences(references); + + // if the refactored function is a constructor, we must also go through the references to its class + if (isConstructorDeclaration(functionDeclaration)) { + const className = getClassName(functionDeclaration); + groupedReferences = groupClassReferences(groupedReferences, className); + } + + validateReferences(groupedReferences); return groupedReferences; + function getClassName(constructorDeclaration: ValidConstructor): Identifier { + switch (constructorDeclaration.parent.kind) { + case SyntaxKind.ClassDeclaration: + return constructorDeclaration.parent.name; + case SyntaxKind.ClassExpression: + return constructorDeclaration.parent.parent.name; + } + } + function groupReferences(referenceEntries: ReadonlyArray | undefined): GroupedReferences { - const references: GroupedReferences = { calls: [], declarations: [], unhandled: [] }; + const groupedReferences: GroupedReferences = { functionCalls: [], declarations: [], unhandled: [], valid: true }; + forEach(referenceEntries, (entry) => { - const decl = entryToDeclarationName(entry); + const decl = entryToDeclaration(entry); if (decl) { - references.declarations.push(decl); + groupedReferences.declarations.push(decl); return; } + const call = entryToFunctionCall(entry); if (call) { - references.calls.push(call); + groupedReferences.functionCalls.push(call); return; } - const node = entryToNode(entry); - if (node) { - references.unhandled.push(node); - } - }); - return references; - function entryToFunctionCall(entry: FindAllReferences.Entry): CallExpression | NewExpression | undefined { - if (entry.kind !== FindAllReferences.EntryKind.Span && entry.node && entry.node.parent) { - const functionReference = entry.node; - const parent = functionReference.parent; - switch (parent.kind) { - // Function call (foo(...) or super(...)) - case SyntaxKind.CallExpression: - const callExpression = tryCast(parent, isCallExpression); - if (callExpression && callExpression.expression === functionReference) { - return callExpression; - } - break; - // Constructor call (new Foo(...)) - case SyntaxKind.NewExpression: - const newExpression = tryCast(parent, isNewExpression); - if (newExpression && newExpression.expression === functionReference) { - return newExpression; - } - break; - // Method call (x.foo(...)) - case SyntaxKind.PropertyAccessExpression: - const propertyAccessExpression = tryCast(parent, isPropertyAccessExpression); - if (propertyAccessExpression && propertyAccessExpression.parent && propertyAccessExpression.name === functionReference) { - const callExpression = tryCast(propertyAccessExpression.parent, isCallExpression); - if (callExpression && callExpression.expression === propertyAccessExpression) { - return callExpression; - } - } - break; - // Method call (x['foo'](...)) - case SyntaxKind.ElementAccessExpression: - const elementAccessExpression = tryCast(parent, isElementAccessExpression); - if (elementAccessExpression && elementAccessExpression.parent && elementAccessExpression.argumentExpression === functionReference) { - const callExpression = tryCast(elementAccessExpression.parent, isCallExpression); - if (callExpression && callExpression.expression === elementAccessExpression) { - return callExpression; - } - } - break; + groupedReferences.unhandled.push(entry); + }); + return groupedReferences; + } + + function groupClassReferences(groupedReferences: GroupedReferences, className: Identifier): GroupedReferences { + const classReferences: ClassReferences = { accessExpressions: [], typeUsages: [] }; + const unhandledEntries = groupedReferences.unhandled; + const newUnhandledEntries: FindAllReferences.Entry[] = []; + + forEach(unhandledEntries, (entry) => { + if (entry.kind === FindAllReferences.EntryKind.Node && entry.node.symbol === className.symbol) { + const accessExpression = entryToAccessExpression(entry); + if (accessExpression) { + classReferences.accessExpressions.push(accessExpression); + return; + } + + // Only class declarations are allowed to be used as a type (in a heritage clause), + // otherwise `findAllReferences` might not be able to track constructor calls. + if (isClassDeclaration(functionDeclaration.parent)) { + const type = entryToType(entry); + if (type) { + classReferences.typeUsages.push(type); + return; + } } } - return undefined; - } + newUnhandledEntries.push(entry); + }); - function entryToDeclarationName(entry: FindAllReferences.Entry): Node | undefined { - if (entry.kind !== FindAllReferences.EntryKind.Span && entry.node && contains(functionNames, entry.node)) { - return entry.node; + return { ...groupedReferences, classReferences, unhandled: newUnhandledEntries }; + } + + function validateReferences(groupedReferences: GroupedReferences): void { + if (groupedReferences.unhandled.length > 0) { + groupedReferences.valid = false; + } + if (!every(groupedReferences.declarations, decl => contains(names, decl))) { + groupedReferences.valid = false; + } + } + + function entryToFunctionCall(entry: FindAllReferences.Entry): CallExpression | NewExpression | undefined { + if (entry.kind === FindAllReferences.EntryKind.Node && entry.node.parent) { + const functionReference = entry.node; + const parent = functionReference.parent; + switch (parent.kind) { + // Function call (foo(...) or super(...)) + case SyntaxKind.CallExpression: + const callExpression = tryCast(parent, isCallExpression); + if (callExpression && callExpression.expression === functionReference) { + return callExpression; + } + break; + // Constructor call (new Foo(...)) + case SyntaxKind.NewExpression: + const newExpression = tryCast(parent, isNewExpression); + if (newExpression && newExpression.expression === functionReference) { + return newExpression; + } + break; + // Method call (x.foo(...)) + case SyntaxKind.PropertyAccessExpression: + const propertyAccessExpression = tryCast(parent, isPropertyAccessExpression); + if (propertyAccessExpression && propertyAccessExpression.parent && propertyAccessExpression.name === functionReference) { + const callExpression = tryCast(propertyAccessExpression.parent, isCallExpression); + if (callExpression && callExpression.expression === propertyAccessExpression) { + return callExpression; + } + } + break; + // Method call (x["foo"](...)) + case SyntaxKind.ElementAccessExpression: + const elementAccessExpression = tryCast(parent, isElementAccessExpression); + if (elementAccessExpression && elementAccessExpression.parent && elementAccessExpression.argumentExpression === functionReference) { + const callExpression = tryCast(elementAccessExpression.parent, isCallExpression); + if (callExpression && callExpression.expression === elementAccessExpression) { + return callExpression; + } + } + break; } - return undefined; } + return undefined; + } - function entryToNode(entry: FindAllReferences.Entry): Node | undefined { - if (entry.kind !== FindAllReferences.EntryKind.Span && entry.node) { - return entry.node; + function entryToDeclaration(entry: FindAllReferences.Entry): Node | undefined { + if (entry.kind === FindAllReferences.EntryKind.Node && contains(names, entry.node)) { + return entry.node; + } + return undefined; + } + + function entryToAccessExpression(entry: FindAllReferences.Entry): ElementAccessExpression | PropertyAccessExpression | undefined { + if (entry.kind === FindAllReferences.EntryKind.Node && entry.node.parent) { + const reference = entry.node; + const parent = reference.parent; + switch (parent.kind) { + // `C.foo` + case SyntaxKind.PropertyAccessExpression: + const propertyAccessExpression = tryCast(parent, isPropertyAccessExpression); + if (propertyAccessExpression && propertyAccessExpression.expression === reference) { + return propertyAccessExpression; + } + break; + // `C["foo"]` + case SyntaxKind.ElementAccessExpression: + const elementAccessExpression = tryCast(parent, isElementAccessExpression); + if (elementAccessExpression && elementAccessExpression.expression === reference) { + return elementAccessExpression; + } + break; } - return undefined; } + return undefined; } - } - function checkReferences(functionNames: Node[], groupedReferences: GroupedReferences): boolean { - if (groupedReferences.unhandled.length > 0) { - return false; + function entryToType(entry: FindAllReferences.Entry): Node | undefined { + if (entry.kind === FindAllReferences.EntryKind.Node) { + const reference = entry.node; + if (getMeaningFromLocation(reference) === SemanticMeaning.Type || isExpressionWithTypeArgumentsInClassExtendsClause(reference.parent)) { + return reference; + } + } + return undefined; } - if (groupedReferences.declarations.length > functionNames.length) { - return false; - } - return true; } function getFunctionDeclarationAtPosition(file: SourceFile, startPosition: number, checker: TypeChecker): ValidFunctionDeclaration | undefined { @@ -180,7 +251,7 @@ namespace ts.refactor.convertToNamedParameters { return !!functionDeclaration.name && !!functionDeclaration.body && !checker.isImplementationOfOverload(functionDeclaration); case SyntaxKind.Constructor: if (isClassDeclaration(functionDeclaration.parent)) { - return !!functionDeclaration.body && !checker.isImplementationOfOverload(functionDeclaration); + return !!functionDeclaration.body && !!functionDeclaration.parent.name && !checker.isImplementationOfOverload(functionDeclaration); } else { return isValidVariableDeclaration(functionDeclaration.parent.parent) && !!functionDeclaration.body && !checker.isImplementationOfOverload(functionDeclaration); @@ -200,7 +271,7 @@ namespace ts.refactor.convertToNamedParameters { } function isValidVariableDeclaration(node: Node): node is ValidVariableDeclaration { - return isVariableDeclaration(node) && isVarConst(node) && !node.type; + return isVariableDeclaration(node) && isVarConst(node) && isIdentifier(node.name) && !node.type; } } @@ -359,7 +430,7 @@ namespace ts.refactor.convertToNamedParameters { return getTextOfIdentifierOrLiteral(paramDeclaration.name); } - function getFunctionDeclarationNames(functionDeclaration: ValidFunctionDeclaration): Node[] { + function getDeclarationNames(functionDeclaration: ValidFunctionDeclaration): Node[] { switch (functionDeclaration.kind) { case SyntaxKind.FunctionDeclaration: case SyntaxKind.MethodDeclaration: @@ -368,10 +439,14 @@ namespace ts.refactor.convertToNamedParameters { const ctrKeyword = findChildOfKind(functionDeclaration, SyntaxKind.ConstructorKeyword, functionDeclaration.getSourceFile())!; switch (functionDeclaration.parent.kind) { case SyntaxKind.ClassDeclaration: - return [ctrKeyword]; + const classDeclaration = functionDeclaration.parent; + return [classDeclaration.name, ctrKeyword]; case SyntaxKind.ClassExpression: - const name = functionDeclaration.parent.parent.name; - return [ctrKeyword, name]; + const classExpression = functionDeclaration.parent; + const variableDeclaration = functionDeclaration.parent.parent; + const className = classExpression.name; + if (className) return [className, ctrKeyword, variableDeclaration.name]; + return [ctrKeyword, variableDeclaration.name]; default: return Debug.assertNever(functionDeclaration.parent); } case SyntaxKind.ArrowFunction: @@ -384,10 +459,10 @@ namespace ts.refactor.convertToNamedParameters { type ValidParameterNodeArray = NodeArray; - type ValidVariableDeclaration = VariableDeclaration & { type: undefined }; + type ValidVariableDeclaration = VariableDeclaration & { name: Identifier, type: undefined }; interface ValidConstructor extends ConstructorDeclaration { - parent: ClassDeclaration | (ClassExpression & { parent: ValidVariableDeclaration }); + parent: (ClassDeclaration & { name: Identifier }) | (ClassExpression & { parent: ValidVariableDeclaration }); parameters: NodeArray; body: FunctionBody; } @@ -422,8 +497,14 @@ namespace ts.refactor.convertToNamedParameters { } interface GroupedReferences { - calls: (CallExpression | NewExpression)[]; + functionCalls: (CallExpression | NewExpression)[]; declarations: Node[]; - unhandled: Node[]; + classReferences?: ClassReferences; + unhandled: FindAllReferences.Entry[]; + valid: boolean; + } + interface ClassReferences { + accessExpressions: Node[]; + typeUsages: Node[]; } } \ No newline at end of file