Code Llama指令调优秘籍:让AI帮你写单元测试的3种高阶玩法(Python版)

如果你是一名测试工程师或者全栈开发者,最近肯定没少听说Code Llama这个名字。这个基于Llama 2的代码大模型,在开源社区里掀起了一阵不小的波澜。但说实话,大多数教程还停留在“如何安装”和“基础使用”的层面,真正能把它用到自动化测试实战中的深度内容并不多见。

我花了近一个月时间,在几个实际项目中系统性地测试了Code Llama在单元测试生成方面的表现。从最初的简单函数测试,到复杂的类继承和多线程场景,再到需要结合业务上下文的集成测试——我发现,用好Code Llama的关键不在于模型本身,而在于如何“调教”它。就像训练一名优秀的测试工程师,你需要告诉它你的测试哲学、框架偏好、甚至代码风格。

今天我要分享的,就是我在实战中总结出的三套高阶玩法。这些方法不是简单的API调用,而是结合了系统提示词工程、上下文填充技巧、以及长上下文优化策略的完整工作流。无论你是想提升现有测试覆盖率,还是为遗留代码库快速补充测试,这些方法都能给你带来实实在在的效率提升。

1. 系统提示词定制:打造你的专属测试风格

很多人用Code Llama生成测试,就是简单地把函数代码扔进去,然后说“写个测试”。结果往往不尽如人意——生成的测试要么过于简单,要么风格混乱,要么完全不符合项目规范。问题出在哪里?Code Llama不知道你的“测试文化”是什么

1.1 构建测试风格模板

每个团队、每个项目都有自己的测试习惯。有的喜欢用pytestassert语句,有的偏好unittestTestCase类;有的要求每个测试用例都有详细的注释,有的则追求最小化的测试代码。你需要把这些偏好“教”给Code Llama。

我建议创建一个测试风格配置文件,用YAML或JSON定义你的测试规范。下面是一个我常用的模板:

test_style:
  framework: "pytest"
  assertions:
    primary: "assert"
    use_rich_comparison: true
    custom_matchers: ["pytest.approx"]
  fixtures:
    auto_use: true
    naming_convention: "test_"
  mocking:
    library: "unittest.mock"
    prefer_patch_over_MagicMock: true
  coverage:
    require_docstrings: true
    include_edge_cases: true
    boundary_testing: true
  structure:
    arrange_act_assert: true
    one_assert_per_test: false
    test_class_per_module: true

把这个配置转换成Code Llama能理解的系统提示词:

system_prompt = """
你是一个专业的Python测试工程师,专门为项目生成高质量、可维护的单元测试。
请严格按照以下规范编写测试代码:

1. **测试框架**: 使用pytest,不要使用unittest
2. **断言风格**: 
   - 主要使用简单的assert语句
   - 浮点数比较使用pytest.approx
   - 异常断言使用pytest.raises
3. **测试结构**:
   - 每个测试函数以test_开头
   - 使用AAA模式(Arrange-Act-Assert)
   - 为复杂测试添加docstring说明
4. **Mock使用**:
   - 优先使用unittest.mock.patch
   - 为每个mock对象添加明确的spec
5. **测试覆盖**:
   - 必须包含正常路径测试
   - 必须包含边界条件测试
   - 必须包含异常路径测试
6. **代码风格**:
   - 遵循PEP 8
   - 使用有意义的变量名
   - 避免魔法数字

现在,请为以下函数生成单元测试:
"""

1.2 动态提示词生成器

在实际使用中,你不可能每次都手动编写这么长的提示词。我开发了一个提示词生成器,可以根据不同的测试场景动态调整提示词内容:

class TestPromptGenerator:
    def __init__(self, style_config):
        self.style = style_config
    
    def generate_for_function(self, func_code, func_name, 
                            include_docstring=True,
                            include_edge_cases=True,
                            mock_dependencies=False):
        """为特定函数生成测试提示词"""
        
        prompt_parts = []
        
        # 1. 系统指令部分
        prompt_parts.append(self._generate_system_instruction())
        
        # 2. 函数上下文
        prompt_parts.append(f"## 待测试函数\n```python\n{func_code}\n```")
        
        # 3. 具体要求
        requirements = []
        if include_docstring:
            requirements.append("每个测试函数必须有清晰的docstring")
        if include_edge_cases:
            requirements.append("必须包含边界条件测试")
        if mock_dependencies:
            requirements.append("使用适当的mock隔离外部依赖")
            
        if requirements:
            prompt_parts.append(f"## 测试要求\n- " + "\n- ".join(requirements))
        
        # 4. 输出格式
        prompt_parts.append("""
## 输出格式
请直接输出完整的测试代码,以```python开始,以```结束。
不要添加任何解释性文字。
""")
        
        return "\n\n".join(prompt_parts)
    
    def _generate_system_instruction(self):
        """生成系统指令部分"""
        return f"""你是一个专业的{self.style['framework']}测试工程师。
请严格按照以下规范编写测试代码:
- 断言方式: {self.style['assertions']['primary']}
- 测试命名: {self.style['fixtures']['naming_convention']}*
- 测试结构: {'AAA模式' if self.style['structure']['arrange_act_assert'] else '自由格式'}
- Mock库: {self.style['mocking']['library']}
"""

