"""库存原子化服务。 核心规则: 1. 所有库存变动必须走这里的函数,**禁止业务层直接 update Stock** 2. 每个变动包 transaction.atomic + select_for_update 锁行 3. 同步写 StockMovement 流水 4. raise InsufficientStock 当可用库存不足 """ from __future__ import annotations from dataclasses import dataclass, field from datetime import date from decimal import Decimal from django.db import IntegrityError, transaction from django.db.models import F from .models import Stock, StockBatch, StockMovement, Warehouse class InsufficientStock(Exception): """可用库存不足。""" pass class WarehouseNotFound(Exception): """仓库不存在。""" @dataclass class StockChangeResult: stock: Stock movement: StockMovement batch_allocations: list = field(default_factory=list) def _get_locked(tenant, warehouse, product): """select_for_update 取 stock;不存在则创建。返回刷新过的 stock。 **并发安全**:`select_for_update()` 对"不存在的行"无法上锁——两个事务同时 发现无记录时会各自 create,撞唯一约束(PG 实测复现)。 因此创建路径用 `get_or_create` 并捕获 IntegrityError 重取: 竞态失败方在对手提交后重查即可拿到同一行。 """ stock = ( Stock.objects.select_for_update() .filter(tenant=tenant, warehouse=warehouse, product=product) .first() ) if stock is not None: return stock # 创建路径:允许"被并发抢先"——用保存点隔离失败,避免整个外层事务被标记 aborted try: with transaction.atomic(): # 保存点:回滚只影响本段 stock, _created = Stock.objects.get_or_create( tenant=tenant, warehouse=warehouse, product=product, defaults={"on_hand": Decimal("0"), "locked": Decimal("0"), "avg_cost": Decimal("0")}, ) except IntegrityError: # 竞态:对手已插入并提交 → 重取(此时行一定存在) stock = ( Stock.objects.select_for_update() .filter(tenant=tenant, warehouse=warehouse, product=product) .first() ) if stock is None: # 对手回滚了 → 自己创建 stock = Stock.objects.create( tenant=tenant, warehouse=warehouse, product=product, on_hand=Decimal("0"), locked=Decimal("0"), avg_cost=Decimal("0"), ) return stock def get_or_create_stock(*, tenant, warehouse, product) -> Stock: """非事务环境下的便捷版(用于初始化)。""" stock, _ = Stock.objects.get_or_create( tenant=tenant, warehouse=warehouse, product=product, defaults={"on_hand": Decimal("0"), "locked": Decimal("0"), "avg_cost": Decimal("0")}, ) return stock def _get_locked_batch(tenant, warehouse, product, batch_no, defaults: dict): """select_for_update 取批次账;不存在则按 defaults 创建。 并发安全同 `_get_locked`:创建路径用 get_or_create + IntegrityError 重取, 避免"首次并发入库同一批次"撞唯一约束。 """ batch = ( StockBatch.objects.select_for_update() .filter(tenant=tenant, warehouse=warehouse, product=product, batch_no=batch_no) .first() ) if batch is None: try: with transaction.atomic(): # 保存点:隔离并发冲突 batch, _created = StockBatch.objects.get_or_create( tenant=tenant, warehouse=warehouse, product=product, batch_no=batch_no, defaults={ "production_date": defaults.get("production_date"), "expiry_date": defaults.get("expiry_date"), "on_hand": Decimal("0"), "locked": Decimal("0"), "unit_cost": defaults.get("unit_cost", Decimal("0")), }, ) except IntegrityError: batch = ( StockBatch.objects.select_for_update() .filter(tenant=tenant, warehouse=warehouse, product=product, batch_no=batch_no) .first() ) if batch is None: batch = StockBatch.objects.create( tenant=tenant, warehouse=warehouse, product=product, batch_no=batch_no, production_date=defaults.get("production_date"), expiry_date=defaults.get("expiry_date"), on_hand=Decimal("0"), locked=Decimal("0"), unit_cost=defaults.get("unit_cost", Decimal("0")), ) else: # 批次信息允许补录(生产日期/到期日为空时回填) fill = {} if batch.production_date is None and defaults.get("production_date"): fill["production_date"] = defaults["production_date"] if batch.expiry_date is None and defaults.get("expiry_date"): fill["expiry_date"] = defaults["expiry_date"] if fill: StockBatch.objects.filter(pk=batch.pk).update(**fill) batch.refresh_from_db() return batch def inbound( *, tenant, warehouse, product, quantity: Decimal, unit_cost: Decimal = Decimal("0"), source_type: str = "manual", source_ref: str = "", batch_no: str = "", production_date=None, expiry_date=None, ) -> StockChangeResult: """入库:on_hand += quantity;写流水 + 加权平均成本。 批次管理商品(product.is_batch_managed)必须带 batch_no,同时维护批次账。 """ if quantity <= 0: raise ValueError("quantity must be positive") is_batch = getattr(product, "is_batch_managed", False) if is_batch and not batch_no: raise ValueError(f"product {product.code} is batch-managed: batch_no required") with transaction.atomic(): stock = _get_locked(tenant, warehouse, product) # 用 F() 表达式做增量 update,避免 race condition Stock.objects.filter(pk=stock.pk).update(on_hand=F("on_hand") + quantity) # 加权平均成本 if unit_cost and unit_cost > 0: old_qty = stock.on_hand old_cost = stock.avg_cost new_cost = ( (old_cost * old_qty + unit_cost * quantity) / (old_qty + quantity) if (old_qty + quantity) > 0 else unit_cost ) Stock.objects.filter(pk=stock.pk).update(avg_cost=new_cost) # 刷新 in-memory stock.refresh_from_db() movement = StockMovement.objects.create( tenant=tenant, warehouse=warehouse, product=product, movement_type=StockMovement.MOVEMENT_INBOUND, source_type=source_type, source_ref=source_ref, quantity=quantity, unit_cost=unit_cost, ) batch_allocations = [] if is_batch: if expiry_date is None and product.shelf_life_days: base = production_date or date.today() expiry_date = base + __import__("datetime").timedelta( days=product.shelf_life_days ) batch = _get_locked_batch( tenant, warehouse, product, batch_no, defaults={ "production_date": production_date, "expiry_date": expiry_date, "unit_cost": unit_cost, }, ) StockBatch.objects.filter(pk=batch.pk).update( on_hand=F("on_hand") + quantity, unit_cost=unit_cost if unit_cost > 0 else batch.unit_cost, ) batch.refresh_from_db() batch_allocations.append({"batch_no": batch_no, "quantity": str(quantity)}) return StockChangeResult( stock=stock, movement=movement, batch_allocations=batch_allocations ) def _fefo_allocate(tenant, warehouse, product, quantity: Decimal) -> list: """批次商品出库分摊:FEFO(近效期优先,空到期日最后,再按先入先出 id)。 返回 [{batch_no, quantity, expiry_date}];总量不足抛 InsufficientStock。 调用方须处于事务中。 """ batches = list( StockBatch.objects.select_for_update() .filter(tenant=tenant, warehouse=warehouse, product=product) .order_by(F("expiry_date").asc(nulls_last=True), "id") ) remaining = quantity allocations = [] for b in batches: if remaining <= 0: break available = b.on_hand - b.locked if available <= 0: continue take = min(available, remaining) allocations.append({"batch": b, "take": take}) remaining -= take if remaining > 0: total_available = sum((b.on_hand - b.locked) for b in batches) raise InsufficientStock( f"insufficient batch stock for {warehouse.code}/{product.code}: " f"available={total_available}, requested={quantity}" ) return allocations def outbound( *, tenant, warehouse, product, quantity: Decimal, source_type: str = "manual", source_ref: str = "", ) -> StockChangeResult: """出库:available = on_hand - locked 必须 >= quantity;on_hand -= quantity。 批次管理商品自动按 FEFO 分摊到批次账(近效期优先)。 """ if quantity <= 0: raise ValueError("quantity must be positive") is_batch = getattr(product, "is_batch_managed", False) with transaction.atomic(): stock = _get_locked(tenant, warehouse, product) available = stock.on_hand - stock.locked if available < quantity: raise InsufficientStock( f"insufficient stock for {warehouse.code}/{product.code}: " f"available={available}, requested={quantity}" ) batch_detail = [] if is_batch: for alloc in _fefo_allocate(tenant, warehouse, product, quantity): b = alloc["batch"] take = alloc["take"] StockBatch.objects.filter(pk=b.pk).update(on_hand=F("on_hand") - take) batch_detail.append( { "batch_no": b.batch_no, "quantity": str(take), "expiry_date": b.expiry_date.isoformat() if b.expiry_date else None, } ) Stock.objects.filter(pk=stock.pk).update(on_hand=F("on_hand") - quantity) stock.refresh_from_db() movement = StockMovement.objects.create( tenant=tenant, warehouse=warehouse, product=product, movement_type=StockMovement.MOVEMENT_OUTBOUND, source_type=source_type, source_ref=source_ref, quantity=quantity, unit_cost=stock.avg_cost, batch_detail=batch_detail, ) return StockChangeResult( stock=stock, movement=movement, batch_allocations=batch_detail ) def lock( *, tenant, warehouse, product, quantity: Decimal, source_type: str = "manual", source_ref: str = "", ) -> StockChangeResult: """锁定:locked += quantity(不直接出库)。""" if quantity <= 0: raise ValueError("quantity must be positive") with transaction.atomic(): stock = _get_locked(tenant, warehouse, product) available = stock.on_hand - stock.locked if available < quantity: raise InsufficientStock( f"insufficient available stock for {warehouse.code}/{product.code}: " f"available={available}, requested={quantity}" ) Stock.objects.filter(pk=stock.pk).update(locked=F("locked") + quantity) stock.refresh_from_db() movement = StockMovement.objects.create( tenant=tenant, warehouse=warehouse, product=product, movement_type=StockMovement.MOVEMENT_LOCK, source_type=source_type, source_ref=source_ref, quantity=quantity, ) return StockChangeResult(stock=stock, movement=movement) def unlock( *, tenant, warehouse, product, quantity: Decimal, source_type: str = "manual", source_ref: str = "", ) -> StockChangeResult: """释放锁定:locked -= quantity。""" if quantity <= 0: raise ValueError("quantity must be positive") with transaction.atomic(): stock = _get_locked(tenant, warehouse, product) if stock.locked < quantity: raise InsufficientStock( f"locked less than unlock for {warehouse.code}/{product.code}: " f"locked={stock.locked}, requested={quantity}" ) Stock.objects.filter(pk=stock.pk).update(locked=F("locked") - quantity) stock.refresh_from_db() movement = StockMovement.objects.create( tenant=tenant, warehouse=warehouse, product=product, movement_type=StockMovement.MOVEMENT_UNLOCK, source_type=source_type, source_ref=source_ref, quantity=quantity, ) return StockChangeResult(stock=stock, movement=movement) def adjust( *, tenant, warehouse, product, new_on_hand: Decimal, source_type: str = "adjust", source_ref: str = "", ) -> StockChangeResult: """盘点调整:把 on_hand 改成 new_on_hand,差值记一条 adjust 流水。""" with transaction.atomic(): stock = _get_locked(tenant, warehouse, product) diff = new_on_hand - stock.on_hand Stock.objects.filter(pk=stock.pk).update(on_hand=new_on_hand) stock.refresh_from_db() movement = StockMovement.objects.create( tenant=tenant, warehouse=warehouse, product=product, movement_type=StockMovement.MOVEMENT_ADJUST, source_type=source_type, source_ref=source_ref, quantity=abs(diff), unit_cost=stock.avg_cost, ) return StockChangeResult(stock=stock, movement=movement)