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

如何为argparse命令行工具编写单元测试?附代码问题求助

问题分析与解决方案

测试失败的核心原因

  1. 参数传递错误:main函数期望接收字符串列表(如['--filter', 'ry']),但你的测试代码传递的是单个字符串(如'--filter=ry'),导致argparse无法正确解析参数。
  2. 硬编码文件路径:主程序中path = /path/to/data.json file存在语法错误,且硬编码路径无法在测试环境中灵活替换测试数据。
  3. 测试缺少断言:当前测试仅调用函数,未验证输出结果是否符合预期,无法完成回归测试的核心目标。
  4. 主程序逻辑缺陷:elif count的写法导致--filter和--count无法互斥,且未处理无有效参数的情况。

第一步:修复主程序代码

优化要点:

  • 修复路径语法错误,改为可配置路径(支持测试时传入模拟数据路径)
  • 添加参数互斥组,确保--filter和--count只能二选一
  • 提取数据读取逻辑,便于测试时替换为模拟数据
  • 优化过滤结果的结构,保留原始层级,提升可读性
import json
import argparse
import sys

def load_data(path):
    """提取数据读取逻辑,方便测试替换"""
    with open(path) as f:
        return json.load(f)

def main(args, data_path=None):
    # 创建参数互斥组,确保--filter和--count只能选一个
    my_parser = argparse.ArgumentParser(description='过滤包含指定模式的动物条目,或统计人员与动物总数')
    group = my_parser.add_mutually_exclusive_group(required=True)
    group.add_argument('--filter',
                       metavar='PATTERN',
                       type=str,
                       help='过滤出动物名称包含指定模式的条目')
    group.add_argument('--count',
                       action='store_true',
                       help='统计每个分组下的人员与动物总数,并更新名称显示')
    
    # 解析入参
    parsed_args = my_parser.parse_args(args)
    
    # 处理数据路径:优先使用测试传入的路径,否则用默认路径
    path = data_path or '/path/to/data.json'
    data = load_data(path)
    
    if parsed_args.filter:
        # 过滤逻辑:保留层级结构,便于后续验证
        filtered_result = [
            {
                "group": dico['name'],
                "person": person['name'],
                "animal": animal
            }
            for dico in data 
            for person in dico['people'] 
            for animal in person['animals'] 
            if parsed_args.filter in animal['name']
        ]
        return filtered_result if filtered_result else []
                
    elif parsed_args.count:
        count_result = []
        for dico in data:
            total = 0
            for person in dico['people']:
                animal_num = len(person['animals'])
                person_total = 1 + animal_num
                total += person_total
                person['name'] += f" [{animal_num}]"
            dico['name'] += f" [{total}]"
            count_result.append(dico)
        return count_result

if __name__ == '__main__':
    result = main(sys.argv[1:])
    print(json.dumps(result, indent=2))

第二步:编写正确的单元测试

测试思路:

  • 使用临时测试文件注入模拟数据,避免依赖真实环境
  • 传递符合要求的参数列表给main函数
  • 添加断言验证输出结果与预期一致
  • 覆盖过滤匹配、过滤无匹配、统计功能三种核心场景
import unittest
import json
import os
from tempfile import NamedTemporaryFile
from your_module_name import main  # 替换为你的主程序模块名

class TestCommandLineTool(unittest.TestCase):
    # 定义模拟测试数据
    TEST_DATA = [
        {
            "name": "Group 1",
            "people": [
                {
                    "name": "Alice",
                    "animals": [{"name": "Rex"}, {"name": "Mittens"}]
                },
                {
                    "name": "Bob",
                    "animals": [{"name": "Buddy"}]
                }
            ]
        },
        {
            "name": "Group 2",
            "people": [
                {
                    "name": "Charlie",
                    "animals": [{"name": "Whiskers"}]
                }
            ]
        }
    ]

    def setUp(self):
        # 创建临时JSON文件,写入测试数据
        self.temp_file = NamedTemporaryFile(mode='w', delete=False, suffix='.json')
        json.dump(self.TEST_DATA, self.temp_file)
        self.temp_file.close()

    def tearDown(self):
        # 清理临时文件
        os.unlink(self.temp_file.name)

    def test_filter_pattern_matches(self):
        # 测试过滤包含"tt"的动物
        args = ['--filter', 'tt']
        result = main(args, data_path=self.temp_file.name)
        
        expected = [
            {"group": "Group 1", "person": "Alice", "animal": {"name": "Mittens"}},
            {"group": "Group 2", "person": "Charlie", "animal": {"name": "Whiskers"}}
        ]
        self.assertEqual(result, expected)

    def test_filter_pattern_no_match(self):
        # 测试过滤不存在的模式
        args = ['--filter', 'xyz']
        result = main(args, data_path=self.temp_file.name)
        self.assertEqual(result, [])

    def test_count_function(self):
        # 测试统计功能
        args = ['--count']
        result = main(args, data_path=self.temp_file.name)
        
        expected = [
            {
                "name": "Group 1 [5]",  # Alice(1+2) + Bob(1+1) = 5
                "people": [
                    {"name": "Alice [2]", "animals": [{"name": "Rex"}, {"name": "Mittens"}]},
                    {"name": "Bob [1]", "animals": [{"name": "Buddy"}]}
                ]
            },
            {
                "name": "Group 2 [2]",  # Charlie(1+1) = 2
                "people": [
                    {"name": "Charlie [1]", "animals": [{"name": "Whiskers"}]}
                ]
            }
        ]
        self.assertEqual(result, expected)

if __name__ == '__main__':
    unittest.main()

第三步:运行测试

执行测试文件,验证功能正确性:

python -m unittest test_your_module.py -v

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.10 05:05:21