1.3 实战案例:数据库操作函数测试

让我们看一个具体的例子。假设我们有一个数据库操作函数:

def get_user_by_id(user_id: int, db_connection) -> dict:
    """根据用户ID从数据库获取用户信息"""
    if not isinstance(user_id, int) or user_id <= 0:
        raise ValueError("user_id必须是正整数")
    
    cursor = db_connection.cursor()
    cursor.execute("SELECT * FROM users WHERE id = %s", (user_id,))
    result = cursor.fetchone()
    
    if not result:
        raise ValueError(f"用户ID {user_id} 不存在")
    
    return {
        'id': result[0],
        'name': result[1],
        'email': result[2],
        'created_at': result[3]
    }

使用我们的提示词生成器,可以得到这样的测试代码:

import pytest
from unittest.mock import Mock, MagicMock
from datetime import datetime

def test_get_user_by_id_valid_input():
    """测试正常情况下的用户查询"""
    # Arrange
    mock_cursor = MagicMock()
    mock_cursor.fetchone.return_value = (1, '张三', 'zhangsan@example.com', datetime.now())
    
    mock_conn = Mock()
    mock_conn.cursor.return_value = mock_cursor
    
    # Act
    result = get_user_by_id(1, mock_conn)
    
    # Assert
    assert result['id'] == 1
    assert result['name'] == '张三'
    assert result['email'] == 'zhangsan@example.com'
    assert isinstance(result['created_at'], datetime)
    mock_cursor.execute.assert_called_once_with(
        "SELECT * FROM users WHERE id = %s", (1,)
    )

def test_get_user_by_id_invalid_user_id():
    """测试无效的用户ID"""
    mock_conn = Mock()
    
    # 测试非整数ID
    with pytest.raises(ValueError, match="user_id必须是正整数"):
        get_user_by_id("invalid", mock_conn)
    
    # 测试零或负数ID
    with pytest.raises(ValueError, match="user_id必须是正整数"):
        get_user_by_id(0, mock_conn)
    
    with pytest.raises(ValueError, match="user_id必须是正整数"):
        get_user_by_id(-1, mock_conn)

def test_get_user_by_id_user_not_found():
    """测试用户不存在的情况"""
    # Arrange
    mock_cursor = MagicMock()
    mock_cursor.fetchone.return_value = None
    
    mock_conn = Mock()
    mock_conn.cursor.return_value = mock_cursor
    
    # Act & Assert
    with pytest.raises(ValueError, match="用户ID 999 不存在"):
        get_user_by_id(999, mock_conn)

提示:在实际使用中,你可以把这个提示词生成器集成到你的CI/CD流程中,每次代码提交时自动为新增函数生成测试模板,大大减少手动编写测试的工作量。

2. 结合pytest框架的上下文填充技巧

Code Llama有一个非常强大的功能叫Fill-in-the-Middle(FIM),也就是代码填充。这个功能在编写测试时特别有用,因为测试代码往往有固定的模式,你只需要模型填充关键部分。

2.1 FIM在测试生成中的应用

FIM的基本格式是<PRE> {前缀} <SUF> {后缀} <MID>,模型会在<PRE><SUF>之间生成内容。在测试场景中,我们可以这样使用:

from transformers import AutoTokenizer, AutoModelForCausalLM
import torch

def generate_test_with_fim(model, tokenizer, function_code, test_framework="pytest"):
    """使用FIM生成测试代码"""
    
    # 构建FIM提示
    prefix = f'''import pytest
from unittest.mock import Mock, patch

{function_code}

def test_'''
    
    suffix = '''
    # 这里由模型填充具体的测试逻辑
    pass
'''
    
    prompt = f"<PRE> {prefix} <SUF>{suffix} <MID>"
    
    inputs = tokenizer(prompt, return_tensors="pt").to("cuda")
    
    # 生成参数配置
    generation_config = {
        "max_new_tokens": 500,
        "temperature": 0.2,  # 较低的温度保证代码质量
        "top_p": 0.95,
        "do_sample": True,
        "pad_token_id": tokenizer.eos_token_id,
    }
    
    output = model.generate(**inputs, **generation_config)
    generated_text = tokenizer.decode(output[0], skip_special_tokens=True)
    
    # 提取生成的测试代码
    # 这里需要解析生成的内容,提取<PRE>和<SUF>之间的部分
    return extract_fim_output(generated_text)

2.2 上下文感知的测试生成

单纯的FIM还不够智能。在实际项目中,测试往往需要了解类的结构、模块的导入、以及依赖关系。我们可以利用Code Llama的100k上下文窗口,把整个模块甚至整个包的结构都喂给模型。

下面是一个更高级的示例,展示如何为整个类生成测试:

