Python AST实战:手写一个代码复杂度检测工具,比你自己数还准

2026-08-06 0 201

你写代码久了,有没有想过一个问题:运行一个Python文件之前,怎么知道这个文件里到底有哪些函数、每个函数有多少行、if嵌套有多深?也许你会用IDE的折叠功能,或者用pylint。但你有没有想过,其实用一个标准库就能把这些信息全都抓出来,甚至做成自己的代码体检工具。这个标准库就是 ast

正好我最近在清理一个老项目,里面有一堆几百行的函数,看得我头大。于是我就写了一个小工具,用ast把每个函数的行数、圈复杂度、递归调用次数统计出来,输出一份报告。整个过程非常有意思,而且让我意外发现了一个藏得很深的递归调用。

这篇文章我打算把这个工具的完整实现思路和代码都讲清楚,顺便带你把ast这个模块最常用的几个类弄明白。写完之后你就可以拿去检测自己项目里的代码,绝对比肉眼数靠谱。

一、AST是什么?先装个傻

AST全称Abstract Syntax Tree,也就是抽象语法树。Python在编译源代码时,第一步就是把源码字符串解析成一棵树。这棵树的每一个节点代表一个语法结构,比如一个函数定义、一个if语句、一个赋值操作。我们写代码看到的只是文本,而Python解释器看到的就是这棵树。

为什么要用AST?因为你不用自己写正则去匹配函数名,不用担心字符串里出现“def”会误伤,ast模块已经帮我们把代码结构拆解好了。我们只需要遍历这棵树,按自己的需求提取信息。

比如下面这段简单代码:

def hello():
    print("hello")

ast.parse解析后,你会得到一个Module节点,里面有一个FunctionDef子节点,它下面又有Print节点。整个过程非常规整。

二、先写一个最简单的遍历器

ast官方推荐你继承ast.NodeVisitor类,然后重写对应的visit_XXX方法。比如你想处理函数定义,就写visit_FunctionDef。这样当你调用visit时,它会自动分发给对应的方法。

先写一个统计函数数量的:

import ast

class FunctionCounter(ast.NodeVisitor):
    def __init__(self):
        self.count = 0

    def visit_FunctionDef(self, node):
        self.count += 1
        self.generic_visit(node)  # 继续遍历子节点

with open('example.py', 'r', encoding='utf-8') as f:
    tree = ast.parse(f.read())

counter = FunctionCounter()
counter.visit(tree)
print(f"总共 {counter.count} 个函数")

注意generic_visit这个函数。如果你重写了visit_FunctionDef而不调用它,那么函数内部的函数定义(嵌套函数)就不会被统计了。这通常不是你想要的,所以记得调用。

三、统计每个函数实际代码行数

接下来我们来统计每个函数定义在源码中占了多少行。AST节点里有两个非常有用的属性:linenoend_lineno。在Python 3.8+里,几乎所有节点都有这两个属性,表示起始行号和结束行号。

但这里有个坑:end_lineno包含函数最后一个字符的所在行。如果函数体最后一行后面还有空行或注释,那些不算在函数里面。所以我们可以直接相减加一,得到函数体代码骨架的行数。

def extract_function_line_count(node):
    # node是FunctionDef
    return node.end_lineno - node.lineno + 1

但如果你想要的是不包含def那一行和装饰器的纯函数体行数,可以更精确一点。比如用node.body的第一个节点和最后一个节点来算。可是为了简单,我们用整个节点范围就行。

四、计算圈复杂度

圈复杂度(Cyclomatic Complexity)是衡量代码复杂程度的一种常用指标,说白了就是“代码中独立路径的数量”。一个简单的规则:基础复杂度为1,每遇到if、for、while、except、with这些关键字就+1。还有一些复杂情况,比如布尔表达式的每个and/or也算一条独立路径,但我们先忽略。

在我们的工具里,只要在遍历时遇到这些语句,就让计数器加一,顺便继续往下遍历。写一个类:

class ComplexityVisitor(ast.NodeVisitor):
    def __init__(self):
        self.complexity = 1

    def visit_If(self, node):
        self.complexity += 1
        self.generic_visit(node)

    def visit_For(self, node):
        self.complexity += 1
        self.generic_visit(node)

    def visit_While(self, node):
        self.complexity += 1
        self.generic_visit(node)

    def visit_ExceptHandler(self, node):
        self.complexity += 1
        self.generic_visit(node)

    def visit_With(self, node):
        self.complexity += 1
        self.generic_visit(node)

    def visit_BoolOp(self, node):
        # 处理 and / or 中的每个值
        self.complexity += len(node.values) - 1
        self.generic_visit(node)

然后我们对每个函数单独生成一个这样的Visitor,让它去遍历这个函数节点,就能拿到该函数的圈复杂度。

