From 27eda8cb26be475992d18b00c38c488cc03b172c Mon Sep 17 00:00:00 2001 From: Kirill Gusev Date: Wed, 29 Apr 2026 00:34:48 +0300 Subject: [PATCH 1/3] Add MoveMembersToExtension code action --- .../CodeActions/MoveMembersToExtension.swift | 169 +++++++++ .../CodeActions/SyntaxCodeActions.swift | 2 + .../MoveMembersToExtensionTests.swift | 340 ++++++++++++++++++ 3 files changed, 511 insertions(+) create mode 100644 Sources/SwiftLanguageService/CodeActions/MoveMembersToExtension.swift create mode 100644 Tests/SourceKitLSPTests/MoveMembersToExtensionTests.swift diff --git a/Sources/SwiftLanguageService/CodeActions/MoveMembersToExtension.swift b/Sources/SwiftLanguageService/CodeActions/MoveMembersToExtension.swift new file mode 100644 index 000000000..1abc49334 --- /dev/null +++ b/Sources/SwiftLanguageService/CodeActions/MoveMembersToExtension.swift @@ -0,0 +1,169 @@ +//===----------------------------------------------------------------------===// +// +// This source file is part of the Swift.org open source project +// +// Copyright (c) 2014 - 2026 Apple Inc. and the Swift project authors +// Licensed under Apache License v2.0 with Runtime Library Exception +// +// See https://swift.org/LICENSE.txt for license information +// See https://swift.org/CONTRIBUTORS.txt for the list of Swift project authors +// +//===----------------------------------------------------------------------===// + +@_spi(SourceKitLSP) import LanguageServerProtocol +import SwiftRefactor +import SwiftSyntax + +private enum ValidationResult: CustomStringConvertible { + case accessor + case deinitializer + case enumCase + case storedProperty + + var description: String { + switch self { + case .accessor: return "accessor" + case .deinitializer: return "deinitializer" + case .enumCase: return "enum case" + case .storedProperty: return "stored property" + } + } + + /// Validates that `member` can be moved to an extension. If it can, return `nil`, otherwise return the reason why + /// `member` cannot be moved to an extension. + init?(_ member: MemberBlockItemSyntax) { + switch member.decl.kind { + case .accessorDecl: + self = .accessor + case .deinitializerDecl: + self = .deinitializer + case .enumCaseDecl: + self = .enumCase + default: + if let varDecl = member.decl.as(VariableDeclSyntax.self), + varDecl.bindings.contains(where: { $0.accessorBlock == nil || $0.initializer != nil }) + { + self = .storedProperty + } else { + return nil + } + } + } +} + +struct MoveMembersToExtension: SyntaxRefactoringProvider { + struct Context { + let range: Range + + init(range: Range) { + self.range = range + } + } + + static func refactor(syntax: SourceFileSyntax, in context: Context) throws -> SourceFileSyntax { + guard + let statement = syntax.statements.first(where: { $0.item.range.contains(context.range) }), + let decl = statement.item.asProtocol((any NamedDeclSyntax).self), + let declGroup = statement.item.asProtocol((any DeclGroupSyntax).self), + let statementIndex = syntax.statements.index(of: statement) + else { + throw RefactoringNotApplicableError("Type declaration not found") + } + + let selectedMembers = Array(declGroup.memberBlock.members).filter { context.range.overlaps($0.trimmedRange) } + .map { (member: $0, validationResult: ValidationResult($0)) } + + var membersToMove = selectedMembers.filter({ $0.validationResult == nil }).map(\.member) + + guard !membersToMove.isEmpty else { + let notMovedMembers = Set(selectedMembers.compactMap(\.validationResult)) + .map(\.description) + .sorted().joined(separator: ", ") + throw RefactoringNotApplicableError( + "Cannot move \(notMovedMembers) to extension" + ) + } + + var updatedDeclGroup = declGroup + var remainingMembers = Array(declGroup.memberBlock.members).filter { !membersToMove.contains($0) } + membersToMove[0].decl.leadingTrivia = membersToMove[0].decl.leadingTrivia.trimmingPrefix(while: \.isSpaceOrTab) + + if remainingMembers.isEmpty { + updatedDeclGroup.memberBlock.rightBrace.leadingTrivia = Trivia() + } else { + remainingMembers[0].leadingTrivia = .newline.merging( + remainingMembers[0].leadingTrivia.trimmingPrefix(while: \.isNewline) + ) + remainingMembers[remainingMembers.count - 1].trailingTrivia = remainingMembers[remainingMembers.count - 1] + .trailingTrivia.trimmingSuffix(while: \.isNewline) + } + + updatedDeclGroup.memberBlock.members = MemberBlockItemListSyntax(remainingMembers) + membersToMove[0].leadingTrivia = .newline.merging(membersToMove[0].leadingTrivia.trimmingPrefix(while: \.isNewline)) + let extensionMemberBlockSyntax = declGroup.memberBlock.with(\.members, MemberBlockItemListSyntax(membersToMove)) + + var declName = decl.name + declName.trailingTrivia = declName.trailingTrivia.merging(.space) + + let extensionDecl = ExtensionDeclSyntax( + leadingTrivia: .newlines(2), + extendedType: IdentifierTypeSyntax( + leadingTrivia: .space, + name: declName + ), + memberBlock: extensionMemberBlockSyntax + ) + + var syntax = syntax + let updatedStatement = statement.with(\.item, .decl(DeclSyntax(updatedDeclGroup))) + syntax.statements[statementIndex] = updatedStatement + syntax.statements.insert( + CodeBlockItemSyntax(item: .decl(DeclSyntax(extensionDecl))), + at: syntax.statements.index(after: statementIndex) + ) + return syntax + } +} + +extension MoveMembersToExtension: SyntaxRefactoringCodeActionProvider { + static var title: String { "Move to extension" } + + static func refactoringContext(for scope: SyntaxCodeActionScope) -> Context { + Context(range: scope.range) + } + + static func nodeToRefactor(in scope: SyntaxCodeActionScope) -> SourceFileSyntax? { + guard scope.request.range.lowerBound != scope.request.range.upperBound else { + return nil + } + + return scope.file + } + + static func textRefactor(syntax: SourceFileSyntax, in context: Context) throws -> [SourceEdit] { + let updatedSyntax = try self.refactor(syntax: syntax, in: context) + + return [ + .replace(syntax, with: updatedSyntax.description) + ] + } +} + +fileprivate extension Trivia { + func trimmingPrefix( + while predicate: (TriviaPiece) -> Bool + ) -> Trivia { + Trivia(pieces: self.drop(while: predicate)) + } + + func trimmingSuffix( + while predicate: (TriviaPiece) -> Bool + ) -> Trivia { + Trivia( + pieces: self[...] + .reversed() + .drop(while: predicate) + .reversed() + ) + } +} diff --git a/Sources/SwiftLanguageService/CodeActions/SyntaxCodeActions.swift b/Sources/SwiftLanguageService/CodeActions/SyntaxCodeActions.swift index a481a5950..3a5a856ef 100644 --- a/Sources/SwiftLanguageService/CodeActions/SyntaxCodeActions.swift +++ b/Sources/SwiftLanguageService/CodeActions/SyntaxCodeActions.swift @@ -28,6 +28,7 @@ let allSyntaxCodeActions: [any SyntaxCodeActionProvider.Type] = { ConvertZeroParameterFunctionToComputedProperty.self, FormatRawStringLiteral.self, MigrateToNewIfLetSyntax.self, + MoveMembersToExtension.self, OpaqueParameterToGeneric.self, RemoveRedundantParentheses.self, RemoveSeparatorsFromIntegerLiteral.self, @@ -41,5 +42,6 @@ let allSyntaxCodeActions: [any SyntaxCodeActionProvider.Type] = { let supersededSourcekitdRefactoringActions: Set = [ "source.refactoring.kind.convert.to.computed.property", // Superseded by ConvertStoredPropertyToComputed + "source.refactoring.kind.move.members.to.extension", // Superseded by MoveMembersToExtension "source.refactoring.kind.simplify.long.number.literal", // Superseded by AddSeparatorsToIntegerLiteral ] diff --git a/Tests/SourceKitLSPTests/MoveMembersToExtensionTests.swift b/Tests/SourceKitLSPTests/MoveMembersToExtensionTests.swift new file mode 100644 index 000000000..07141d972 --- /dev/null +++ b/Tests/SourceKitLSPTests/MoveMembersToExtensionTests.swift @@ -0,0 +1,340 @@ +//===----------------------------------------------------------------------===// +// +// This source file is part of the Swift.org open source project +// +// Copyright (c) 2014 - 2026 Apple Inc. and the Swift project authors +// Licensed under Apache License v2.0 with Runtime Library Exception +// +// See https://swift.org/LICENSE.txt for license information +// See https://swift.org/CONTRIBUTORS.txt for the list of Swift project authors +// +//===----------------------------------------------------------------------===// + +@_spi(SourceKitLSP) import LanguageServerProtocol +import SKLogging +import SKTestSupport +import SourceKitLSP +import SwiftExtensions +@_spi(Testing) import SwiftLanguageService +import SwiftParser +import SwiftRefactor +import SwiftSyntax +import SwiftSyntaxBuilder +import XCTest + +private typealias CodeActionCapabilities = TextDocumentClientCapabilities.CodeAction +private typealias CodeActionLiteralSupport = CodeActionCapabilities.CodeActionLiteralSupport +private typealias CodeActionKindCapabilities = CodeActionLiteralSupport.CodeActionKindValueSet + +private let clientCapabilitiesWithCodeActionSupport: ClientCapabilities = { + var documentCapabilities = TextDocumentClientCapabilities() + var codeActionCapabilities = CodeActionCapabilities() + let codeActionKinds = CodeActionKindCapabilities(valueSet: [.refactor, .quickFix]) + let codeActionLiteralSupport = CodeActionLiteralSupport(codeActionKind: codeActionKinds) + codeActionCapabilities.codeActionLiteralSupport = codeActionLiteralSupport + documentCapabilities.codeAction = codeActionCapabilities + documentCapabilities.completion = .init(completionItem: .init(snippetSupport: true)) + return ClientCapabilities(workspace: nil, textDocument: documentCapabilities) +}() + +final class MoveMembersToExtensionTests: SourceKitLSPTestCase { + func testMoveMembersToExtension() async throws { + try await assertMoveMembersToExtensionCodeAction( + """ + 1️⃣class Foo { + 2️⃣func foo() { + print("Hello world!") + }3️⃣ + + func bar() { + print("Hello world!") + } + }4️⃣ + """, + expected: + """ + class Foo { + func bar() { + print("Hello world!") + } + } + + extension Foo { + func foo() { + print("Hello world!") + } + } + """ + ) + } + + func testMoveParticiallySelectedFunctionFromClass() async throws { + try await assertMoveMembersToExtensionCodeAction( + """ + 1️⃣class Foo { + func foo() { + print("Hello world!") + } + + func bar() { + 2️⃣print("Hello world!") + }3️⃣ + } + + struct Bar { + func foo() {} + }4️⃣ + """, + expected: + """ + class Foo { + func foo() { + print("Hello world!") + } + } + + extension Foo { + func bar() { + print("Hello world!") + } + } + + struct Bar { + func foo() {} + } + """ + ) + } + + func testMoveSelectedFromClass() async throws { + try await assertMoveMembersToExtensionCodeAction( + """ + 1️⃣class Foo {2️⃣ + func foo() { + print("Hello world!") + } + + deinit() {} + + func bar() { + print("Hello world!") + }3️⃣ + } + + struct Bar { + func foo() {} + }4️⃣ + """, + expected: + """ + class Foo { + deinit() {} + } + + extension Foo { + func foo() { + print("Hello world!") + } + + func bar() { + print("Hello world!") + } + } + + struct Bar { + func foo() {} + } + """ + ) + } + + func testMoveNestedFromStruct() async throws { + try await assertMoveMembersToExtensionCodeAction( + """ + 1️⃣struct Outer {2️⃣ + struct Inner { + func moveThis() {} + }3️⃣ + }4️⃣ + """, + expected: + """ + struct Outer {} + + extension Outer { + struct Inner { + func moveThis() {} + } + } + """ + ) + } + + func testMoveNestedFromStruct2() async throws { + try await assertMoveMembersToExtensionCodeAction( + """ + 1️⃣struct Outer {2️⃣ + struct Inner { + func moveThis() {} + }3️⃣ + }4️⃣ + """, + expected: + """ + struct Outer {} + + extension Outer { + struct Inner { + func moveThis() {} + } + } + """ + ) + } + + func testMoveSelectedFunctionName() async throws { + try await assertMoveMembersToExtensionCodeAction( + """ + 1️⃣struct Outer { + struct Inner { + func 2️⃣moveThis()3️⃣ {} + } + }4️⃣ + """, + expected: + """ + struct Outer {} + + extension Outer { + struct Inner { + func moveThis() {} + } + } + """ + ) + } + + func testSelectedDeinitializerMember() async throws { + let source = """ + 1️⃣class Foo { + func foo() { + print("Hello world!") + } + + 2️⃣deinit() {}3️⃣ + + func bar() { + print("Hello world!") + } + } + + struct Bar { + func foo() {} + }4️⃣ + """ + + let testClient = try await TestSourceKitLSPClient(capabilities: clientCapabilitiesWithCodeActionSupport) + let uri = DocumentURI(for: .swift) + + let positions = testClient.openDocument(source, uri: uri) + + let request = CodeActionRequest( + range: positions["2️⃣"].. Date: Fri, 7 Aug 2026 14:52:19 +0300 Subject: [PATCH 2/3] Update branch, update tests --- Package.swift | 3 +- .../MoveMembersToExtension.swift | 39 +- .../MoveMembersToExtensionTests.swift | 340 ------------------ .../MoveMembersToExtensionTests.swift | 266 ++++++++++++++ 4 files changed, 290 insertions(+), 358 deletions(-) delete mode 100644 Tests/SourceKitLSPTests/MoveMembersToExtensionTests.swift create mode 100644 Tests/SwiftSyntaxCodeActionsTests/MoveMembersToExtensionTests.swift diff --git a/Package.swift b/Package.swift index 4c9c5cae3..0b249309b 100644 --- a/Package.swift +++ b/Package.swift @@ -543,7 +543,8 @@ var targets: [Target] = [ .testTarget( name: "SwiftSyntaxCodeActionsTests", dependencies: [ - "SwiftSyntaxCodeActions" + "SwiftSyntaxCodeActions", + "SKTestSupport" ] + swiftSyntaxDependencies([ "SwiftParser", diff --git a/Sources/SwiftSyntaxCodeActions/MoveMembersToExtension.swift b/Sources/SwiftSyntaxCodeActions/MoveMembersToExtension.swift index 1abc49334..781907386 100644 --- a/Sources/SwiftSyntaxCodeActions/MoveMembersToExtension.swift +++ b/Sources/SwiftSyntaxCodeActions/MoveMembersToExtension.swift @@ -12,7 +12,7 @@ @_spi(SourceKitLSP) import LanguageServerProtocol import SwiftRefactor -import SwiftSyntax +package import SwiftSyntax private enum ValidationResult: CustomStringConvertible { case accessor @@ -51,16 +51,16 @@ private enum ValidationResult: CustomStringConvertible { } } -struct MoveMembersToExtension: SyntaxRefactoringProvider { - struct Context { +package struct MoveMembersToExtension: SyntaxRefactoringProvider { + package struct Context { let range: Range - init(range: Range) { + package init(range: Range) { self.range = range } } - static func refactor(syntax: SourceFileSyntax, in context: Context) throws -> SourceFileSyntax { + package static func refactor(syntax: SourceFileSyntax, in context: Context) throws -> SourceFileSyntax { guard let statement = syntax.statements.first(where: { $0.item.range.contains(context.range) }), let decl = statement.item.asProtocol((any NamedDeclSyntax).self), @@ -125,28 +125,33 @@ struct MoveMembersToExtension: SyntaxRefactoringProvider { } } -extension MoveMembersToExtension: SyntaxRefactoringCodeActionProvider { - static var title: String { "Move to extension" } +extension MoveMembersToExtension: ResolvableSyntaxRefactoringCodeActionProvider { + static func refactoringContext( + for node: SwiftSyntax.SourceFileSyntax, + in scope: SyntaxCodeActionScope + ) -> RefactoringContext { + .context(Context(range: scope.range)) + } - static func refactoringContext(for scope: SyntaxCodeActionScope) -> Context { + static func resolveContext( + for data: UnresolvedData, + in scope: SyntaxCodeActionScope, + symbolInfo: (_ position: Position) async throws -> [SymbolDetails] + ) async throws -> Context { Context(range: scope.range) } + typealias UnresolvedData = EmptyLSPCodable + + static var title: String { "Move to extension" } + static func nodeToRefactor(in scope: SyntaxCodeActionScope) -> SourceFileSyntax? { - guard scope.request.range.lowerBound != scope.request.range.upperBound else { + guard scope.range.lowerBound != scope.range.upperBound else { return nil } return scope.file } - - static func textRefactor(syntax: SourceFileSyntax, in context: Context) throws -> [SourceEdit] { - let updatedSyntax = try self.refactor(syntax: syntax, in: context) - - return [ - .replace(syntax, with: updatedSyntax.description) - ] - } } fileprivate extension Trivia { diff --git a/Tests/SourceKitLSPTests/MoveMembersToExtensionTests.swift b/Tests/SourceKitLSPTests/MoveMembersToExtensionTests.swift deleted file mode 100644 index 07141d972..000000000 --- a/Tests/SourceKitLSPTests/MoveMembersToExtensionTests.swift +++ /dev/null @@ -1,340 +0,0 @@ -//===----------------------------------------------------------------------===// -// -// This source file is part of the Swift.org open source project -// -// Copyright (c) 2014 - 2026 Apple Inc. and the Swift project authors -// Licensed under Apache License v2.0 with Runtime Library Exception -// -// See https://swift.org/LICENSE.txt for license information -// See https://swift.org/CONTRIBUTORS.txt for the list of Swift project authors -// -//===----------------------------------------------------------------------===// - -@_spi(SourceKitLSP) import LanguageServerProtocol -import SKLogging -import SKTestSupport -import SourceKitLSP -import SwiftExtensions -@_spi(Testing) import SwiftLanguageService -import SwiftParser -import SwiftRefactor -import SwiftSyntax -import SwiftSyntaxBuilder -import XCTest - -private typealias CodeActionCapabilities = TextDocumentClientCapabilities.CodeAction -private typealias CodeActionLiteralSupport = CodeActionCapabilities.CodeActionLiteralSupport -private typealias CodeActionKindCapabilities = CodeActionLiteralSupport.CodeActionKindValueSet - -private let clientCapabilitiesWithCodeActionSupport: ClientCapabilities = { - var documentCapabilities = TextDocumentClientCapabilities() - var codeActionCapabilities = CodeActionCapabilities() - let codeActionKinds = CodeActionKindCapabilities(valueSet: [.refactor, .quickFix]) - let codeActionLiteralSupport = CodeActionLiteralSupport(codeActionKind: codeActionKinds) - codeActionCapabilities.codeActionLiteralSupport = codeActionLiteralSupport - documentCapabilities.codeAction = codeActionCapabilities - documentCapabilities.completion = .init(completionItem: .init(snippetSupport: true)) - return ClientCapabilities(workspace: nil, textDocument: documentCapabilities) -}() - -final class MoveMembersToExtensionTests: SourceKitLSPTestCase { - func testMoveMembersToExtension() async throws { - try await assertMoveMembersToExtensionCodeAction( - """ - 1️⃣class Foo { - 2️⃣func foo() { - print("Hello world!") - }3️⃣ - - func bar() { - print("Hello world!") - } - }4️⃣ - """, - expected: - """ - class Foo { - func bar() { - print("Hello world!") - } - } - - extension Foo { - func foo() { - print("Hello world!") - } - } - """ - ) - } - - func testMoveParticiallySelectedFunctionFromClass() async throws { - try await assertMoveMembersToExtensionCodeAction( - """ - 1️⃣class Foo { - func foo() { - print("Hello world!") - } - - func bar() { - 2️⃣print("Hello world!") - }3️⃣ - } - - struct Bar { - func foo() {} - }4️⃣ - """, - expected: - """ - class Foo { - func foo() { - print("Hello world!") - } - } - - extension Foo { - func bar() { - print("Hello world!") - } - } - - struct Bar { - func foo() {} - } - """ - ) - } - - func testMoveSelectedFromClass() async throws { - try await assertMoveMembersToExtensionCodeAction( - """ - 1️⃣class Foo {2️⃣ - func foo() { - print("Hello world!") - } - - deinit() {} - - func bar() { - print("Hello world!") - }3️⃣ - } - - struct Bar { - func foo() {} - }4️⃣ - """, - expected: - """ - class Foo { - deinit() {} - } - - extension Foo { - func foo() { - print("Hello world!") - } - - func bar() { - print("Hello world!") - } - } - - struct Bar { - func foo() {} - } - """ - ) - } - - func testMoveNestedFromStruct() async throws { - try await assertMoveMembersToExtensionCodeAction( - """ - 1️⃣struct Outer {2️⃣ - struct Inner { - func moveThis() {} - }3️⃣ - }4️⃣ - """, - expected: - """ - struct Outer {} - - extension Outer { - struct Inner { - func moveThis() {} - } - } - """ - ) - } - - func testMoveNestedFromStruct2() async throws { - try await assertMoveMembersToExtensionCodeAction( - """ - 1️⃣struct Outer {2️⃣ - struct Inner { - func moveThis() {} - }3️⃣ - }4️⃣ - """, - expected: - """ - struct Outer {} - - extension Outer { - struct Inner { - func moveThis() {} - } - } - """ - ) - } - - func testMoveSelectedFunctionName() async throws { - try await assertMoveMembersToExtensionCodeAction( - """ - 1️⃣struct Outer { - struct Inner { - func 2️⃣moveThis()3️⃣ {} - } - }4️⃣ - """, - expected: - """ - struct Outer {} - - extension Outer { - struct Inner { - func moveThis() {} - } - } - """ - ) - } - - func testSelectedDeinitializerMember() async throws { - let source = """ - 1️⃣class Foo { - func foo() { - print("Hello world!") - } - - 2️⃣deinit() {}3️⃣ - - func bar() { - print("Hello world!") - } - } - - struct Bar { - func foo() {} - }4️⃣ - """ - - let testClient = try await TestSourceKitLSPClient(capabilities: clientCapabilitiesWithCodeActionSupport) - let uri = DocumentURI(for: .swift) - - let positions = testClient.openDocument(source, uri: uri) - - let request = CodeActionRequest( - range: positions["2️⃣"].. {1️⃣ + struct Inner { + func moveThis() {} + }2️⃣ + } + """, + expected: + """ + struct Outer {} + + extension Outer { + struct Inner { + func moveThis() {} + } + } + """ + ) + } + + func testMoveSelectedFunctionName() throws { + try assertMoveMembersToExtension( + """ + struct Outer { + struct Inner { + func 1️⃣moveThis()2️⃣ {} + } + } + """, + expected: + """ + struct Outer {} + + extension Outer { + struct Inner { + func moveThis() {} + } + } + """ + ) + } + + func testSelectedDeinitializerMember() async throws { + try assertMoveMembersToExtension ( + """ + class Foo { + func foo() { + print("Hello world!") + } + + 1️⃣deinit() {}2️⃣ + + func bar() { + print("Hello world!") + } + } + + struct Bar { + func foo() {} + } + """, + expected: nil + ) + } + + func testMoveEmptySelection() throws { + try assertMoveMembersToExtension( + """ + class Foo { + func foo() { + print("Hello world!") + } + 1️⃣2️⃣ + func bar() { + print("Hello world!") + } + } + + struct Bar { + func foo() {} + } + """, + expected: nil + ) + } +} + +private func assertMoveMembersToExtension( + _ source: String, + expected: SourceFileSyntax?, + file: StaticString = #filePath, + line: UInt = #line +) throws { + let (markers, source) = extractMarkers(source.description) + let positions = markers.mapValues { $0 } + var parser = Parser(source) + let tree = SourceFileSyntax.parse(from: &parser) + + let range = try XCTUnwrap(positions["1️⃣"]).. Date: Mon, 10 Aug 2026 12:06:51 +0300 Subject: [PATCH 3/3] Fix review and update tests --- Package.swift | 2 +- .../MoveMembersToExtension.swift | 363 ++++++++++++++++-- .../MoveMembersToExtensionTests.swift | 186 ++++++--- 3 files changed, 454 insertions(+), 97 deletions(-) diff --git a/Package.swift b/Package.swift index 0b249309b..bdc6b8dd8 100644 --- a/Package.swift +++ b/Package.swift @@ -544,7 +544,7 @@ var targets: [Target] = [ name: "SwiftSyntaxCodeActionsTests", dependencies: [ "SwiftSyntaxCodeActions", - "SKTestSupport" + "SKTestSupport", ] + swiftSyntaxDependencies([ "SwiftParser", diff --git a/Sources/SwiftSyntaxCodeActions/MoveMembersToExtension.swift b/Sources/SwiftSyntaxCodeActions/MoveMembersToExtension.swift index 781907386..592b9facc 100644 --- a/Sources/SwiftSyntaxCodeActions/MoveMembersToExtension.swift +++ b/Sources/SwiftSyntaxCodeActions/MoveMembersToExtension.swift @@ -11,9 +11,42 @@ //===----------------------------------------------------------------------===// @_spi(SourceKitLSP) import LanguageServerProtocol +import SwiftBasicFormat import SwiftRefactor package import SwiftSyntax +package struct MoveMembersToExtension: SyntaxRefactoringProvider { + package struct Context { + let range: Range + + package init(range: Range) { + self.range = range + } + } + + package static func refactor(syntax: SourceFileSyntax, in context: Context) throws -> SourceFileSyntax { + let sourceDecl = try findSourceDecl(syntax: syntax, range: context.range) + let selectedMembers = findSelectedMembers(declGroup: sourceDecl.declGroup, range: context.range) + let membersToMove = try validateMovableMembers(selectedMembers: selectedMembers) + let remainingMembers = updateRemainingMembers(declGroup: sourceDecl.declGroup, membersToMove: membersToMove) + + var updatedDeclGroup = sourceDecl.declGroup + updatedDeclGroup.memberBlock = updateMemberBlock( + sourceDecl.declGroup.memberBlock, + members: remainingMembers + ) + + let extensionDecl = makeExtension(sourceDecl: sourceDecl, membersToMove: membersToMove) + + return updateSyntax( + syntax, + sourceDecl: sourceDecl, + updatedDeclGroup: updatedDeclGroup, + extensionDecl: extensionDecl + ) + } +} + private enum ValidationResult: CustomStringConvertible { case accessor case deinitializer @@ -51,32 +84,93 @@ private enum ValidationResult: CustomStringConvertible { } } -package struct MoveMembersToExtension: SyntaxRefactoringProvider { - package struct Context { - let range: Range +private extension MoveMembersToExtension { + struct MemberToMove { + struct NestedMember: Equatable { + let updatedMember: MemberBlockItemSyntax + let partsToMove: [MemberBlockItemSyntax] + let validationResults: [ValidationResult] + let path: [String] + let indentationToRemove: Trivia + } - package init(range: Range) { - self.range = range + enum Scope: Equatable { + case inner + case nested(NestedMember) + } + + let member: MemberBlockItemSyntax + let scope: Scope + + var validationResults: [ValidationResult] { + switch scope { + case .inner: + return [ValidationResult(member)].compactMap { $0 } + + case .nested(let nested): + return nested.validationResults + } + } + + var hasMovableMembers: Bool { + switch scope { + case .inner: + return ValidationResult(member) == nil + + case .nested(let nested): + return !nested.partsToMove.isEmpty + } + } + + init( + member: MemberBlockItemSyntax, + scope: Scope = .inner + ) { + self.member = member + self.scope = scope } } - package static func refactor(syntax: SourceFileSyntax, in context: Context) throws -> SourceFileSyntax { + struct SourceDecl { + let statement: CodeBlockItemSyntax + let index: CodeBlockItemListSyntax.Index + let declGroup: any DeclGroupSyntax + let declName: TokenSyntax + } + + static func findSourceDecl(syntax: SourceFileSyntax, range: Range) throws -> SourceDecl { guard - let statement = syntax.statements.first(where: { $0.item.range.contains(context.range) }), + let statement = syntax.statements.first(where: { $0.item.range.contains(range) }), let decl = statement.item.asProtocol((any NamedDeclSyntax).self), let declGroup = statement.item.asProtocol((any DeclGroupSyntax).self), - let statementIndex = syntax.statements.index(of: statement) + let index = syntax.statements.index(of: statement) else { throw RefactoringNotApplicableError("Type declaration not found") } - let selectedMembers = Array(declGroup.memberBlock.members).filter { context.range.overlaps($0.trimmedRange) } - .map { (member: $0, validationResult: ValidationResult($0)) } + return SourceDecl(statement: statement, index: index, declGroup: declGroup, declName: decl.name) + } + + static func findSelectedMembers(declGroup: any DeclGroupSyntax, range: Range) -> [MemberToMove] { + declGroup.memberBlock.members.compactMap { member in + guard range.overlaps(member.trimmedRange) else { return nil } + + if let nestedMember = findNestedMovableMember(member: member, range: range) { + return MemberToMove( + member: member, + scope: .nested(nestedMember) + ) + } + + return MemberToMove(member: member) + } + } - var membersToMove = selectedMembers.filter({ $0.validationResult == nil }).map(\.member) + static func validateMovableMembers(selectedMembers: [MemberToMove]) throws -> [MemberToMove] { + let membersToMove = selectedMembers.filter(\.hasMovableMembers) guard !membersToMove.isEmpty else { - let notMovedMembers = Set(selectedMembers.compactMap(\.validationResult)) + let notMovedMembers = Set(selectedMembers.flatMap(\.validationResults)) .map(\.description) .sorted().joined(separator: ", ") throw RefactoringNotApplicableError( @@ -84,42 +178,241 @@ package struct MoveMembersToExtension: SyntaxRefactoringProvider { ) } - var updatedDeclGroup = declGroup - var remainingMembers = Array(declGroup.memberBlock.members).filter { !membersToMove.contains($0) } - membersToMove[0].decl.leadingTrivia = membersToMove[0].decl.leadingTrivia.trimmingPrefix(while: \.isSpaceOrTab) + return membersToMove + } - if remainingMembers.isEmpty { - updatedDeclGroup.memberBlock.rightBrace.leadingTrivia = Trivia() - } else { - remainingMembers[0].leadingTrivia = .newline.merging( - remainingMembers[0].leadingTrivia.trimmingPrefix(while: \.isNewline) + static func updateRemainingMembers( + declGroup: any DeclGroupSyntax, + membersToMove: [MemberToMove] + ) -> [MemberBlockItemSyntax] { + + let wholeMembers = membersToMove.compactMap { move in + if case .inner = move.scope { + return move.member + } + return nil + } + + var remainingMembers = Array(declGroup.memberBlock.members).filter { !wholeMembers.contains($0) } + + for index in remainingMembers.indices { + guard let move = membersToMove.first(where: { $0.member == remainingMembers[index] }), + case .nested(let nested) = move.scope + else { + continue + } + + remainingMembers[index] = nested.updatedMember + } + + return remainingMembers + } + + static func updateMemberBlock( + _ memberBlock: MemberBlockSyntax, + members: [MemberBlockItemSyntax] + ) -> MemberBlockSyntax { + var memberBlock = memberBlock + var members = members + + if members.isEmpty { + memberBlock.members = MemberBlockItemListSyntax() + memberBlock.rightBrace.leadingTrivia = Trivia() + return memberBlock + } + + members[0].leadingTrivia = .newline.merging( + members[0] + .leadingTrivia + .trimmingPrefix(while: \.isNewline) + ) + + let lastIndex = members.index(before: members.endIndex) + + members[lastIndex].trailingTrivia = + members[lastIndex] + .trailingTrivia + .trimmingSuffix(while: \.isNewline) + + memberBlock.members = MemberBlockItemListSyntax(members) + return memberBlock + } + + static func findNestedMovableMember( + member: MemberBlockItemSyntax, + range: Range + ) -> MemberToMove.NestedMember? { + guard let memberGroup = member.decl.asProtocol((any DeclGroupSyntax).self), + let memberName = member.decl.asProtocol((any NamedDeclSyntax).self)?.name.text + else { + return nil + } + + var currentGroup = memberGroup + var path = [memberName] + var declGroups = [(declGroup: any DeclGroupSyntax, index: Int)]() + + while true { + let members = Array(currentGroup.memberBlock.members) + + let selectedIndices = members.indices.filter { range.overlaps(members[$0].trimmedRange) } + + guard selectedIndices.count == 1, + let selectedIndex = selectedIndices.first, + let childGroup = members[selectedIndex].decl.asProtocol((any DeclGroupSyntax).self), + let childName = members[selectedIndex].decl.asProtocol((any NamedDeclSyntax).self)?.name.text, + childGroup.memberBlock.members.contains(where: { range.overlaps($0.trimmedRange) }) + else { + break + } + + declGroups.append((currentGroup, selectedIndex)) + + currentGroup = childGroup + path.append(childName) + } + + let selectedMembers = Array(currentGroup.memberBlock.members).filter { range.overlaps($0.trimmedRange) } + + guard !selectedMembers.isEmpty else { return nil } + + let partsToMove = selectedMembers.filter { ValidationResult($0) == nil } + + let validationResults = selectedMembers.compactMap(ValidationResult.init) + + let remainingMembers = selectedMembers.filter { !partsToMove.contains($0) } + + var updatedGroup = currentGroup + updatedGroup.memberBlock = updateMemberBlock(currentGroup.memberBlock, members: remainingMembers) + + for (declGroup, index) in declGroups.reversed() { + var parentMembers = Array(declGroup.memberBlock.members) + + parentMembers[index].decl = DeclSyntax(updatedGroup) + + var updatedParent = declGroup + updatedParent.memberBlock.members = MemberBlockItemListSyntax(parentMembers) + updatedGroup = updatedParent + } + + var updatedMember = member + updatedMember.decl = DeclSyntax(updatedGroup) + + let indentationToRemove = + currentGroup + .firstToken(viewMode: .sourceAccurate)? + .indentationOfLine + ?? Trivia() + + return MemberToMove.NestedMember( + updatedMember: updatedMember, + partsToMove: partsToMove, + validationResults: validationResults, + path: path, + indentationToRemove: indentationToRemove + ) + } + + static func unindentExtensionMembers( + _ members: [MemberBlockItemSyntax], + by indentation: Trivia + ) -> [MemberBlockItemSyntax] { + members.map { member in + let remover = IndentationRemover( + indentation: indentation, + indentFirstLine: true + ) + + return + remover + .rewrite(member) + .as(MemberBlockItemSyntax.self) + ?? member + } + } + + static func makeExtensionName( + rootName: TokenSyntax, + path: [String] + ) -> TypeSyntax { + var rootName = rootName + rootName.leadingTrivia = Trivia() + rootName.trailingTrivia = Trivia() + + var type = TypeSyntax( + IdentifierTypeSyntax( + leadingTrivia: .space, + name: rootName + ) + ) + + for component in path { + type = TypeSyntax( + MemberTypeSyntax( + baseType: type, + period: .periodToken(), + name: .identifier(component) + ) ) - remainingMembers[remainingMembers.count - 1].trailingTrivia = remainingMembers[remainingMembers.count - 1] - .trailingTrivia.trimmingSuffix(while: \.isNewline) } - updatedDeclGroup.memberBlock.members = MemberBlockItemListSyntax(remainingMembers) - membersToMove[0].leadingTrivia = .newline.merging(membersToMove[0].leadingTrivia.trimmingPrefix(while: \.isNewline)) - let extensionMemberBlockSyntax = declGroup.memberBlock.with(\.members, MemberBlockItemListSyntax(membersToMove)) + type.trailingTrivia = .space + return type + } - var declName = decl.name + static func makeExtension(sourceDecl: SourceDecl, membersToMove: [MemberToMove]) -> ExtensionDeclSyntax { + var extensionMembers = [MemberBlockItemSyntax]() + var declName = sourceDecl.declName declName.trailingTrivia = declName.trailingTrivia.merging(.space) - let extensionDecl = ExtensionDeclSyntax( + let extendedType: TypeSyntax + + if membersToMove.count == 1, + case .nested(let nested) = membersToMove[0].scope + { + extensionMembers = unindentExtensionMembers( + nested.partsToMove, + by: nested.indentationToRemove + ) + + extendedType = makeExtensionName( + rootName: sourceDecl.declName, + path: nested.path + ) + } else { + extensionMembers = membersToMove.map(\.member) + + extendedType = makeExtensionName( + rootName: sourceDecl.declName, + path: [] + ) + } + + extensionMembers[0].leadingTrivia = .newline.merging( + extensionMembers[0].leadingTrivia.trimmingPrefix(while: \.isNewline) + ) + + var memberBlock = sourceDecl.declGroup.memberBlock + memberBlock.members = MemberBlockItemListSyntax(extensionMembers) + + return ExtensionDeclSyntax( leadingTrivia: .newlines(2), - extendedType: IdentifierTypeSyntax( - leadingTrivia: .space, - name: declName - ), - memberBlock: extensionMemberBlockSyntax + extendedType: extendedType, + memberBlock: memberBlock ) + } + static func updateSyntax( + _ syntax: SourceFileSyntax, + sourceDecl: SourceDecl, + updatedDeclGroup: any DeclGroupSyntax, + extensionDecl: ExtensionDeclSyntax + ) -> SourceFileSyntax { var syntax = syntax - let updatedStatement = statement.with(\.item, .decl(DeclSyntax(updatedDeclGroup))) - syntax.statements[statementIndex] = updatedStatement + syntax.statements[sourceDecl.index] = sourceDecl.statement.with(\.item, .decl(DeclSyntax(updatedDeclGroup))) syntax.statements.insert( CodeBlockItemSyntax(item: .decl(DeclSyntax(extensionDecl))), - at: syntax.statements.index(after: statementIndex) + at: syntax.statements.index(after: sourceDecl.index) ) return syntax } diff --git a/Tests/SwiftSyntaxCodeActionsTests/MoveMembersToExtensionTests.swift b/Tests/SwiftSyntaxCodeActionsTests/MoveMembersToExtensionTests.swift index 4e73d5045..26195f506 100644 --- a/Tests/SwiftSyntaxCodeActionsTests/MoveMembersToExtensionTests.swift +++ b/Tests/SwiftSyntaxCodeActionsTests/MoveMembersToExtensionTests.swift @@ -10,14 +10,13 @@ // //===----------------------------------------------------------------------===// +import SKTestSupport import SwiftParser import SwiftRefactor import SwiftSyntax import SwiftSyntaxBuilder import SwiftSyntaxCodeActions import XCTest -import SKTestSupport - final class MoveMembersToExtensionTests: XCTestCase { func testMoveFunctionFromClass() throws { @@ -34,19 +33,19 @@ final class MoveMembersToExtensionTests: XCTestCase { } """, expected: - """ - class Foo { - func bar() { - print("Hello world!") + """ + class Foo { + func bar() { + print("Hello world!") + } } - } - extension Foo { - func foo() { - print("Hello world!") + extension Foo { + func foo() { + print("Hello world!") + } } - } - """ + """ ) } @@ -57,7 +56,7 @@ final class MoveMembersToExtensionTests: XCTestCase { func foo() { 1️⃣print("Hello world!") }2️⃣ - + func bar() { print("Hello world!") } @@ -68,23 +67,23 @@ final class MoveMembersToExtensionTests: XCTestCase { } """, expected: - """ - class Foo { - func bar() { - print("Hello world!") + """ + class Foo { + func bar() { + print("Hello world!") + } } - } - extension Foo { - func foo() { - print("Hello world!") + extension Foo { + func foo() { + print("Hello world!") + } } - } - struct Bar { - func foo() {} - } - """ + struct Bar { + func foo() {} + } + """ ) } @@ -97,7 +96,7 @@ final class MoveMembersToExtensionTests: XCTestCase { } deinit() {} - + func bar() { print("Hello world!") }2️⃣ @@ -108,25 +107,25 @@ final class MoveMembersToExtensionTests: XCTestCase { } """, expected: - """ - class Foo { - deinit() {} - } - - extension Foo { - func foo() { - print("Hello world!") + """ + class Foo { + deinit() {} } - func bar() { - print("Hello world!") + extension Foo { + func foo() { + print("Hello world!") + } + + func bar() { + print("Hello world!") + } } - } - struct Bar { - func foo() {} - } - """ + struct Bar { + func foo() {} + } + """ ) } @@ -140,15 +139,15 @@ final class MoveMembersToExtensionTests: XCTestCase { } """, expected: - """ - struct Outer {} - - extension Outer { - struct Inner { + """ + struct Outer { + struct Inner {} + } + + extension Outer.Inner { func moveThis() {} } - } - """ + """ ) } @@ -162,19 +161,19 @@ final class MoveMembersToExtensionTests: XCTestCase { } """, expected: - """ - struct Outer {} - - extension Outer { - struct Inner { + """ + struct Outer { + struct Inner {} + } + + extension Outer.Inner { func moveThis() {} } - } - """ + """ ) } - func testMoveSelectedFunctionName() throws { + func testMoveNestedFunctionNameFromGeneric() throws { try assertMoveMembersToExtension( """ struct Outer { @@ -184,20 +183,85 @@ final class MoveMembersToExtensionTests: XCTestCase { } """, expected: + """ + struct Outer { + struct Inner {} + } + + extension Outer.Inner { + func moveThis() {} + } + """ + ) + } + + func testMoveNestedFunctionName2() throws { + try assertMoveMembersToExtension( """ - struct Outer {} + struct Outer { + struct Middle { + struct Inner { + func 1️⃣moveThis()2️⃣ {} + } + } + } + """, + expected: + """ + struct Outer { + struct Middle { + struct Inner {} + } + } - extension Outer { - struct Inner { + extension Outer.Middle.Inner { func moveThis() {} } + """ + ) + } + + func testNestedStoredPropertyIsNotMoved() throws { + try assertMoveMembersToExtension( + """ + struct Outer { + struct Inner { + 1️⃣var value = 12️⃣ + } } + """, + expected: nil + ) + } + + func testNestedInvalidMemberRemains() throws { + try assertMoveMembersToExtension( """ + struct Outer { + struct Inner {1️⃣ + var value = 1 + + func moveThis() {}2️⃣ + } + } + """, + expected: + """ + struct Outer { + struct Inner { + var value = 1 + } + } + + extension Outer.Inner { + func moveThis() {} + } + """ ) } func testSelectedDeinitializerMember() async throws { - try assertMoveMembersToExtension ( + try assertMoveMembersToExtension( """ class Foo { func foo() {