Files
dealerhub/backend/tests/test_concurrency.py
T

434 lines
18 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""库存与过账的安全不变量测试(迭代深挖)。
验证"库存 = 钱"的核心约束:不超卖、不为负、过账幂等、失败整单回滚。
**关于真并发**:SQLite 用数据库级写锁(`database table is locked`),多线程不可靠;
生产是 PostgreSQL(行级锁)。因此分两层:
1. **顺序化的"超量请求"**(SQLite 可跑):覆盖同一业务约束;
2. **`@pytest.mark.postgres` 真并发**:默认在 SQLite 下跳过,连 PG 时执行。
PG 并发验证命令(本机 PG 在 5433):
python scripts/setup_pg.py # 一次性建库
DJANGO_SETTINGS_MODULE=config.settings.pgtest python -m pytest tests/ -p no:cacheprovider
"""
import threading
from datetime import date
from decimal import Decimal
import pytest
from model_bakery import baker
from apps.catalog.models import Product
from apps.inventory import services as inv_services
from apps.inventory.models import Stock, Warehouse
from apps.partner.models import Customer
from apps.sales import services as sales_services
from apps.sales.services import CreditLimitExceeded
@pytest.fixture
def warehouse(db, tenant):
return baker.make(Warehouse, tenant=tenant, code="WH01", name="主仓")
@pytest.fixture
def product(db, tenant):
return baker.make(Product, tenant=tenant, code="P001", name="可乐",
sale_price=Decimal("10"))
@pytest.fixture
def customer(db, tenant):
return baker.make(Customer, tenant=tenant, code="C001", name="张三",
credit_limit=Decimal("0")) # 不限额,专测库存
# ============================================================
# 库存原子性
# ============================================================
def test_no_oversell_sequential(db, tenant, warehouse, product):
"""连续超量出库必须被拒(第 2 次应抛 InsufficientStock)。"""
inv_services.inbound(tenant=tenant, warehouse=warehouse, product=product,
quantity=Decimal("10"), unit_cost=Decimal("5"))
inv_services.outbound(tenant=tenant, warehouse=warehouse, product=product,
quantity=Decimal("10"))
with pytest.raises(inv_services.InsufficientStock):
inv_services.outbound(tenant=tenant, warehouse=warehouse, product=product,
quantity=Decimal("1"))
def test_stock_never_negative_on_rejected_outbound(db, tenant, warehouse, product):
"""被拒的出库不得改动库存(事务回滚)。"""
inv_services.inbound(tenant=tenant, warehouse=warehouse, product=product,
quantity=Decimal("5"), unit_cost=Decimal("3"))
before = Stock.objects.get(product=product).on_hand
with pytest.raises(inv_services.InsufficientStock):
inv_services.outbound(tenant=tenant, warehouse=warehouse, product=product,
quantity=Decimal("100"))
after = Stock.objects.get(product=product).on_hand
assert before == after == Decimal("5")
@pytest.mark.postgres
def test_concurrent_outbound_respects_limit(transactional_db, tenant, warehouse, product):
"""并发出库:总量不得超过库存(超出的请求应失败,不产生负库存)。
用 `transactional_db`(真实提交)——普通 `db` fixture 把测试包在事务里,
多线程连接看不到彼此的数据,并发语义无从谈起。
"""
inv_services.inbound(tenant=tenant, warehouse=warehouse, product=product,
quantity=Decimal("20"), unit_cost=Decimal("2"))
results = {"ok": 0, "fail": 0, "errors": []}
lock = threading.Lock()
def worker():
try:
inv_services.outbound(tenant=tenant, warehouse=warehouse,
product=product, quantity=Decimal("6"))
with lock:
results["ok"] += 1
except Exception as exc:
with lock:
results["fail"] += 1
results["errors"].append(type(exc).__name__)
threads = [threading.Thread(target=worker) for _ in range(6)]
for t in threads:
t.start()
for t in threads:
t.join()
final = Stock.objects.get(product=product).on_hand
# 核心不变量:不超卖、不为负
assert final >= 0, f"库存为负:{final}"
assert results["ok"] <= 3, f"成功次数超过库存上限:{results}"
assert final == Decimal("20") - Decimal("6") * results["ok"]
print(f"并发出库:成功 {results['ok']} 次,失败 {results['fail']} 次"
f"({set(results['errors'])}),剩余库存 {final}")
@pytest.mark.postgres
def test_concurrent_inbound_sums_correctly(transactional_db, tenant, warehouse, product):
"""并发入库:总量必须等于各次之和(F() 自增不能丢更新)。"""
errors = []
lock = threading.Lock()
def worker():
try:
inv_services.inbound(tenant=tenant, warehouse=warehouse,
product=product, quantity=Decimal("7"),
unit_cost=Decimal("1"))
except Exception as exc:
with lock:
errors.append(f"{type(exc).__name__}: {exc}")
threads = [threading.Thread(target=worker) for _ in range(10)]
for t in threads:
t.start()
for t in threads:
t.join()
assert not errors, f"并发入库出现异常:{errors[:3]}"
final = Stock.objects.get(product=product).on_hand
assert final == Decimal("70"), f"入库丢更新:{final}(期望 70)"
@pytest.mark.postgres
def test_concurrent_oversell_blocked_across_products(transactional_db, tenant, warehouse):
"""并发场景下多商品同时抢库存:各自不超卖(验证行锁不互相阻塞到违约)。"""
prods = [
baker.make(Product, tenant=tenant, code=f"CP{i}", name=f"并发品{i}",
sale_price=Decimal("10"))
for i in range(3)
]
for p in prods:
inv_services.inbound(tenant=tenant, warehouse=warehouse, product=p,
quantity=Decimal("10"), unit_cost=Decimal("1"))
lock = threading.Lock()
ok = {"n": 0}
def worker(prod):
try:
inv_services.outbound(tenant=tenant, warehouse=warehouse,
product=prod, quantity=Decimal("4"))
with lock:
ok["n"] += 1
except Exception:
pass
threads = [threading.Thread(target=worker, args=(p,))
for p in prods for _ in range(4)]
for t in threads:
t.start()
for t in threads:
t.join()
for p in prods:
remain = Stock.objects.get(product=p).on_hand
assert remain >= 0
# 每个商品库存 10,每次出 4 → 最多成功 2 次
assert remain in (Decimal("10"), Decimal("6"), Decimal("2")), f"{p.code} 库存异常:{remain}"
def test_f_expression_no_lost_update_sequential(db, tenant, warehouse, product):
"""F() 自增语义:连续多次入库,总量必须精确累加(不丢更新)。
这是并发安全的基础——`update(on_hand=F("on_hand") + q)` 在数据库侧求值,
而非读-改-写,因此即便请求交错也不会覆盖彼此的增量。
"""
for _ in range(10):
inv_services.inbound(tenant=tenant, warehouse=warehouse, product=product,
quantity=Decimal("7"), unit_cost=Decimal("1"))
assert Stock.objects.get(product=product).on_hand == Decimal("70")
def test_available_check_uses_locked(db, tenant, warehouse, product):
"""可用量 = on_hand - locked;锁定部分不可出库。"""
inv_services.inbound(tenant=tenant, warehouse=warehouse, product=product,
quantity=Decimal("10"), unit_cost=Decimal("2"))
Stock.objects.filter(product=product).update(locked=Decimal("8"))
# 可用只剩 2,出 5 必须失败
with pytest.raises(inv_services.InsufficientStock):
inv_services.outbound(tenant=tenant, warehouse=warehouse, product=product,
quantity=Decimal("5"))
# 出 2 可以
inv_services.outbound(tenant=tenant, warehouse=warehouse, product=product,
quantity=Decimal("2"))
# ============================================================
# 单据过账幂等
# ============================================================
def test_confirm_twice_rejected(db, tenant, warehouse, product, customer):
"""同一张单不能过账两次(否则重复扣库存 + 重复生成应收)。"""
inv_services.inbound(tenant=tenant, warehouse=warehouse, product=product,
quantity=Decimal("50"), unit_cost=Decimal("2"))
bill = sales_services.create_sales_bill(
tenant=tenant, customer=customer, warehouse=warehouse,
lines=[{"product": product, "quantity": 5, "unit_price": 10}],
)
sales_services.confirm_sales_bill(bill)
assert Stock.objects.get(product=product).on_hand == Decimal("45")
with pytest.raises(ValueError):
sales_services.confirm_sales_bill(bill)
# 库存不能再被扣
assert Stock.objects.get(product=product).on_hand == Decimal("45")
def test_credit_reject_rolls_back_everything(db, tenant, warehouse, product):
"""信用超限拒绝后:库存、应收、状态全部不变(整事务回滚)。"""
from apps.finance.models import Receivable
limited = baker.make(Customer, tenant=tenant, code="C002", name="受限",
credit_limit=Decimal("100"))
inv_services.inbound(tenant=tenant, warehouse=warehouse, product=product,
quantity=Decimal("50"), unit_cost=Decimal("2"))
baker.make(Receivable, tenant=tenant, customer=limited, bill_no="RC-OLD",
total_amount=Decimal("80"), status="open")
bill = sales_services.create_sales_bill(
tenant=tenant, customer=limited, warehouse=warehouse,
lines=[{"product": product, "quantity": 5, "unit_price": 10}], # 50 → 超限
)
before_receivables = Receivable.objects.filter(tenant=tenant).count()
with pytest.raises(CreditLimitExceeded):
sales_services.confirm_sales_bill(bill)
assert Stock.objects.get(product=product).on_hand == Decimal("50") # 未扣
assert Receivable.objects.filter(tenant=tenant).count() == before_receivables
bill.refresh_from_db()
assert bill.state == "draft"
# ============================================================
# 库存不足时整单回滚(多行部分失败)
# ============================================================
def test_multi_line_partial_shortage_rolls_back(db, tenant, warehouse, customer):
"""多行单据中某一行库存不足 → 整单回滚,不能只扣成功的行。"""
p1 = baker.make(Product, tenant=tenant, code="P001", name="有货",
sale_price=Decimal("10"))
p2 = baker.make(Product, tenant=tenant, code="P002", name="缺货",
sale_price=Decimal("10"))
inv_services.inbound(tenant=tenant, warehouse=warehouse, product=p1,
quantity=Decimal("100"), unit_cost=Decimal("2"))
inv_services.inbound(tenant=tenant, warehouse=warehouse, product=p2,
quantity=Decimal("1"), unit_cost=Decimal("2"))
bill = sales_services.create_sales_bill(
tenant=tenant, customer=customer, warehouse=warehouse,
lines=[
{"product": p1, "quantity": 10, "unit_price": 10}, # 够
{"product": p2, "quantity": 50, "unit_price": 10}, # 不够
],
)
with pytest.raises(inv_services.InsufficientStock):
sales_services.confirm_sales_bill(bill)
assert Stock.objects.get(product=p1).on_hand == Decimal("100") # 未被扣
assert Stock.objects.get(product=p2).on_hand == Decimal("1")
bill.refresh_from_db()
assert bill.state == "draft"
def test_multi_line_all_success(db, tenant, warehouse, customer):
p1 = baker.make(Product, tenant=tenant, code="P001", name="A",
sale_price=Decimal("10"))
p2 = baker.make(Product, tenant=tenant, code="P002", name="B",
sale_price=Decimal("10"))
for p in (p1, p2):
inv_services.inbound(tenant=tenant, warehouse=warehouse, product=p,
quantity=Decimal("100"), unit_cost=Decimal("2"))
bill = sales_services.create_sales_bill(
tenant=tenant, customer=customer, warehouse=warehouse,
lines=[
{"product": p1, "quantity": 10, "unit_price": 10},
{"product": p2, "quantity": 20, "unit_price": 10},
],
)
sales_services.confirm_sales_bill(bill)
assert Stock.objects.get(product=p1).on_hand == Decimal("90")
assert Stock.objects.get(product=p2).on_hand == Decimal("80")
# ============================================================
# 采购入库的加权成本并发正确性
# ============================================================
def test_weighted_cost_after_sequential_inbounds(db, tenant, warehouse, product):
"""两次入库的加权平均成本正确(10@2 + 10@4 = 3)。"""
inv_services.inbound(tenant=tenant, warehouse=warehouse, product=product,
quantity=Decimal("10"), unit_cost=Decimal("2"))
inv_services.inbound(tenant=tenant, warehouse=warehouse, product=product,
quantity=Decimal("10"), unit_cost=Decimal("4"))
st = Stock.objects.get(product=product)
assert st.on_hand == Decimal("20")
assert st.avg_cost == Decimal("3.0000")
# ============================================================
# 单号生成竞态(迭代第 4 轮)
# ============================================================
@pytest.mark.postgres
def test_concurrent_bill_no_generation_unique(transactional_db, tenant, warehouse, customer):
"""并发建单:单号必须唯一(`count()+1` 取名方式存在竞态)。
背景:`_generate_bill_no` 用 `filter(...).count() + 1` 取序号——
两个事务同时 count 会拿到同一个数,生成相同单号,撞
`unique_together = [("tenant", "bill_no")]`。
"""
from apps.catalog.models import Product
from apps.sales.models import SalesBill
products = [
baker.make(Product, tenant=tenant, code=f"BN{i}", name=f"单号测试品{i}",
sale_price=Decimal("10"))
for i in range(3)
]
# 每个线程用自己的商品,避免库存相互影响,只测单号生成
errors = []
created_nos = []
lock = threading.Lock()
def worker(product):
try:
bill = sales_services.create_sales_bill(
tenant=tenant, customer=customer, warehouse=warehouse,
lines=[{"product": product, "quantity": 1, "unit_price": 10}],
)
with lock:
created_nos.append(bill.bill_no)
except Exception as exc:
with lock:
errors.append(f"{type(exc).__name__}: {str(exc)[:120]}")
threads = [threading.Thread(target=worker, args=(p,)) for p in products for _ in range(3)]
for t in threads:
t.start()
for t in threads:
t.join()
# 单号不得重复(这是核心不变量)
assert len(created_nos) == len(set(created_nos)), \
f"生成了重复单号:{[n for n in created_nos if created_nos.count(n) > 1]}"
# 不应有异常(并发下应全部成功或优雅重试)
assert not errors, f"并发建单出现异常:{errors[:3]}"
# 数据库里的单号也唯一
nos = list(SalesBill.objects.filter(tenant=tenant).values_list("bill_no", flat=True))
assert len(nos) == len(set(nos))
def test_bill_no_not_reused_after_deletion(db, tenant, warehouse, customer, product):
"""删除中间单据后,新单号不得复用已用过的号。
这是比并发更隐蔽的问题:旧实现用 `count()+1`,
建了 001/002 后删掉 001 → count=1 → 下一张又是 002 → 撞唯一约束。
改用 MAX(序号)+1 后不会。
"""
from apps.sales.models import SalesBill
from apps.inventory import services as inv
inv.inbound(tenant=tenant, warehouse=warehouse, product=product,
quantity=Decimal("100"), unit_cost=Decimal("2"))
bills = []
for _ in range(3):
b = sales_services.create_sales_bill(
tenant=tenant, customer=customer, warehouse=warehouse,
lines=[{"product": product, "quantity": 1, "unit_price": 10}],
)
bills.append(b)
nos = [b.bill_no for b in bills]
assert len(set(nos)) == 3, f"单号重复:{nos}"
assert nos[0].endswith("0001") and nos[2].endswith("0003")
# 删掉第一张
bills[0].lines.all().delete()
bills[0].delete()
# 新建一张:不得复用 0002/0003
b4 = sales_services.create_sales_bill(
tenant=tenant, customer=customer, warehouse=warehouse,
lines=[{"product": product, "quantity": 1, "unit_price": 10}],
)
assert b4.bill_no.endswith("0004"), f"删除后单号复用/回退:{b4.bill_no}"
existing = set(SalesBill.objects.filter(tenant=tenant).values_list("bill_no", flat=True))
assert len(existing) == 3 # 0002/0003/0004
def test_bill_no_survives_gap(db, tenant, warehouse, customer, product):
"""序号有跳号(历史删除造成)时仍取最大值 +1,不回退填空洞。"""
from apps.sales.models import SalesBill
# 手工造一个"序号很大"的单(模拟历史删除后的空洞)
from apps.inventory import services as inv
inv.inbound(tenant=tenant, warehouse=warehouse, product=product,
quantity=Decimal("100"), unit_cost=Decimal("2"))
today = date.today().strftime("%Y%m%d")
baker.make(SalesBill, tenant=tenant, bill_no=f"XS{today}0099",
customer=customer, warehouse=warehouse, bill_date=date.today(),
total_amount=Decimal("0"), state="draft")
b = sales_services.create_sales_bill(
tenant=tenant, customer=customer, warehouse=warehouse,
lines=[{"product": product, "quantity": 1, "unit_price": 10}],
)
assert b.bill_no.endswith("0100"), f"未取 MAX+1:{b.bill_no}"