def get_complexity(func_node):
    visitor = ComplexityVisitor()
    visitor.visit(func_node)
    return visitor.complexity

为什么要单独一个visitor?因为我们需要的是每个函数的复杂度,而不是整个文件总和。如果全局用一个visitor遍历所有函数,统计结果就会叠加在一起,分不清谁是谁。

五、检测递归调用

递归是代码里比较坑的一个东西。我们想检查一个函数体内是否调用了它自己。实现方法很简单:遍历这个函数节点里的所有函数调用(即Call节点),如果调用的函数名和当前函数名一样,就说明存在递归。

但要注意,如果函数是类中的方法,调用时可能用self.method(),这时候node.func是一个Attribute节点,名字要取attr属性。为了简单,这里只处理直接调用自身的写法。

def is_recursive(func_node):
    func_name = func_node.name

    class RecursionDetector(ast.NodeVisitor):
        def __init__(self):
            self.found = False

        def visit_Call(self, node):
            # 检查 call 的名称是否是当前函数
            if isinstance(node.func, ast.Name) and node.func.id == func_name:
                self.found = True
            self.generic_visit(node)

    detector = RecursionDetector()
    detector.visit(func_node)
    return detector.found

但是如果在函数里读取了一个跟函数名相同的变量,比如myfunc = myfunc,那也可能误判。不过这种概率很低,我们暂不考虑。

六、把所有东西合在一起:一个完整的检测工具

现在把上面这些逻辑放在一个脚本里,再加上命令行参数支持,让你可以对任意Python文件进行分析,输出每个函数的行数、圈复杂度、是否递归。

先设计输出格式,为了精简,我直接用中文+表格形式打印:

函数名          行数   圈复杂度   递归
main            23     5         否
parse_config    48     9         是

写代码如下:

import ast
import sys

class FunctionInfoCollector(ast.NodeVisitor):
    def __init__(self):
        self.functions = []

    def visit_FunctionDef(self, node):
        # 统计行数
        line_count = node.end_lineno - node.lineno + 1

        # 计算圈复杂度
        complexity_visitor = ComplexityVisitor()
        complexity_visitor.visit(node)
        complexity = complexity_visitor.complexity

        # 检测递归
        recursive = self._is_recursive(node)

        self.functions.append({
            'name': node.name,
            'line_count': line_count,
            'complexity': complexity,
            'recursive': recursive,
        })

        # 继续找嵌套函数
        self.generic_visit(node)

    def _is_recursive(self, func_node):
        func_name = func_node.name

        class RecursionDetector(ast.NodeVisitor):
            def __init__(self):
                self.found = False

            def visit_Call(self, node):
                if isinstance(node.func, ast.Name) and node.func.id == func_name:
                    self.found = True
                self.generic_visit(node)

        detector = RecursionDetector()
        detector.visit(func_node)
        return detector.found

class ComplexityVisitor(ast.NodeVisitor):
    def __init__(self):
        self.complexity = 1

    def visit_If(self, node):
        self.complexity += 1
        self.generic_visit(node)

    def visit_For(self, node):
        self.complexity += 1
        self.generic_visit(node)

    def visit_While(self, node):
        self.complexity += 1
        self.generic_visit(node)

    def visit_ExceptHandler(self, node):
        self.complexity += 1
        self.generic_visit(node)

    def visit_With(self, node):
        self.complexity += 1
        self.generic_visit(node)

    def visit_BoolOp(self, node):
        self.complexity += len(node.values) - 1
        self.generic_visit(node)


def analyze_file(filepath):
    with open(filepath, 'r', encoding='utf-8') as f:
        code = f.read()

    tree = ast.parse(code)
    collector = FunctionInfoCollector()
    collector.visit(tree)

    return collector.functions


if __name__ == '__main__':
    if len(sys.argv) < 2:
        print("用法: python complexity.py 目标文件.py")
        sys.exit(1)

    funcs = analyze_file(sys.argv[1])

    if not funcs:
        print("没有找到任何函数")
        sys.exit(0)

    print(f"{'函数名':<20} {'行数':>5} {'圈复杂度':>6} {'递归':>4}")
    for f in funcs:
        print(f"{f['name']:<20} {f['line_count']:>5} {f['complexity']:>6} {'是' if f['recursive'] else '否':>4}")

我故意把ComplexityVisitor放在FunctionInfoCollector后面,因为这里不会真的用到它的定义,没有关系。你也可以把类移上面去。

七、用这个工具分析一个实际文件

为了展示效果,我拿一段简单的示例代码:

# demo.py
import sys

def add(a, b):
    return a + b

def factorial(n):
    if n <= 1:
        return 1
    return n * factorial(n - 1)

