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

如何修复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

关键修复点说明

  1. 类映射构建:build_class_map遍历所有文件,将每个类的「模块名+类名」作为唯一键,关联对应的ClassDef节点,支持跨模块继承查找。
  2. 基类标识解析:_get_class_identifier处理两种常见基类节点:
    • ast.Name:同一模块内的类,使用当前模块名+类名作为键
    • ast.Attribute:跨模块类(如foo.Bar),拆分出模块名和类名
  3. 缓存机制:使用lru_cache缓存已计算的类层级,避免重复递归计算,提升效率。
  4. 多继承支持:遍历所有基类,取最大层级加1,符合多继承下取最深继承链的需求。

测试验证

对于示例中的继承链C → B → A:

  • C的层级为0(无基类)
  • B的层级为C的层级+1=1
  • A的层级为B的层级+1=2
    完全符合预期结果。

内容的提问来源于stack exchange,提问作者Giammaria GIORDANO

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.24 10:55:11