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

如何用NodeTransformer移除Python AST中函数定义的指定默认参数?

移除函数定义参数中的特定默认赋值

问题背景

给定代码的AST片段,需从函数定义的参数中移除所有包含在列表vars_to_remove中的变量对应的默认赋值。

举个例子:

  • 待移除变量列表:vars_to_remove = ['sum1']
  • 原函数定义:def do_smth(sum = sum1):
  • 修改后目标:def do_smth(sum):

已尝试的无效方法

  • 父节点遍历法:重写visit_FunctionDef(self, node)方法,虽能定位到目标参数,但返回None会删除整个FunctionDef节点,无法仅移除指定默认值。
  • 子节点遍历法:重写visit_Name(self, node)方法,返回None可删除节点,但会全局匹配所有Name节点,误删代码中其他位置的id:'sum1'节点。

可行解决方案

核心思路是在visit_FunctionDef方法内直接处理当前函数的参数与默认值绑定关系,精准过滤目标默认值,无需递归删除子节点或删除整个函数节点。

实现逻辑

  1. 参数与默认值的对应关系:函数的node.args.args是参数列表,node.args.defaults是默认值列表,两者从后往前一一对应(比如3个参数+2个默认值,对应最后2个参数拥有默认值)。
  2. 筛选保留项:遍历参数列表,判断每个带默认值的参数是否需要移除默认值——若默认值是vars_to_remove中的变量,则只保留参数,丢弃默认值;否则同时保留参数和默认值。
  3. 更新函数节点:将筛选后的参数和默认值列表重新赋值给node.args.args和node.args.defaults,返回修改后的函数节点。

示例代码

import ast

class RemoveDefaultAssignments(ast.NodeTransformer):
    def __init__(self, vars_to_remove):
        self.vars_to_remove = vars_to_remove

    def visit_FunctionDef(self, node):
        # 先递归处理函数内部节点
        self.generic_visit(node)
        
        args = node.args.args
        defaults = node.args.defaults
        if not defaults:
            return node
        
        # 计算带默认值的参数起始索引
        start_idx = len(args) - len(defaults)
        new_args = []
        new_defaults = []
        
        for idx, arg in enumerate(args):
            if idx >= start_idx:
                # 取出当前参数对应的默认值
                default = defaults[idx - start_idx]
                # 检查默认值是否为待移除的变量
                if isinstance(default, ast.Name) and default.id in self.vars_to_remove:
                    new_args.append(arg)
                else:
                    new_args.append(arg)
                    new_defaults.append(default)
            else:
                # 无默认值的参数直接保留
                new_args.append(arg)
        
        # 更新函数的参数和默认值
        node.args.args = new_args
        node.args.defaults = new_defaults
        return node

# 测试用例
test_code = """
sum1 = 100
def do_smth(sum = sum1, x=5):
    print(sum, x)
def another_func(a=sum1, b=20):
    print(a, b)
"""

# 解析、转换、输出代码
tree = ast.parse(test_code)
transformer = RemoveDefaultAssignments(vars_to_remove=['sum1'])
modified_tree = transformer.visit(tree)
print(ast.unparse(modified_tree))

输出结果

sum1 = 100
def do_smth(sum, x=5):
    print(sum, x)
def another_func(a, b=20):
    print(a, b)

内容的提问来源于stack exchange,提问作者JJ Kam

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.08 02:25:25