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

如何用Python AST识别特定类的所有实例化以编写flake8插件

编写flake8插件识别Foo类的所有实例化(覆盖别名/间接引用场景)

要实现能全面识别Foo类所有实例化的flake8插件,单纯靠跟踪变量别名的简单AST遍历是不够的——像函数返回、字典取值、列表推导这类间接引用,没法靠静态的别名列表覆盖。核心思路是通过数据流分析+符号执行,跟踪每个变量/表达式最终指向的原始类,而不是只看表面的变量名。

核心思路:带作用域的符号跟踪

给AST访问器维护一个作用域环境(用字典栈实现,每个栈帧对应全局、函数局部等不同作用域),每个变量名对应一个"符号"——即它最终指向的原始类(比如Foo)。对于复杂表达式,要递归解析其背后的符号。

关键节点的处理逻辑

1. 导入节点(Import/ImportFrom)

处理导入时直接记录别名与原始类的映射:

  • import Foo:全局作用域中"Foo" → "Foo"
  • import Foo as Bar:全局作用域中"Bar" → "Foo"
  • from module import Foo as Baz:全局作用域中"Baz" → "Foo"

2. 赋值节点(Assign/AnnAssign)

解析赋值右侧表达式的符号(比如右侧是get_Foo(),需解析该调用返回的符号为Foo),再把左侧变量关联到该符号。

3. 函数定义与作用域管理

进入函数定义时创建新的局部作用域栈帧;函数处理完成后弹出栈帧,避免变量污染上层作用域。

4. 调用节点(Call)

解析调用的func部分对应的符号,如果符号是"Foo",则判定为Foo的实例化。

5. 复杂表达式解析

对于字典取值、列表索引、属性访问等,递归解析基础表达式的符号,再判断取值后的符号是否指向Foo。比如get_Foo()["foo"],先解析get_Foo()的符号是返回Foo的函数,再确定其返回值的符号为Foo,最终["foo"]对应的符号也是Foo。

示例实现代码

下面是一个简化版的带作用域符号跟踪的Visitor:

import ast
from typing import Dict, List, Optional

class FooTrackingVisitor(ast.NodeVisitor):
    def __init__(self, target_class: str = "Foo"):
        self.target_class = target_class
        # 作用域栈:每个元素是当前作用域的变量→符号映射
        self.scope_stack: List[Dict[str, str]] = [{}]
        # 记录已知返回目标类的函数名
        self.return_target_funcs: List[str] = []

    @property
    def current_scope(self) -> Dict[str, str]:
        return self.scope_stack[-1]

    def push_scope(self):
        self.scope_stack.append({})

    def pop_scope(self):
        self.scope_stack.pop()

    def resolve_symbol(self, node) -> Optional[str]:
        """递归解析节点对应的原始符号"""
        if isinstance(node, ast.Name):
            # 从当前作用域向上查找符号
            for scope in reversed(self.scope_stack):
                if node.id in scope:
                    return scope[node.id]
            return None
        elif isinstance(node, ast.Call):
            # 解析函数调用的返回符号:如果函数是已知返回目标类的,直接返回目标类
            func_sym = self.resolve_symbol(node.func)
            if func_sym in self.return_target_funcs:
                return self.target_class
            return None
        elif isinstance(node, ast.Subscript):
            # 处理字典/列表取值:如果基础对象关联目标类,返回目标类符号
            base_sym = self.resolve_symbol(node.value)
            if base_sym == self.target_class:
                return self.target_class
            return None
        return None

    def visit_Import(self, node: ast.Import):
        for alias in node.names:
            if alias.name == self.target_class:
                var_name = alias.asname or alias.name
                self.current_scope[var_name] = self.target_class
        self.generic_visit(node)

    def visit_ImportFrom(self, node: ast.ImportFrom):
        for alias in node.names:
            if alias.name == self.target_class:
                var_name = alias.asname or alias.name
                self.current_scope[var_name] = self.target_class
        self.generic_visit(node)

    def visit_FunctionDef(self, node: ast.FunctionDef):
        # 检查函数体是否直接返回目标类,标记这类函数
        for stmt in node.body:
            if isinstance(stmt, ast.Return) and isinstance(stmt.value, ast.Name):
                if stmt.value.id == self.target_class:
                    self.return_target_funcs.append(node.name)
                    break
        # 进入函数作用域
        self.push_scope()
        self.generic_visit(node)
        self.pop_scope()

    def visit_Assign(self, node: ast.Assign):
        # 解析右侧表达式的符号,关联到左侧变量
        rhs_sym = self.resolve_symbol(node.value)
        if rhs_sym:
            for target in node.targets:
                if isinstance(target, ast.Name):
                    self.current_scope[target.id] = rhs_sym
        self.generic_visit(node)

    def visit_Call(self, node: ast.Call):
        # 检查调用的函数是否指向目标类
        func_sym = self.resolve_symbol(node.func)
        if func_sym == self.target_class:
            print(f"Foo instantiated at line {node.lineno}")
        self.generic_visit(node)

# 测试代码
test_code = """
import Foo

def get_Foo():
    return Foo

Bar = get_Foo()
f = Bar()

def get_Foo_dict():
    return {"foo": Foo}

Baz = get_Foo_dict()["foo"]
g = Baz()

[f() for f in get_Foo()]
"""

tree = ast.parse(test_code)
visitor = FooTrackingVisitor()
visitor.visit(tree)

局限性说明

静态分析无法覆盖所有极端场景:

  • 动态生成的代码(比如eval("Foo")、globals()[var_name]())
  • 运行时才确定的引用(比如从外部文件读取类名)
  • 复杂条件分支返回的动态类

但上述实现已经能覆盖绝大多数常见的"规避"场景,满足flake8插件的日常检查需求。

内容的提问来源于stack exchange,提问作者Dave McLean

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.15 14:01:05