如何合规提取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
相关产品推荐
相关产品推荐

