如何用pytest测试Django模型过滤类方法?替代全流程集成测试
Django模型类方法的测试困境与解决方案
我在Django的models.py中编写了大量用于数据过滤的classmethod,这些方法会在视图中被调用以实现多种功能。例如下面的Order模型:
class Order(models.Model): user = models.ForeignKey(User, null=True, on_delete=models.SET_NULL) total = models.DecimalField(max_digits=12, decimal_places=2, default=0) order_date = models.DateTimeField(auto_now_add=True) cancelled = models.BooleanField(default=False) items = models.ManyToManyField(Items, through='OrderItems') status = models.CharField(max_length=30, choices=choice_status, default='waiting') restaurant = models.ForeignKey(Restaurant, on_delete=models.SET_NULL, related_name='orders', null=True) razorpay_payment_id = models.CharField(max_length=100, null=True, blank=True) razorpay_order_id = models.CharField(max_length=100, null=True, blank=True) mode = models.CharField(max_length=10, choices=choice_mode, default='COD', null=True, blank=True) paid = models.BooleanField(default=False, null=True, blank=True) delivery_charges = models.DecimalField(max_digits=12, decimal_places=2, default=0) user_address = models.TextField(max_length=500) user_address_lat = models.DecimalField(max_digits=9, decimal_places=6) user_address_long = models.DecimalField(max_digits=9, decimal_places=6) restaurant_address = models.TextField(max_length=500) restaurant_address_lat = models.DecimalField(max_digits=9, decimal_places=6) restaurant_address_long = models.DecimalField(max_digits=9, decimal_places=6) @classmethod def get_order_with_search_params(cls, params, queryset): search_order_id = None with contextlib.suppress(ValueError): search_order_id = int(params) return queryset.filter( Q(user__username__icontains=params) | Q(restaurant__name__icontains=params) | Q(id=search_order_id) ) @classmethod def get_order_from_status(cls, params, queryset): return queryset.filter(Q(status=params.lower())) @classmethod def get_order_from_restaurant_id(cls, params, queryset): restaurant_id = None with contextlib.suppress(ValueError): restaurant_id = int(params) return queryset.filter(Q(restaurant_id=restaurant_id)) @classmethod def get_object_from_pk(cls, pk): return get_object_or_404(cls, pk=pk) @classmethod def get_total_of_all_orders(cls, queryset=None): if not queryset: queryset = cls.objects.all() return queryset.filter(status='delivered').aggregate(total_sum=Sum('total'))['total_sum']
当前测试难题
现在测试这些类方法必须走完整业务流程:
- 注册并激活餐厅类型用户
- 餐厅用户添加商品
- 注册并激活顾客类型用户
- 顾客将商品加入购物车并下单
只有生成Order对象后才能测试过滤功能,这属于集成测试,还会影响多个关联表(比如注册用户生成Address记录、添加商品填充CartItem表等),导致我没法直接用Order.objects.create(**kwargs)创建对象来单独测试模型类方法。我是否应该跳过这些类方法的测试?或者有其他更高效的测试方式?
目前我只编写了测试模型字段和__str__方法的用例,示例如下:
def test_str_method(self): obj = State.objects.create(name='TestState') assert str(obj) == obj.name
def test_document_data(self, get_agent_document_form_initial_data): agent = User.objects.create( **{'username': 'test_agent1', 'mobile_number': '+917894541211', 'email': 'test_agent1@gmail.com'}) document_obj = Document.objects.create(agent=agent, **get_agent_document_form_initial_data(pancard_document=filename, license_document=filename)) assert document_obj.agent == agent assert document_obj.is_verified is False assert document_obj.account_no == 1234567890 assert document_obj.ifsc_code == 'ABCD0123456' assert document_obj.pancard_number == 'BNZAA2318J' assert document_obj.license_number == 'HR-0619850034761' assert document_obj.pancard_document == './media/default.jpg' assert document_obj.license_document == './media/default.jpg' assert document_obj.razorpay_contact_id is None assert document_obj.razorpay_fund_account_id is None assert str(document_obj) == f'{document_obj.id} | {document_obj.agent} | {document_obj.is_verified}'
我希望为模型中与业务流程相关的其他方法编写测试用例。
解决方案:不要跳过,用最小化测试数据单独验证类方法
这些类方法是业务逻辑的核心部分,单独测试能快速定位问题,不需要依赖完整业务流程。可以通过以下两种方式实现:
1. 最小化创建依赖对象
只创建测试所需的关联对象(比如User、Restaurant),不需要走注册激活等流程,直接用objects.create()生成最小化实例:
from django.test import TestCase from django.http import Http404 from .models import Order, User, Restaurant from django.db.models import Sum class OrderModelTests(TestCase): def setUp(self): # 创建测试用的最小化用户和餐厅 self.customer = User.objects.create(username="test_customer") self.restaurant = Restaurant.objects.create(name="test_restaurant") # 创建测试订单 self.order_delivered = Order.objects.create( user=self.customer, total=100.00, status="delivered", restaurant=self.restaurant, user_address="测试地址", user_address_lat=12.345678, user_address_long=98.765432, restaurant_address="餐厅地址", restaurant_address_lat=11.111111, restaurant_address_long=22.222222 ) self.order_waiting = Order.objects.create( user=User.objects.create(username="another_user"), total=50.00, status="waiting", restaurant=self.restaurant, user_address="测试地址2", user_address_lat=13.345678, user_address_long=99.765432, restaurant_address="餐厅地址", restaurant_address_lat=11.111111, restaurant_address_long=22.222222 ) def test_get_order_with_search_params(self): # 测试按用户名搜索 result = Order.get_order_with_search_params("test_customer", Order.objects.all()) self.assertEqual(result.count(), 1) self.assertEqual(result.first().id, self.order_delivered.id) # 测试按餐厅名称搜索 result = Order.get_order_with_search_params("test_restaurant", Order.objects.all()) self.assertEqual(result.count(), 2) # 测试按订单ID搜索 result = Order.get_order_with_search_params(str(self.order_delivered.id), Order.objects.all()) self.assertEqual(result.count(), 1) def test_get_order_from_status(self): # 测试大写状态参数 result = Order.get_order_from_status("DELIVERED", Order.objects.all()) self.assertEqual(result.count(), 1) self.assertEqual(result.first().status, "delivered") # 测试小写状态参数 result = Order.get_order_from_status("waiting", Order.objects.all()) self.assertEqual(result.count(), 1) def test_get_order_from_restaurant_id(self): # 测试有效餐厅ID result = Order.get_order_from_restaurant_id(str(self.restaurant.id), Order.objects.all()) self.assertEqual(result.count(), 2) # 测试无效ID result = Order.get_order_from_restaurant_id("abc", Order.objects.all()) self.assertEqual(result.count(), 0) def test_get_object_from_pk(self): # 测试存在的ID order = Order.get_object_from_pk(self.order_delivered.id) self.assertEqual(order.id, self.order_delivered.id) # 测试不存在的ID,应抛出404 with self.assertRaises(Http404): Order.get_object_from_pk(9999) def test_get_total_of_all_orders(self): # 测试默认查询集 total = Order.get_total_of_all_orders() self.assertEqual(total, 100.00) # 测试自定义查询集 queryset = Order.objects.filter(id=self.order_waiting.id) total = Order.get_total_of_all_orders(queryset) self.assertIsNone(total)
2. 使用工厂类简化测试数据创建
如果依赖对象较多,可以用factory_boy库定义工厂,快速生成测试对象:
import factory from .models import User, Restaurant, Order class UserFactory(factory.django.DjangoModelFactory): class Meta: model = User username = factory.Sequence(lambda n: f"user_{n}") class RestaurantFactory(factory.django.DjangoModelFactory): class Meta: model = Restaurant name = factory.Sequence(lambda n: f"restaurant_{n}") class OrderFactory(factory.django.DjangoModelFactory): class Meta: model = Order user = factory.SubFactory(UserFactory) restaurant = factory.SubFactory(RestaurantFactory) total = 100.00 status = "delivered" user_address = "测试地址" user_address_lat = 12.345678 user_address_long = 98.765432 restaurant_address = "餐厅地址" restaurant_address_lat = 11.111111 restaurant_address_long = 22.222222
在测试中直接调用工厂创建对象,代码会更简洁:
def setUp(self): self.order1 = OrderFactory() self.order2 = OrderFactory(status="waiting", user__username="test_user")
核心原则
测试时只关注类方法的过滤/计算逻辑,不需要模拟完整业务流程(比如用户激活、购物车操作等无关步骤),只需要保证测试数据能覆盖方法的各种分支场景即可。
内容的提问来源于stack exchange,提问作者Parth
相关产品推荐
相关产品推荐