class TestGeneratorWithContext:
    def __init__(self, model, tokenizer, max_context_length=80000):
        self.model = model
        self.tokenizer = tokenizer
        self.max_context_length = max_context_length
    
    def generate_tests_for_class(self, class_file_path, output_test_path):
        """为整个Python类生成测试"""
        
        # 1. 读取类文件
        with open(class_file_path, 'r', encoding='utf-8') as f:
            class_code = f.read()
        
        # 2. 读取相关依赖
        imports = self._extract_imports(class_code)
        dependent_modules = self._find_dependent_modules(class_file_path)
        
        # 3. 构建上下文
        context = self._build_context(
            class_code=class_code,
            imports=imports,
            dependent_modules=dependent_modules,
            test_framework="pytest"
        )
        
        # 4. 生成测试
        prompt = self._create_class_test_prompt(context)
        tests = self._generate_with_model(prompt)
        
        # 5. 保存测试文件
        self._write_test_file(tests, output_test_path)
    
    def _build_context(self, class_code, imports, dependent_modules, test_framework):
        """构建包含所有必要上下文的提示"""
        
        context_parts = []
        
        # 添加框架导入
        context_parts.append(f"# 测试框架: {test_framework}")
        if test_framework == "pytest":
            context_parts.append("import pytest")
            context_parts.append("from unittest.mock import Mock, MagicMock, patch")
        
        # 添加原始导入
        context_parts.append("\n# 原始模块导入")
        context_parts.extend(imports)
        
        # 添加依赖模块的关键部分
        context_parts.append("\n# 相关依赖代码")
        for module in dependent_modules[:3]:  # 限制数量避免过长
            context_parts.append(f"# {module['name']} 的关键函数:")
            context_parts.append(module['key_functions'])
        
        # 添加要测试的类
        context_parts.append("\n# 要测试的类")
        context_parts.append(class_code)
        
        return "\n".join(context_parts)

2.3 实战:为Flask应用生成集成测试

让我们看一个更复杂的例子——为Flask应用的视图函数生成测试。这里的关键是提供足够的上下文,让模型理解整个应用的结构。

# app.py 的一部分
from flask import Flask, request, jsonify
from database import get_db_connection
from auth import require_auth

app = Flask(__name__)

@app.route('/api/users/<int:user_id>', methods=['GET'])
@require_auth
def get_user(user_id):
    """获取用户信息API"""
    try:
        conn = get_db_connection()
        user = get_user_by_id(user_id, conn)
        return jsonify({
            'success': True,
            'data': user
        }), 200
    except ValueError as e:
        return jsonify({
            'success': False,
            'error': str(e)
        }), 404
    except Exception as e:
        return jsonify({
            'success': False,
            'error': '服务器内部错误'
        }), 500

使用我们的上下文感知生成器,可以得到这样的集成测试:

import pytest
from unittest.mock import Mock, patch, MagicMock
from app import app

class TestGetUserAPI:
    """测试用户获取API"""
    
    def setup_method(self):
        self.client = app.test_client()
        self.app_context = app.app_context()
        self.app_context.push()
    
    def teardown_method(self):
        self.app_context.pop()
    
    @patch('app.get_db_connection')
    @patch('app.get_user_by_id')
    def test_get_user_success(self, mock_get_user, mock_get_connection):
        """测试成功获取用户信息"""
        # Arrange
        mock_user_data = {
            'id': 1,
            'name': '测试用户',
            'email': 'test@example.com'
        }
        mock_get_user.return_value = mock_user_data
        
        mock_conn = Mock()
        mock_get_connection.return_value = mock_conn
        
        # 模拟认证头
        headers = {'Authorization': 'Bearer valid_token'}
        
        # Act
        response = self.client.get('/api/users/1', headers=headers)
        
        # Assert
        assert response.status_code == 200
        data = response.get_json()
        assert data['success'] is True
        assert data['data'] == mock_user_data
        mock_get_user.assert_called_once_with(1, mock_conn)
    
    @patch('app.get_db_connection')
    def test_get_user_not_found(self, mock_get_connection):
        """测试用户不存在的情况"""
        # Arrange
        mock_conn = Mock()
        mock_get_connection.return_value = mock_conn
        
        # 模拟get_user_by_id抛出ValueError
        with patch('app.get_user_by_id') as mock_get_user:
            mock_get_user.side_effect = ValueError("用户ID 999 不存在")
            
            headers = {'Authorization': 'Bearer valid_token'}
            
            # Act
            response = self.client.get('/api/users/999', headers=headers)
            
            # Assert
            assert response.status_code == 404
            data = response.get_json()
            assert data['success'] is False
            assert "用户ID 999 不存在" in data['error']
    
    def test_get_user_unauthorized(self):
        """测试未授权访问"""
        # Act
        response = self.client.get('/api/users/1')  # 没有认证头
        
        # Assert
        assert response.status_code == 401
    
    @patch('app.get_db_connection')
    def test_get_user_server_error(self, mock_get_connection):
        """测试服务器内部错误"""
        # Arrange
        mock_conn = Mock()
        mock_get_connection.return_value = mock_conn
        
        # 模拟get_user_by_id抛出非ValueError异常
        with patch('app.get_user_by_id') as mock_get_user:
            mock_get_user.side_effect = Exception("数据库连接失败")
            
            headers = {'Authorization': 'Bearer valid_token'}
            
            # Act
            response = self.client.get('/api/users/1', headers=headers)
            
            # Assert
            assert response.status_code == 500
            data = response.get_json()
            assert data['success'] is False
            assert data['error'] == '服务器内部错误'

