from concurrent.futures import ThreadPoolExecutor from django.contrib.auth import get_user_model from django.db import close_old_connections, connections from django.test import TransactionTestCase from django.utils import timezone from basic_info import models as basic_models from business import models as business_models, services from .fixtures import create_sales_fixtures class SalesOrderConcurrencyTestCase(TransactionTestCase): reset_sequences = True def setUp(self): ( self.merchant, self.customer, self.warehouse_strict, self.warehouse_relaxed, self.warehouse_strict_out, self.product, self.operator, ) = create_sales_fixtures() basic_models.MerchantSetting.objects.filter( merchant=self.merchant, key=basic_models.MerchantSettingKeyEnum.AUTO_CREATE_STOCK_CHANGE_TASKS, ).update(val_bool=False) User = get_user_model() self.user = User.objects.create_user(username='concurrent', password='pass123') self.sales_order = services.create_sales_order( merchant=self.merchant, customer=self.customer, order_date=timezone.now().date(), warehouse=self.warehouse_strict, operator=self.operator, items=[{'product_id': self.product.id, 'numbers': [5], 'price': '12', 'unit': '米'}], created_by=self.user, ) def test_concurrent_sales_order_approval_updates_balance_once(self): def approve(): # ThreadPoolExecutor 会复用线程;Django 的 DB connection 是线程局部的, # 若不显式关闭,可能导致测试 DB 在 teardown 时仍被占用,无法 DROP。 close_old_connections() try: services.review_sales_order( sales_order_id=self.sales_order.id, target_status=business_models.SalesOrderStatusEnum.APPROVED, reviewed_by=self.user, ) finally: connections.close_all() with ThreadPoolExecutor(max_workers=2) as executor: futures = [executor.submit(approve) for _ in range(2)] for future in futures: future.result() balance = business_models.CustomerBalance.objects.get( merchant=self.merchant, customer=self.customer, ) self.assertEqual(balance.balance, self.sales_order.get_total_amount()) records = business_models.BalanceChangeRecord.objects.filter( merchant=self.merchant, source_type=business_models.BalanceChangeSourceEnum.SALES_ORDER, source_id=self.sales_order.id, ) self.assertEqual(records.count(), 1) self.assertEqual(records.first().balance_after, balance.balance) def tearDown(self): connections.close_all()