如何修复Python AST计算类最大继承层级的错误?
修复Python AST类继承层级计算问题
问题根源
原代码存在两个核心问题:
- 错误处理基节点类型:
class_def_node.bases中的元素是ast.Name或ast.Attribute节点(表示基类的名称/路径),而非ast.ClassDef节点。直接递归这些节点会触发AttributeError返回0,导致间接继承的层级无法累积。 - 缺少类节点映射:没有提前收集项目中所有类的
ClassDef节点,无法通过基类名称找到对应的定义,也就无法计算间接继承的层级。
修复方案
步骤1:预收集项目中所有类的节点映射
先遍历所有Python文件,建立「类标识(模块+类名)→ ClassDef节点」的映射,避免同名类冲突,同时支持跨模块继承的查找。
步骤2:修改层级计算函数
- 解析基类的完整标识(处理直接类名和跨模块类名)
- 通过映射找到基类的
ClassDef节点 - 加入缓存机制,避免重复计算同一类的层级
- 处理内置类(如
object)的情况,其层级为0
完整修复代码
import ast import os import pandas as pd from functools import lru_cache def _get_class_identifier(base_node, module_name): """解析基类节点的完整标识,返回(模块名, 类名)""" if isinstance(base_node, ast.Name): # 同一模块的类,模块名使用当前文件的模块名 return (module_name, base_node.id) elif isinstance(base_node, ast.Attribute): # 跨模块的类,如module.ClassName parts = [] node = base_node while isinstance(node, ast.Attribute): parts.append(node.attr) node = node.value if isinstance(node, ast.Name): parts.append(node.id) full_name = '.'.join(reversed(parts)) # 拆分模块名和类名,比如"module.submodule.Class" → ("module.submodule", "Class") module_part, class_part = full_name.rsplit('.', 1) return (module_part, class_part) # 无法解析的基类(如表达式),视为内置类 return (None, base_node.id if isinstance(base_node, ast.Name) else str(base_node)) def build_class_map(project_path): """构建项目中所有类的(模块名, 类名) → ClassDef节点的映射""" class_map = {} for root, _, files in os.walk(project_path): for file in files: if file.endswith(".py"): file_path = os.path.join(root, file) # 计算模块名:相对于项目根目录的路径替换为点分隔,去掉.py后缀 relative_path = os.path.relpath(file_path, project_path) module_name = os.path.splitext(relative_path)[0].replace(os.sep, '.') with open(file_path, "r", encoding="utf-8") as f: try: tree = ast.parse(f.read()) except: continue for node in ast.walk(tree): if isinstance(node, ast.ClassDef): class_key = (module_name, node.name) class_map[class_key] = node return class_map @lru_cache(maxsize=None) def calculate_inheritance(class_key, class_map): """计算指定类的最大继承层级,使用缓存避免重复计算""" class_node = class_map.get(class_key) if not class_node: # 内置类或未找到的类,层级为0 return 0 bases = class_node.bases if not bases: return 0 module_name = class_key[0] max_base_level = 0 for base in bases: base_key = _get_class_identifier(base, module_name) base_level = calculate_inheritance(base_key, class_map) if base_level > max_base_level: max_base_level = base_level return max_base_level + 1 def createAST(project_path): project_name = os.path.basename(project_path) class_map = build_class_map(project_path) data = [] for root, _, files in os.walk(project_path): for file in files: if file.endswith(".py"): file_path = os.path.join(root, file) relative_path = os.path.relpath(file_path, project_path) module_name = os.path.splitext(relative_path)[0].replace(os.sep, '.') with open(file_path, "r", encoding="utf-8") as f: try: tree = ast.parse(f.read()) except: continue for node in tree.body: if isinstance(node, ast.ClassDef): class_name = node.name class_key = (module_name, class_name) inheritance_level = calculate_inheritance(class_key, class_map) data.append((project_name, file_path, class_name, inheritance_level)) df = pd.DataFrame(data, columns=["Project", "File Path", "Class Name", "Inheritance Level"]) return df
关键修复点说明
- 类映射构建:
build_class_map遍历所有文件,将每个类的「模块名+类名」作为唯一键,关联对应的ClassDef节点,支持跨模块继承查找。 - 基类标识解析:
_get_class_identifier处理两种常见基类节点:ast.Name:同一模块内的类,使用当前模块名+类名作为键ast.Attribute:跨模块类(如foo.Bar),拆分出模块名和类名
- 缓存机制:使用
lru_cache缓存已计算的类层级,避免重复递归计算,提升效率。 - 多继承支持:遍历所有基类,取最大层级加1,符合多继承下取最深继承链的需求。
测试验证
对于示例中的继承链C → B → A:
- C的层级为0(无基类)
- B的层级为
C的层级+1=1 - A的层级为
B的层级+1=2
完全符合预期结果。
内容的提问来源于stack exchange,提问作者Giammaria GIORDANO
相关产品推荐
相关产品推荐