注意:生成集成测试时,一定要确保模型理解整个请求-响应周期,包括中间件、认证、数据库连接等。这就是为什么我们需要提供完整的上下文。

3. 利用100k长上下文处理复杂测试用例

Code Llama最大的优势之一就是支持100k tokens的上下文窗口。这意味着我们可以把整个代码库的相关部分都放进去,让模型基于完整的上下文生成测试。

3.1 长上下文优化策略

但是,直接把100k tokens扔给模型并不总是最佳选择。我们需要智能地选择哪些上下文是相关的。下面是我总结的上下文选择策略

上下文类型 包含内容 重要性 最大长度
目标代码 要测试的函数/类 必须 完整代码
直接依赖 直接调用的函数/类 关键部分
间接依赖 间接相关的模块 摘要或接口
测试模式 项目中的测试示例 2-3个示例
配置信息 测试配置、环境变量 必要部分

实现这个策略的代码:

class SmartContextBuilder:
    def __init__(self, project_root, max_tokens=80000):
        self.project_root = project_root
        self.max_tokens = max_tokens
        self.tokenizer = AutoTokenizer.from_pretrained("codellama/CodeLlama-34b-Instruct-hf")
    
    def build_context_for_testing(self, target_file_path, target_element):
        """为目标元素构建测试上下文"""
        
        context_parts = []
        used_tokens = 0
        
        # 1. 添加目标代码
        target_code = self._read_file(target_file_path)
        target_tokens = len(self.tokenizer.encode(target_code))
        
        if target_tokens > self.max_tokens * 0.3:  # 不超过30%的上下文
            target_code = self._summarize_code(target_code)
            target_tokens = len(self.tokenizer.encode(target_code))
        
        context_parts.append(f"# 目标代码: {target_file_path}")
        context_parts.append(target_code)
        used_tokens += target_tokens
        
        # 2. 分析依赖并添加
        dependencies = self._analyze_dependencies(target_file_path, target_element)
        
        for dep in dependencies:
            if used_tokens >= self.max_tokens * 0.8:  # 保留20%给生成
                break
                
            dep_code = self._read_file(dep['path'])
            dep_tokens = len(self.tokenizer.encode(dep_code))
            
            # 如果依赖太大,只取相关部分
            if dep_tokens > 5000:
                dep_code = self._extract_relevant_parts(dep_code, target_element)
                dep_tokens = len(self.tokenizer.encode(dep_code))
            
            if used_tokens + dep_tokens < self.max_tokens * 0.8:
                context_parts.append(f"\n# 相关依赖: {dep['name']}")
                context_parts.append(dep_code)
                used_tokens += dep_tokens
        
        # 3. 添加测试示例
        test_examples = self._find_test_examples()
        for example in test_examples[:2]:  # 只取2个示例
            example_tokens = len(self.tokenizer.encode(example))
            if used_tokens + example_tokens < self.max_tokens * 0.9:
                context_parts.append("\n# 测试示例参考")
                context_parts.append(example)
                used_tokens += example_tokens
        
        return "\n".join(context_parts)

3.2 复杂业务逻辑的测试生成

让我们看一个电商系统中的复杂业务逻辑——购物车结算:

# cart.py
class ShoppingCart:
    def __init__(self, user_id, inventory_service, discount_service):
        self.user_id = user_id
        self.items = []
        self.inventory_service = inventory_service
        self.discount_service = discount_service
    
    def add_item(self, product_id, quantity):
        """添加商品到购物车"""
        if quantity <= 0:
            raise ValueError("数量必须大于0")
        
        # 检查库存
        available = self.inventory_service.check_stock(product_id, quantity)
        if not available:
            raise ValueError(f"商品 {product_id} 库存不足")
        
        # 查找是否已存在
        for item in self.items:
            if item['product_id'] == product_id:
                item['quantity'] += quantity
                break
        else:
            self.items.append({
                'product_id': product_id,
                'quantity': quantity,
                'unit_price': self.inventory_service.get_price(product_id)
            })
    
    def remove_item(self, product_id, quantity=None):
        """从购物车移除商品"""
        for i, item in enumerate(self.items):
            if item['product_id'] == product_id:
                if quantity is None or quantity >= item['quantity']:
                    self.items.pop(i)
                else:
                    item['quantity'] -= quantity
                break
    
    def calculate_total(self):
        """计算购物车总价"""
        subtotal = sum(item['quantity'] * item['unit_price'] 
                      for item in self.items)
        
        # 应用折扣
        discount = self.discount_service.calculate_discount(
            self.user_id, subtotal, self.items
        )
        
        return {
            'subtotal': subtotal,
            'discount': discount,
            'total': subtotal - discount,
            'item_count': len(self.items)
        }
    
    def checkout(self, payment_service, shipping_address):
        """结算购物车"""
        if not self.items:
            raise ValueError("购物车为空")
        
        total_info = self.calculate_total()
        
        # 处理支付
        payment_result = payment_service.process_payment(
            self.user_id, total_info['total'], '购物车结算'
        )
        
        if not payment_result['success']:
            raise ValueError(f"支付失败: {payment_result['reason']}")
        
        # 更新库存
        for item in self.items:
            self.inventory_service.reduce_stock(
                item['product_id'], item['quantity']
            )
        
        # 生成订单
        order = {
            'user_id': self.user_id,
            'items': self.items.copy(),
            'total': total_info['total'],
            'payment_id': payment_result['payment_id'],
            'shipping_address': shipping_address,
            'status': 'paid'
        }
        
        # 清空购物车
        self.items.clear()
        
        return order