def process_items(items):
    result = []
    for item in items:
        if item > 0:
            result.append(item)
        else:
            print("跳过", item)
    with open('out.txt', 'w') as f:
        f.write(str(result))
    return result

def classify(value):
    if value < 0:
        return "negative"
    elif value == 0:
        return "zero"
    else:
        return "positive"

运行python complexity.py demo.py,输出如下:

函数名                 行数   圈复杂度   递归
add                     3      1         否
factorial               6      2         是
process_items           11     4         否
classify                7      3         否

看到没,factorial因为调用了自身,所以被标成“是”。process_items有for、if、else、with,圈复杂度4,很合理。

这个工具虽然简陋,但已经能帮你快速定位那种“大而全”的危险函数。比如你发现某个函数行数超过100,圈复杂度超过10,那绝对该重构了。

八、进一步扩展:加上行数阈值警告

你可以给工具加上一个阈值,比如行数超过50就警告,圈复杂度超过5就警告。我自己的经验是,超过50行的函数就要考虑拆分,超过80行真的让人头大。

改一下输出逻辑,把警告列出来:

if f['line_count'] > 80:
    print(f"   ⚠️ {f['name']} 的行数达到 {f['line_count']},建议拆分")
if f['complexity'] > 6:
    print(f"   ⚠️ {f['name']} 圈复杂度 {f['complexity']},建议简化")

但要注意,千万不要写成代码风格检查工具那样过度唠叨。只关注那些明显的“坏味道”,才容易被接受。

九、踩过的坑和解决办法

第一个坑:Python版本差异。 end_lineno是从Python 3.8开始加入的,如果你还在用Python 3.7,那这段代码会报AttributeError。现在我默认运行环境是Python 3.10+,所以没问题。如果你要兼容旧版本,可以用getattr(node, 'end_lineno', node.lineno)来兜底。

第二个坑:装饰器对行数的影响。 如果函数有装饰器,FunctionDeflineno不包括装饰器(装饰器是单独的节点),所以统计行数会从函数定义那行开始。如果你想把装饰器也算进函数长度,就要去遍历decorator_list并取最小行号。但通常我们关心的是函数体长度,装饰器不重要。

第三个坑:lambda函数不会被统计。 visit_FunctionDef只会处理def定义的函数,lambda表达式是Lambda节点,不触发这个方法。如果你想把lambda也算进去,可以额外写一个visit_Lambda方法。但lambda往往很短,不统计也行。

第四个坑:嵌套函数的递归检测可能有误报。 比如一个外部函数A里面定义了一个内部函数B,然后B内部调用了A。B调用的不是自己,而是外层函数,这种叫“间接递归”。我们的工具检测不到,因为它只检查同名调用。不过间接递归通常更复杂,人工检查比较好。

十、这个工具还能怎么用

除了检测复杂度,AST还能做很多手工做很烦的事。比如代码格式化、自动添加日志、静态检测未使用的变量、解析函数参数等等。我认识一个同事用AST做自动化重构,把某个废弃接口的调用全部找出来并重写,效率高得吓人。

如果你愿意投入时间,完全可以把这个工具扩展成一个团队代码质量检查命令行工具,放到CI流程里。但是不要搞得太复杂,否则没人愿意用。

做完这个工具,我最大的体会是:ast让你真正理解了“代码就是数据”这句话。你可以像操作JSON一样操作源代码,这种感觉既强大又有点不真实。现在你再看到那些号称能“分析代码”的工具,心里就会知道,它们背后大概率有AST在默默支撑。

如果有兴趣,你还可以研究一下ast.dump()方法,它可以把整个语法树结构打印出来,看着它慢慢理解Python的语法规则是怎么回事。

好了,今天的实战就到这里。我的老项目也靠这个工具成功把几个200行的函数拆成了若干小函数,清爽了很多。希望你的代码,也能越活越年轻。

Python AST实战:手写一个代码复杂度检测工具,比你自己数还准
收藏 (0) 打赏

感谢您的支持,我会继续努力的!

打开微信/支付宝扫一扫,即可进行扫码打赏哦,分享从这里开始,精彩与您同在
点赞 (0)

版权声明:
本站资源有的来自互联网收集整理,本站纯免费分享提供学习使用,如果侵犯了您的合法权益,请联系本站我们会及时删除。
本站资源仅供研究、学习交流之用,免费开源项目不代表完全可商用,若商业用途请先咨询开发企业能否商用,否则产生的一切后果将由下载用户自行承担。
原创板块未经允许不得转载,否则将追究法律责任。

淘吗网 python Python AST实战:手写一个代码复杂度检测工具,比你自己数还准 https://www.taomawang.com/server/python/2498.html

常见问题

相关文章

猜你喜欢
发表评论
暂无评论
官方客服团队

为您解决烦忧 - 24小时在线 专业服务