如何在Swift Macro中获取Protocol对应的SimpleTypeIdentifierSyntax的声明方法?
解决Swift Macro获取协议方法并生成实现的问题
要搞定获取协议方法并生成实现的需求,核心是借助MacroExpansionContext的类型查询能力,拿到协议的完整元数据,再从中提取方法声明并生成默认实现。以下是具体实现步骤和代码修改:
1. 获取协议的类型符号
首先要从宏参数中解析出协议的类型,然后通过上下文的typeLookup方法获取协议的类型符号,确认它是协议类型:
// 在expansion方法中添加类型查询逻辑 let protocolTypeName = protocolType.identifier.text guard let protocolTypeSyntax = TypeSyntax(stringLiteral: protocolTypeName), let protocolSymbol = context.typeLookup(.type(protocolTypeSyntax)) as? ProtocolTypeSymbol else { throw MyError.invalidProtocolType }
2. 提取协议的方法成员
遍历协议符号的成员,筛选出所有实例方法(可根据需求调整是否保留static/class方法):
// 筛选协议中的实例方法 let protocolMethods = protocolSymbol.members.compactMap { member -> FunctionDeclSyntax? in guard let funcDecl = member.declaration.as(FunctionDeclSyntax.self) else { return nil } // 排除static/class方法 guard funcDecl.modifiers?.contains(where: { $0.name.text == "static" || $0.name.text == "class" }) == false else { return nil } return funcDecl }
3. 生成方法的默认实现
为每个方法生成符合语法的默认实现,根据返回类型适配不同的默认值:
// 生成方法的默认实现 private static func generateDefaultImplementation(for funcDecl: FunctionDeclSyntax) throws -> FunctionDeclSyntax { var implementedFunc = funcDecl let returnType = funcDecl.signature.returnClause?.type let bodyStatements: CodeBlockItemListSyntax if let returnType = returnType { let defaultExpr: ExprSyntax switch returnType.as(SimpleTypeIdentifierSyntax.self)?.name.text { case "String": defaultExpr = ExprSyntax(StringLiteralExprSyntax(content: "")) case "Int", "Double", "Float": defaultExpr = ExprSyntax(IntegerLiteralExprSyntax(digits: "0")) case "Bool": defaultExpr = ExprSyntax(BooleanLiteralExprSyntax(value: false)) default: // 自定义非可选类型需单独处理,这里仅支持可选类型返回nil guard returnType.as(OptionalTypeSyntax.self) != nil else { throw MyError.unsupportedReturnType } defaultExpr = ExprSyntax(NilLiteralExprSyntax()) } bodyStatements = CodeBlockItemListSyntax { CodeBlockItemSyntax(item: .stmt(StmtSyntax(ReturnStmtSyntax(expression: defaultExpr)))) } } else { // 无返回值(Void),空实现 bodyStatements = CodeBlockItemListSyntax() } implementedFunc.body = CodeBlockSyntax(leftBrace: .leftBraceToken(), statements: bodyStatements, rightBrace: .rightBraceToken()) return implementedFunc }
4. 整合到类声明中
把生成的方法添加到类的成员集合里:
// 生成所有方法的实现 let methodDecls = try protocolMethods.map { try generateDefaultImplementation(for: $0) } // 创建类声明并注入方法 let mockDecl = ClassDeclSyntax(identifier: "MyGeneratedClass", inheritanceClause: inheritance) { for method in methodDecls { DeclSyntax(method) } }
完整修改后的宏实现
public enum MyError: Error { case missingArgument case invalidProtocolType case unsupportedReturnType } public struct MyMacro: DeclarationMacro { public static func expansion(of node: some SwiftSyntax.FreestandingMacroExpansionSyntax, in context: some SwiftSyntaxMacros.MacroExpansionContext) throws -> [SwiftSyntax.DeclSyntax] { let protocolType = try node.extractArgument() let inheritanceType = SimpleTypeIdentifierSyntax(name: protocolType.identifier) let inheritance = TypeInheritanceClauseSyntax { InheritedTypeListSyntax { InheritedTypeSyntax(typeName: inheritanceType) } } // 1. 获取协议类型符号 let protocolTypeName = protocolType.identifier.text guard let protocolTypeSyntax = TypeSyntax(stringLiteral: protocolTypeName), let protocolSymbol = context.typeLookup(.type(protocolTypeSyntax)) as? ProtocolTypeSymbol else { throw MyError.invalidProtocolType } // 2. 筛选协议中的实例方法 let protocolMethods = protocolSymbol.members.compactMap { member -> FunctionDeclSyntax? in guard let funcDecl = member.declaration.as(FunctionDeclSyntax.self) else { return nil } guard funcDecl.modifiers?.contains(where: { $0.name.text == "static" || $0.name.text == "class" }) == false else { return nil } return funcDecl } // 3. 生成方法实现并添加到类中 let methodDecls = try protocolMethods.map { try generateDefaultImplementation(for: $0) } let mockDecl = ClassDeclSyntax(identifier: "MyGeneratedClass", inheritanceClause: inheritance) { for method in methodDecls { DeclSyntax(method) } } return [DeclSyntax(mockDecl)] } private static func generateDefaultImplementation(for funcDecl: FunctionDeclSyntax) throws -> FunctionDeclSyntax { var implementedFunc = funcDecl let returnType = funcDecl.signature.returnClause?.type let bodyStatements: CodeBlockItemListSyntax if let returnType = returnType { let defaultExpr: ExprSyntax switch returnType.as(SimpleTypeIdentifierSyntax.self)?.name.text { case "String": defaultExpr = ExprSyntax(StringLiteralExprSyntax(content: "")) case "Int", "Double", "Float": defaultExpr = ExprSyntax(IntegerLiteralExprSyntax(digits: "0")) case "Bool": defaultExpr = ExprSyntax(BooleanLiteralExprSyntax(value: false)) default: guard returnType.as(OptionalTypeSyntax.self) != nil else { throw MyError.unsupportedReturnType } defaultExpr = ExprSyntax(NilLiteralExprSyntax()) } bodyStatements = CodeBlockItemListSyntax { CodeBlockItemSyntax(item: .stmt(StmtSyntax(ReturnStmtSyntax(expression: defaultExpr)))) } } else { bodyStatements = CodeBlockItemListSyntax() } implementedFunc.body = CodeBlockSyntax(leftBrace: .leftBraceToken(), statements: bodyStatements, rightBrace: .rightBraceToken()) return implementedFunc } } private extension FreestandingMacroExpansionSyntax { func extractArgument() throws -> IdentifierExprSyntax { guard self.argumentList.count == 1, let memberAccessExpr = self.argumentList.first?.expression.as(MemberAccessExprSyntax.self), let identifierSyntax = memberAccessExpr.base?.as(IdentifierExprSyntax.self) else { throw MyError.missingArgument } return identifierSyntax } }
效果验证
使用你提供的示例代码:
protocol MyProtocol { func fooBar() -> String } #myMacro(MyProtocol.self) // 展开后会生成: // class MyGeneratedClass: MyProtocol { // func fooBar() -> String { // return "" // } // }
注意事项
- 上述代码只处理了基本类型和可选类型的默认返回值,对于自定义非可选类型,你可以根据需求调整错误处理逻辑或生成其他默认实现
- 如果协议包含属性、下标等其他成员,你可以用类似的方式扩展处理逻辑
内容的提问来源于stack exchange,提问作者Steve
相关产品推荐
相关产品推荐