使用长上下文策略,我们可以生成覆盖所有边界条件的测试:

import pytest
from unittest.mock import Mock, MagicMock, call
from cart import ShoppingCart

class TestShoppingCart:
    """测试购物车功能"""
    
    def setup_method(self):
        self.mock_inventory = Mock()
        self.mock_discount = Mock()
        self.mock_payment = Mock()
        
        self.cart = ShoppingCart(
            user_id=123,
            inventory_service=self.mock_inventory,
            discount_service=self.mock_discount
        )
    
    def test_add_item_success(self):
        """测试成功添加商品"""
        # Arrange
        self.mock_inventory.check_stock.return_value = True
        self.mock_inventory.get_price.return_value = 100.0
        
        # Act
        self.cart.add_item("PROD001", 2)
        
        # Assert
        assert len(self.cart.items) == 1
        assert self.cart.items[0]['product_id'] == "PROD001"
        assert self.cart.items[0]['quantity'] == 2
        assert self.cart.items[0]['unit_price'] == 100.0
        self.mock_inventory.check_stock.assert_called_once_with("PROD001", 2)
    
    def test_add_item_insufficient_stock(self):
        """测试库存不足的情况"""
        # Arrange
        self.mock_inventory.check_stock.return_value = False
        
        # Act & Assert
        with pytest.raises(ValueError, match="商品 PROD001 库存不足"):
            self.cart.add_item("PROD001", 5)
    
    def test_add_item_invalid_quantity(self):
        """测试无效数量"""
        test_cases = [
            (0, "数量必须大于0"),
            (-1, "数量必须大于0"),
            (-100, "数量必须大于0")
        ]
        
        for quantity, expected_error in test_cases:
            with pytest.raises(ValueError, match=expected_error):
                self.cart.add_item("PROD001", quantity)
    
    def test_add_item_existing_product(self):
        """测试添加已存在的商品"""
        # Arrange
        self.mock_inventory.check_stock.return_value = True
        self.mock_inventory.get_price.return_value = 50.0
        
        # 第一次添加
        self.cart.add_item("PROD002", 1)
        
        # 重置mock调用记录
        self.mock_inventory.check_stock.reset_mock()
        self.mock_inventory.get_price.reset_mock()
        
        # 再次添加同一商品
        self.mock_inventory.check_stock.return_value = True
        
        # Act
        self.cart.add_item("PROD002", 3)
        
        # Assert
        assert len(self.cart.items) == 1
        assert self.cart.items[0]['quantity'] == 4  # 1 + 3
        self.mock_inventory.get_price.assert_not_called()  # 价格只获取一次
    
    def test_remove_item_completely(self):
        """测试完全移除商品"""
        # Arrange
        self.mock_inventory.check_stock.return_value = True
        self.mock_inventory.get_price.return_value = 30.0
        self.cart.add_item("PROD003", 2)
        
        # Act
        self.cart.remove_item("PROD003")
        
        # Assert
        assert len(self.cart.items) == 0
    
    def test_remove_item_partially(self):
        """测试部分移除商品"""
        # Arrange
        self.mock_inventory.check_stock.return_value = True
        self.mock_inventory.get_price.return_value = 25.0
        self.cart.add_item("PROD004", 5)
        
        # Act
        self.cart.remove_item("PROD004", 3)
        
        # Assert
        assert len(self.cart.items) == 1
        assert self.cart.items[0]['quantity'] == 2  # 5 - 3
    
    def test_calculate_total_with_discount(self):
        """测试计算含折扣的总价"""
        # Arrange
        self.mock_inventory.check_stock.return_value = True
        self.mock_inventory.get_price.side_effect = [100.0, 200.0]
        
        self.cart.add_item("PROD005", 1)  # 100元
        self.cart.add_item("PROD006", 2)  # 400元
        
        # 小计应该是500元
        self.mock_discount.calculate_discount.return_value = 50.0  # 50元折扣
        
        # Act
        total_info = self.cart.calculate_total()
        
        # Assert
        assert total_info['subtotal'] == 500.0
        assert total_info['discount'] == 50.0
        assert total_info['total'] == 450.0
        assert total_info['item_count'] == 2
        
        self.mock_discount.calculate_discount.assert_called_once_with(
            123, 500.0, self.cart.items
        )
    
    def test_checkout_success(self):
        """测试成功结算"""
        # Arrange
        self.mock_inventory.check_stock.return_value = True
        self.mock_inventory.get_price.return_value = 150.0
        self.mock_discount.calculate_discount.return_value = 0.0
        
        self.cart.add_item("PROD007", 1)
        
        self.mock_payment.process_payment.return_value = {
            'success': True,
            'payment_id': 'PAY_123456',
            'reason': None
        }
        
        shipping_address = {
            'name': '张三',
            'address': '北京市朝阳区',
            'phone': '13800138000'
        }
        
        # Act
        order = self.cart.checkout(self.mock_payment, shipping_address)
        
        # Assert
        assert order['user_id'] == 123
        assert order['total'] == 150.0
        assert order['payment_id'] == 'PAY_123456'
        assert order['status'] == 'paid'
        assert len(self.cart.items) == 0  # 购物车已清空
        
        # 验证库存更新被调用
        self.mock_inventory.reduce_stock.assert_called_once_with("PROD007", 1)
    
    def test_checkout_empty_cart(self):
        """测试空购物车结算"""
        # Act & Assert
        with pytest.raises(ValueError, match="购物车为空"):
            self.cart.checkout(self.mock_payment, {})
    
    def test_checkout_payment_failed(self):
        """测试支付失败"""
        # Arrange
        self.mock_inventory.check_stock.return_value = True
        self.mock_inventory.get_price.return_value = 200.0
        self.mock_discount.calculate_discount.return_value = 0.0
        
        self.cart.add_item("PROD008", 1)
        
        self.mock_payment.process_payment.return_value = {
            'success': False,
            'payment_id': None,
            'reason': '余额不足'
        }
        
        # Act & Assert
        with pytest.raises(ValueError, match="支付失败: 余额不足"):
            self.cart.checkout(self.mock_payment, {})
        
        # 验证库存没有更新
        self.mock_inventory.reduce_stock.assert_not_called()

