434 lines
18 KiB
Python
434 lines
18 KiB
Python
"""库存与过账的安全不变量测试(迭代深挖)。
|
||
|
||
验证"库存 = 钱"的核心约束:不超卖、不为负、过账幂等、失败整单回滚。
|
||
|
||
**关于真并发**: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}"
|