如何用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
相关产品推荐
相关产品推荐