3.3 性能优化与批量生成

当需要为大量函数生成测试时,我们需要考虑性能优化。下面是一个批量生成测试的脚本,使用了异步处理和缓存:

import asyncio
import hashlib
import json
from pathlib import Path
from typing import List, Dict
import aiohttp

class BatchTestGenerator:
    def __init__(self, model_endpoint, api_key, cache_dir=".test_cache"):
        self.endpoint = model_endpoint
        self.api_key = api_key
        self.cache_dir = Path(cache_dir)
        self.cache_dir.mkdir(exist_ok=True)
    
    async def generate_tests_for_project(self, project_path: str, 
                                       concurrency: int = 5) -> Dict[str, str]:
        """为整个项目生成测试"""
        
        # 1. 扫描项目中的Python文件
        python_files = self._find_python_files(project_path)
        
        # 2. 分析每个文件中的函数和类
        test_targets = []
        for file_path in python_files:
            targets = self._analyze_file(file_path)
            test_targets.extend(targets)
        
        # 3. 分批生成测试
        semaphore = asyncio.Semaphore(concurrency)
        results = {}
        
        async def generate_for_target(target):
            async with semaphore:
                cache_key = self._get_cache_key(target)
                cached = self._get_from_cache(cache_key)
                
                if cached:
                    return target['name'], cached
                
                # 生成测试
                test_code = await self._generate_test_async(target)
                
                # 缓存结果
                self._save_to_cache(cache_key, test_code)
                
                return target['name'], test_code
        
        # 4. 并发执行
        tasks = [generate_for_target(target) for target in test_targets]
        completed = await asyncio.gather(*tasks, return_exceptions=True)
        
        # 5. 处理结果
        for result in completed:
            if isinstance(result, Exception):
                print(f"生成测试时出错: {result}")
                continue
            
            name, test_code = result
            results[name] = test_code
        
        return results
    
    def _get_cache_key(self, target: Dict) -> str:
        """生成缓存键"""
        content = json.dumps(target, sort_keys=True)
        return hashlib.md5(content.encode()).hexdigest()
    
    def _get_from_cache(self, cache_key: str) -> str:
        """从缓存获取测试代码"""
        cache_file = self.cache_dir / f"{cache_key}.txt"
        if cache_file.exists():
            return cache_file.read_text(encoding='utf-8')
        return None
    
    def _save_to_cache(self, cache_key: str, test_code: str):
        """保存测试代码到缓存"""
        cache_file = self.cache_dir / f"{cache_key}.txt"
        cache_file.write_text(test_code, encoding='utf-8')
    
    async def _generate_test_async(self, target: Dict) -> str:
        """异步生成测试代码"""
        headers = {
            "Authorization": f"Bearer {self.api_key}",
            "Content-Type": "application/json"
        }
        
        # 构建提示词
        prompt = self._build_prompt(target)
        
        payload = {
            "model": "codellama-34b-instruct",
            "messages": [
                {"role": "system", "content": "你是一个专业的测试工程师。"},
                {"role": "user", "content": prompt}
            ],
            "temperature": 0.1,
            "max_tokens": 2000
        }
        
        async with aiohttp.ClientSession() as session:
            async with session.post(
                self.endpoint,
                headers=headers,
                json=payload,
                timeout=60
            ) as response:
                if response.status == 200:
                    result = await response.json()
                    return result['choices'][0]['message']['content']
                else:
                    raise Exception(f"API请求失败: {response.status}")

4. HuggingFace TGI服务端的性能优化参数配置

如果你在团队中使用Code Llama,很可能会通过HuggingFace的**Text Generation Inference(TGI)**服务来部署。正确的参数配置可以显著提升生成质量和速度。

4.1 TGI部署配置

下面是一个优化的Docker Compose配置,针对测试生成场景做了特别优化:

# docker-compose.yml
version: '3.8'

