"""Adjustments and count reconciliation always use the immutable ledger writer."""
import uuid
from contextlib import contextmanager
from django.core.exceptions import ValidationError
from django.db.models import Sum
from django.utils import timezone
from apps.access.authorization import require_permission
from apps.parts.locking import persisted_pk
from .locking import inventory_context, lock_positions
from .models import StockAdjustment, StockCount, StockCountUnit, StockLedgerEntry, StockReservation, SerializedStockUnit
from .queries import _on_hand
from .services import _post, check_revision, quantity_value, text_value


def require_position_open(location, part, *, count=None):
    rows = StockCount.objects.filter(location=location, spare_part=part, status="COUNTING")
    if count is not None:
        rows = rows.exclude(pk=count.pk)
    if rows.exists():
        raise ValidationError("Physical counting currently freezes this stock position.")


def position_revision(location, part):
    ledger = StockLedgerEntry.objects.filter(location=location, spare_part=part)
    reservations = StockReservation.objects.filter(location=location, spare_part=part)
    active = reservations.filter(status="ACTIVE").aggregate(value=Sum("quantity"))["value"] or 0
    return f"{location.updated_at.isoformat()}:{part.updated_at.isoformat()}:{ledger.count()}:{reservations.count()}:{active}"


def _adjust(*, actor, company, location, part, delta, reason, command_key, reference, note, units, count=None):
    if type(delta) is not int or delta == 0 or abs(delta)>1000000000:
        raise ValidationError("Supply a nonzero bounded whole-number variance.")
    if reason not in StockAdjustment.Reason.values or (reason == "LOSS" and delta>0) or (reason == "FOUND" and delta<0):
        raise ValidationError("Choose a reason consistent with the adjustment direction.")
    if (reason == "COUNT_VARIANCE") != (count is not None):
        raise ValidationError("Count variances require the reconciliation workflow.")
    if location.location_type in ("TRANSIT", "CUSTODY"):
        raise ValidationError("Managed workflow stock cannot be adjusted directly.")
    require_position_open(location, part, count=count)
    if delta > 0:
        ids = [persisted_pk(unit, SerializedStockUnit) for unit in units]
        if len(ids) != len(set(ids)):
            raise ValidationError("Supply distinct serialized units.")
        units = list(SerializedStockUnit.objects.filter(pk__in=ids).select_for_update(of=("self",)).select_related("current_movement__source__service_center").order_by("pk"))
        if len(units) != len(ids):
            raise ValidationError("A selected unit no longer exists.")
        for unit in units:
            if unit.state == "REGISTERED":
                require_permission(user=actor, permission="inventory.adjust_stock", target=company)
            elif unit.state == "REMOVED":
                origin = unit.current_movement.source
                require_permission(user=actor, permission="inventory.adjust_stock", target=origin.service_center or company)
    movement = _post(actor=actor, company=company, spare_part=part, source=location if delta<0 else None,
        destination=location if delta>0 else None, quantity=abs(delta), reference=reference, note=note,
        idempotency_key=command_key, units=units, kind="ADJUST_IN" if delta>0 else "ADJUST_OUT", count=count)
    row = StockAdjustment(company=company, location=location, spare_part=part, quantity_delta=delta,
        reason=reason, actor=actor, movement=movement, count=count)
    row._persist()
    return row


def adjust_stock(*, actor, location, spare_part, quantity_delta, reason, command_key, reference, note, units=(), expected_revision=None):
    with inventory_context(actor=actor, company=location.company, permission="inventory.adjust_stock", locations=[location], parts=[spare_part]) as (actor, company, locations, parts):
        location, part = locations[location.pk], parts[spare_part.pk]
        lock_positions([location], [part])
        if expected_revision is not None and position_revision(location, part) != expected_revision:
            raise ValidationError("Stock changed; reload and review the adjustment.")
        return _adjust(actor=actor, company=company, location=location, part=part, delta=quantity_delta,
            reason=reason, command_key=command_key, reference=text_value(reference, required=True, maximum=128),
            note=text_value(note, required=True), units=units)


def create_stock_count(*, actor, location, spare_part, note=""):
    with inventory_context(actor=actor, company=location.company, permission="inventory.count_stock", locations=[location], parts=[spare_part]) as (actor, company, locations, parts):
        location = locations[location.pk]
        if not location.is_active or location.location_type in ("TRANSIT", "CUSTODY"):
            raise ValidationError("Counts require an active physical stock location.")
        row = StockCount(company=company, location=location, spare_part=parts[spare_part.pk], created_by=actor, note=text_value(note))
        row._persist()
        return row


