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

如何合规提取mypy类型文件中TYPE_CHECKING条件下的导入模块名?

更规范地提取TYPE_CHECKING块中的导入模块名

你的思路方向是对的——利用AST分析代码是处理这类静态代码提取需求的标准方式,但依赖compile()获取co_names的方法确实比较取巧,容易引入不必要的干扰(比如如果TYPE_CHECKING块里有其他变量引用,也会被包含进去)。下面是两种更规范、可维护的实现方式:

方法一:手动遍历AST节点,精准提取导入

这种方式直接遍历AST树,定位到if TYPE_CHECKING块后,专门处理其中的Import和ImportFrom语句,提取完整的模块/对象路径:

import ast

def extract_type_checking_imports(code: str) -> list[str]:
    imports = []
    tree = ast.parse(code)
    
    for node in ast.walk(tree):
        # 定位到if TYPE_CHECKING的条件语句
        if isinstance(node, ast.If) and isinstance(node.test, ast.Name) and node.test.id == 'TYPE_CHECKING':
            # 遍历if块内的所有语句
            for stmt in node.body:
                # 处理`import xxx`格式的导入
                if isinstance(stmt, ast.Import):
                    for alias in stmt.names:
                        imports.append(alias.name)
                # 处理`from xxx import yyy`格式的导入
                elif isinstance(stmt, ast.ImportFrom):
                    module_prefix = stmt.module or ""
                    for alias in stmt.names:
                        # 拼接完整的路径(处理相对导入和绝对导入)
                        full_import_path = f"{module_prefix}.{alias.name}" if module_prefix else alias.name
                        imports.append(full_import_path)
    return imports

# 测试示例
sample_code = """
from typing import TYPE_CHECKING
if TYPE_CHECKING:
    import abc
    from django.utils import timezone
    from . import local_utils
if 'aaa':
    import os  # 不会被提取
print('hello world')
"""

print(extract_type_checking_imports(sample_code))
# 输出: ['abc', 'django.utils.timezone', 'local_utils']

优势

  • 逻辑直接清晰,只针对导入语句处理,不会混入其他变量名
  • 支持多种导入格式(绝对导入、相对导入、多对象导入)
  • 不需要编译代码,避免了co_names可能带来的冗余内容

方法二:使用ast.NodeVisitor(更符合AST规范的方式)

ast.NodeVisitor是Python AST模块提供的专门用于遍历AST树的工具类,代码结构更模块化,扩展性更强(比如后续要处理嵌套块、注释过滤等需求时,更容易扩展):

import ast

class TypeCheckingImportCollector(ast.NodeVisitor):
    def __init__(self):
        self.collected_imports = []
        self._in_target_block = False

    def visit_If(self, node):
        # 检查当前if块是否是TYPE_CHECKING条件
        if isinstance(node.test, ast.Name) and node.test.id == 'TYPE_CHECKING':
            self._in_target_block = True
            # 遍历块内的所有语句
            for stmt in node.body:
                self.visit(stmt)
            self._in_target_block = False
        else:
            # 非目标块,继续遍历子节点但不改变状态
            self.generic_visit(node)

    def visit_Import(self, node):
        if self._in_target_block:
            for alias in node.names:
                self.collected_imports.append(alias.name)

    def visit_ImportFrom(self, node):
        if self._in_target_block:
            module_prefix = node.module or ""
            for alias in node.names:
                full_path = f"{module_prefix}.{alias.name}" if module_prefix else alias.name
                self.collected_imports.append(full_path)

def extract_type_checking_imports(code: str) -> list[str]:
    visitor = TypeCheckingImportCollector()
    visitor.visit(ast.parse(code))
    return visitor.collected_imports

# 测试扩展场景
sample_code = """
from typing import TYPE_CHECKING
if TYPE_CHECKING:
    import abc
    from django.utils import timezone
    from module.sub import func1, func2
if TYPE_CHECKING:
    from .core import BaseModel
print('hello world')
"""

print(extract_type_checking_imports(sample_code))
# 输出: ['abc', 'django.utils.timezone', 'module.sub.func1', 'module.sub.func2', '.core.BaseModel']

优势

  • 遵循AST处理的标准模式,代码可读性和可维护性更高
  • 状态管理更清晰,避免手动遍历可能出现的遗漏
  • 扩展性强:如果需要处理更复杂的情况(比如嵌套if块、忽略注释导入等),只需新增对应的visit_xxx方法即可

对比原方法的改进

你的原方法依赖compile()后的co_names,虽然能得到结果,但存在以下问题:

  • 会混入TYPE_CHECKING块内的非导入名字(比如如果块内有x = 1,x也会出现在co_names中)
  • 无法区分导入的完整路径(比如from django.utils import timezone,co_names只会给出django、utils、timezone,需要手动拼接)
  • 逻辑不够直观,后续维护者需要理解compile()的内部行为才能修改代码

以上两种方法都解决了这些问题,是更规范的实现方案。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.12 03:54:07