services:
  tgi-codellama:
    image: ghcr.io/huggingface/text-generation-inference:latest
    container_name: codellama-test-generator
    ports:
      - "8080:80"
    volumes:
      - ./models:/data
    environment:
      - MODEL_ID=codellama/CodeLlama-13b-Instruct-hf
      - NUM_SHARD=2
      - QUANTIZE=bitsandbytes-nf4
      - MAX_BATCH_PREFILL_TOKENS=4096
      - MAX_BATCH_TOTAL_TOKENS=8192
      - MAX_INPUT_LENGTH=80000
      - MAX_TOTAL_TOKENS=100000
    deploy:
      resources:
        reservations:
          devices:
            - driver: nvidia
              count: 2
              capabilities: [gpu]
    command: >
      --model-id ${MODEL_ID}
      --num-shard ${NUM_SHARD}
      --quantize ${QUANTIZE}
      --max-batch-prefill-tokens ${MAX_BATCH_PREFILL_TOKENS}
      --max-batch-total-tokens ${MAX_BATCH_TOTAL_TOKENS}
      --max-input-length ${MAX_INPUT_LENGTH}
      --max-total-tokens ${MAX_TOTAL_TOKENS}
      --dtype bfloat16
      --trust-remote-code

4.2 生成参数调优

不同的测试生成场景需要不同的生成参数。下面是一个参数配置表,我根据实际使用经验总结出来的:

场景类型 temperature top_p top_k repetition_penalty 适用情况
单元测试 0.1-0.3 0.9 50 1.1 需要确定性输出,代码必须正确
集成测试 0.2-0.4 0.95 100 1.05 需要一定创造性,但保持一致性
测试数据 0.5-0.7 0.98 200 1.0 生成测试数据,需要多样性
边界测试 0.1-0.2 0.85 30 1.2 严格遵循边界条件

对应的Python客户端配置:

class OptimizedTGIClient:
    def __init__(self, base_url="http://localhost:8080"):
        self.base_url = base_url
        self.session = requests.Session()
    
    def generate_unit_test(self, code_context, function_signature):
        """生成单元测试 - 使用确定性参数"""
        payload = {
            "inputs": self._build_test_prompt(code_context, function_signature),
            "parameters": {
                "temperature": 0.2,  # 低温度保证代码正确性
                "top_p": 0.9,
                "top_k": 50,
                "repetition_penalty": 1.1,
                "max_new_tokens": 1500,
                "do_sample": True,
                "stop_sequences": ["```", "## 测试完成"],
                "seed": 42,  # 固定种子保证可重复性
                "watermark": False,
                "details": False,
                "decoder_input_details": False
            }
        }
        
        response = self.session.post(
            f"{self.base_url}/generate",
            json=payload,
            timeout=30
        )
        
        return self._extract_test_code(response.json())
    
    def generate_integration_test(self, api_spec, dependencies):
        """生成集成测试 - 平衡确定性和创造性"""
        payload = {
            "inputs": self._build_integration_prompt(api_spec, dependencies),
            "parameters": {
                "temperature": 0.3,
                "top_p": 0.95,
                "top_k": 100,
                "repetition_penalty": 1.05,
                "max_new_tokens": 2500,
                "do_sample": True,
                "typical_p": 0.95,  # 使用typical sampling提高质量
                "stop_sequences": ["```end", "## 测试完成"],
                "seed": None,  # 不固定种子,允许一定变化
                "watermark": True
            }
        }
        
        response = self.session.post(
            f"{self.base_url}/generate",
            json=payload,
            timeout=45
        )
        
        return self._extract_test_code(response.json())
    
    def batch_generate_tests(self, test_cases, batch_size=4):
        """批量生成测试 - 优化吞吐量"""
        payload = {
            "inputs": [
                self._build_test_prompt(tc['context'], tc['function'])
                for tc in test_cases[:batch_size]
            ],
            "parameters": {
                "temperature": 0.2,
                "top_p": 0.9,
                "max_new_tokens": 1000,
                "do_sample": True,
                "batch_size": batch_size,
                "watermark": False
            }
        }
        
        response = self.session.post(
            f"{self.base_url}/generate",
            json=payload,
            timeout=60
        )
        
        return [self._extract_test_code(r) for r in response.json()]

4.3 监控与优化

在生产环境中使用TGI服务时,监控和优化是必不可少的。下面是一些关键指标和优化建议:

import time
import psutil
from prometheus_client import start_http_server, Gauge, Histogram

