forked from erp-dev/erp
77 lines
2.9 KiB
Python
77 lines
2.9 KiB
Python
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()
|