@contextmanager
def _count_context(actor, count, expected_revision, *, reconcile=False, require_active_parts=True):
    snapshot = StockCount.objects.select_related("company", "location", "spare_part").get(pk=persisted_pk(count, StockCount))
    with inventory_context(actor=actor, company=snapshot.company, permission="inventory.count_stock",
            locations=[snapshot.location], parts=[snapshot.spare_part], require_active_parts=require_active_parts) as (actor, company, locations, parts):
        location, part = locations[snapshot.location_id], parts[snapshot.spare_part_id]
        if reconcile:
            require_permission(user=actor, permission="inventory.adjust_stock", target=location.service_center or company)
        lock_positions([location], [part])
        row = StockCount.objects.select_for_update().get(pk=snapshot.pk)
        check_revision(row, expected_revision)
        yield actor, company, location, part, row


def start_stock_count(*, actor, count, expected_revision):
    with _count_context(actor, count, expected_revision) as (actor, _, location, part, row):
        if row.status != "DRAFT" or not location.is_active:
            raise ValidationError("Only a draft at an active location can start counting.")
        require_position_open(location, part)
        row.expected_quantity = _on_hand(location, part)
        if not 0 <= row.expected_quantity <= 1000000000:
            raise ValidationError("Stock position requires investigation before counting.")
        row.status, row.started_at, row.started_by = "COUNTING", timezone.now(), actor
        row._persist()
        for unit in SerializedStockUnit.objects.filter(current_location=location, spare_part=part).select_for_update().order_by("pk"):
            StockCountUnit(count=row, unit=unit, expected=True)._persist()
        return row


def record_stock_count(*, actor, count, counted_quantity, units=(), note="", expected_revision):
    if type(counted_quantity) is not int or not 0 <= counted_quantity <= 1000000000:
        raise ValidationError("Count must be a nonnegative bounded whole number.")
    ids = [persisted_pk(unit, SerializedStockUnit) for unit in units]
    if len(ids) != len(set(ids)) or len(ids)>counted_quantity:
        raise ValidationError("Supply distinct counted units within the quantity.")
    with _count_context(actor, count, expected_revision) as (actor, company, location, part, row):
        if row.status != "COUNTING":
            raise ValidationError("Only an active count accepts observations.")
        if (part.serialization_policy == "REQUIRED_SERIAL" and len(ids)!=counted_quantity) or (part.serialization_policy == "NOT_SERIALIZED" and ids):
            raise ValidationError("Counted units must satisfy the current serialization policy.")
        selected = list(SerializedStockUnit.objects.filter(pk__in=ids).select_for_update().order_by("pk"))
        if len(selected)!=len(ids) or any(unit.company_id!=company.pk or unit.spare_part_id!=part.pk or not (
                (unit.state=="IN_STOCK" and unit.current_location_id==location.pk) or (unit.state in ("REGISTERED", "REMOVED") and unit.current_location_id is None)) for unit in selected):
            raise ValidationError("A counted unit belongs elsewhere or cannot be reintroduced.")
        existing = {link.unit_id: link for link in row.units.select_for_update().order_by("pk")}
        for unit_id, link in existing.items():
            link.counted = unit_id in ids
            link._persist()
        for unit in selected:
            if unit.pk not in existing:
                StockCountUnit(count=row, unit=unit, counted=True)._persist()
        row.counted_quantity, row.note = counted_quantity, text_value(note)
        row._persist()
        return row


def cancel_stock_count(*, actor, count, reason, expected_revision):
    with _count_context(actor, count, expected_revision, require_active_parts=False) as (actor, _, _, _, row):
        if row.status not in ("DRAFT", "COUNTING"):
            raise ValidationError("A terminal count cannot be cancelled.")
        row.status, row.finished_at, row.finished_by = "CANCELLED", timezone.now(), actor
        row.reason = text_value(reason, required=True, maximum=500)
        row._persist()
        return row


def reconcile_stock_count(*, actor, count, expected_revision):
    with _count_context(actor, count, expected_revision, reconcile=True) as (actor, company, location, part, row):
        if row.status != "COUNTING" or row.counted_quantity is None:
            raise ValidationError("Record a physical count before reconciliation.")
        if _on_hand(location, part) != row.expected_quantity:
            raise ValidationError("Stock changed despite the count freeze; investigate before reconciliation.")
        links = list(row.units.select_for_update().order_by("pk"))
        expected = {link.unit_id for link in links if link.expected}
        observed = {link.unit_id for link in links if link.counted}
        units = {unit.pk: unit for unit in SerializedStockUnit.objects.filter(pk__in=expected | observed).select_for_update().order_by("pk")}
        missing, found = expected-observed, observed-expected
        anonymous = row.counted_quantity-len(observed)-(row.expected_quantity-len(expected))
        for label, delta, selection in (("missing", -len(missing), missing), ("found", len(found), found), ("anonymous", anonymous, set())):
            if delta:
                _adjust(actor=actor, company=company, location=location, part=part, delta=delta, reason="COUNT_VARIANCE", count=row,
                    command_key=uuid.uuid5(row.pk, label), reference=f"COUNT-{row.pk}", note=row.note or "Approved physical count variance.",
                    units=[units[pk] for pk in selection])
        row.status, row.finished_at, row.finished_by = "RECONCILED", timezone.now(), actor
        row._persist()
        return row