class TGI_Monitor:
    def __init__(self, tgi_url, metrics_port=9090):
        self.tgi_url = tgi_url
        self.session = requests.Session()
        
        # 定义监控指标
        self.request_duration = Histogram(
            'tgi_request_duration_seconds',
            'TGI请求耗时',
            ['endpoint', 'method']
        )
        
        self.token_generation_speed = Gauge(
            'tgi_tokens_per_second',
            '令牌生成速度'
        )
        
        self.batch_utilization = Gauge(
            'tgi_batch_utilization',
            '批次利用率'
        )
        
        self.gpu_memory_usage = Gauge(
            'tgi_gpu_memory_usage_bytes',
            'GPU内存使用量'
        )
        
        # 启动监控服务器
        start_http_server(metrics_port)
    
    def monitor_generation(self, prompt, params):
        """监控生成过程"""
        start_time = time.time()
        
        with self.request_duration.labels('/generate', 'POST').time():
            response = self.session.post(
                f"{self.tgi_url}/generate",
                json={"inputs": prompt, "parameters": params},
                timeout=30
            )
        
        end_time = time.time()
        duration = end_time - start_time
        
        if response.status_code == 200:
            result = response.json()
            generated_tokens = result['details']['generated_tokens']
            
            # 计算令牌生成速度
            tokens_per_second = generated_tokens / duration
            self.token_generation_speed.set(tokens_per_second)
            
            # 记录其他指标
            self._record_additional_metrics(result)
        
        return response
    
    def _record_additional_metrics(self, result):
        """记录额外指标"""
        details = result.get('details', {})
        
        # 预填充令牌数
        prefilled_tokens = details.get('prefilled_tokens', 0)
        
        # 批次统计
        if 'batch_stats' in details:
            batch_size = details['batch_stats'].get('batch_size', 1)
            batch_capacity = details['batch_stats'].get('batch_capacity', 1)
            
            utilization = batch_size / batch_capacity if batch_capacity > 0 else 0
            self.batch_utilization.set(utilization)
        
        # 获取GPU内存使用情况(如果有GPU)
        try:
            import pynvml
            pynvml.nvmlInit()
            handle = pynvml.nvmlDeviceGetHandleByIndex(0)
            info = pynvml.nvmlDeviceGetMemoryInfo(handle)
            self.gpu_memory_usage.set(info.used)
        except:
            pass  # 没有GPU或pynvml不可用
    
    def get_optimization_recommendations(self):
        """获取优化建议"""
        recommendations = []
        
        # 分析当前性能
        current_speed = self.token_generation_speed._value.get()
        current_utilization = self.batch_utilization._value.get()
        
        if current_speed and current_speed < 50:  # 低于50 tokens/秒
            recommendations.append({
                'type': 'performance',
                'message': '令牌生成速度较慢,考虑减少max_new_tokens或使用量化',
                'suggestion': '尝试使用--quantize bitsandbytes-nf4参数'
            })
        
        if current_utilization and current_utilization < 0.5:
            recommendations.append({
                'type': 'efficiency',
                'message': '批次利用率较低,考虑增加批量大小',
                'suggestion': '调整--max-batch-total-tokens参数'
            })
        
        return recommendations

4.4 实际部署经验

在我最近的一个项目中,我们为拥有300多个Python文件的代码库生成了测试覆盖。以下是我们的一些实际经验:

  1. 模型选择:对于测试生成,CodeLlama-13b-Instruct在质量和速度之间取得了最佳平衡。7B版本虽然更快,但在复杂场景下容易出错;34B版本质量更高,但推理速度慢了一倍。

  2. 批处理优化:通过调整--max-batch-total-tokens参数,我们将吞吐量提升了3倍。关键是要根据你的硬件配置找到最佳值:

# 对于2x A100 40GB
MAX_BATCH_TOTAL_TOKENS=16384

# 对于1x RTX 4090
MAX_BATCH_TOTAL_TOKENS=8192

# 对于消费级GPU(如RTX 3090)
MAX_BATCH_TOTAL_TOKENS=4096
  1. 量化策略:使用4位量化(bitsandbytes-nf4)可以将34B模型的内存占用从68GB减少到20GB左右,让它在24GB显存的GPU上也能运行。

  2. 缓存利用:TGI内置了KV缓存,但对于测试生成这种输入长、输出相对短的场景,适当调整缓存策略可以提升性能:

environment:
  - MAX_PREFILL_TOKENS=4096
  - MAX_TOTAL_TOKENS=100000
  - KV_CACHE_MAX_TOKENS=50000
  1. 错误处理与重试:在实际使用中,网络波动或GPU内存不足可能导致生成失败。实现一个健壮的重试机制很重要:
class RobustTestGenerator:
    def __init__(self, tgi_client, max_retries=3):
        self.client = tgi_client
        self.max_retries = max_retries
    
    def generate_with_retry(self, prompt, params, backoff_factor=2):
        """带重试的测试生成"""
        for attempt in range(self.max_retries):
            try:
                return self.client.generate(prompt, params)
            except requests.exceptions.Timeout:
                if attempt == self.max_retries - 1:
                    raise
                wait_time = backoff_factor ** attempt
                time.sleep(wait_time)
                print(f"请求超时,{wait_time}秒后重试...")
            except requests.exceptions.ConnectionError:
                if attempt == self.max_retries - 1:
                    raise
                print("连接错误,检查TGI服务状态...")
                self._restart_tgi_service()
                time.sleep(10)
            except Exception as e:
                if "CUDA out of memory" in str(e):
                    # 减少批次大小或输入长度
                    params['max_new_tokens'] = int(params['max_new_tokens'] * 0.8)
                    print(f"GPU内存不足,减少输出长度到{params['max_new_tokens']}")
                else:
                    raise
        
        raise Exception(f"生成失败,重试{self.max_retries}次后仍不成功")

这些优化策略让我们能够在单台服务器上同时为多个开发者提供测试生成服务,平均响应时间控制在5-10秒,完全满足日常开发需求。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