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