You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何在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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.15 08:44:54