"""Optimistic-lock mixin.

Models inheriting :class:`VersionedModel` carry a monotonically-increasing
``version`` integer. ``save()`` uses ``UPDATE ... WHERE pk=? AND version=?``;
if zero rows are affected it means another writer beat us, and
:class:`simorgh.core.exceptions.ConcurrentUpdate` is raised.

Usage:

    obj = MyModel.objects.get(pk=42)
    obj.name = "x"
    obj.save()           # bumps version, raises if stale

Retry pattern (for update-then-save loops)::

    from simorgh.core.models.versioned import retry_on_conflict

    @retry_on_conflict(max_retries=3)
    def update_name(pk: int, new_name: str) -> None:
        obj = MyModel.objects.get(pk=pk)
        obj.name = new_name
        obj.save()
"""

from __future__ import annotations

import functools
import time
from typing import Callable, TypeVar

from django.db import models, router, transaction
from django.utils.translation import gettext_lazy as _

from simorgh.core.exceptions import ConcurrentUpdate

_F = TypeVar("_F", bound=Callable)


def retry_on_conflict(
    max_retries: int = 3,
    base_delay: float = 0.05,
    backoff: float = 2.0,
) -> Callable[[_F], _F]:
    """Decorator that retries the wrapped function on :exc:`ConcurrentUpdate`.

    Implements exponential back-off: first retry waits ``base_delay`` seconds,
    each subsequent retry doubles the wait up to ``max_retries`` total
    attempts.  If the last attempt still raises :exc:`ConcurrentUpdate`, it
    propagates to the caller.

    Parameters
    ----------
    max_retries:
        Total number of attempts (including the first).  Must be ≥ 1.
    base_delay:
        Seconds to wait before the first retry.
    backoff:
        Multiplier applied to ``base_delay`` after each failed attempt.
    """
    if max_retries < 1:
        raise ValueError("max_retries must be at least 1")

    def decorator(fn: _F) -> _F:
        @functools.wraps(fn)
        def wrapper(*args, **kwargs):
            delay = base_delay
            for attempt in range(max_retries):
                try:
                    return fn(*args, **kwargs)
                except ConcurrentUpdate:
                    if attempt == max_retries - 1:
                        raise
                    time.sleep(delay)
                    delay *= backoff

        return wrapper  # type: ignore[return-value]

    return decorator


class VersionedModel(models.Model):
    """Adds a ``version`` field and optimistic-lock-aware ``save``."""

    version = models.PositiveIntegerField(_("version"), default=0)

    class Meta:
        abstract = True

    def save(  # type: ignore[override]
        self,
        force_insert: bool = False,
        force_update: bool = False,
        using: str | None = None,
        update_fields: list[str] | tuple[str, ...] | None = None,
    ) -> None:
        if self._state.adding or force_insert:
            self.version = 1
            return super().save(
                force_insert=force_insert,
                force_update=force_update,
                using=using,
                update_fields=update_fields,
            )

        using = using or router.db_for_write(self.__class__, instance=self)
        current_version = self.version
        new_version = current_version + 1

        fields_to_update = self._resolve_update_fields(update_fields)
        fields_to_update.add("version")

        values = {name: getattr(self, name) for name in fields_to_update if name != "version"}
        values["version"] = new_version

        with transaction.atomic(using=using):
            updated = (
                type(self)
                ._base_manager.using(using)
                .filter(pk=self.pk, version=current_version)
                .update(**values)
            )
            if updated == 0:
                raise ConcurrentUpdate(
                    f"{type(self).__name__}(pk={self.pk}) was modified by another writer "
                    f"(expected version={current_version})."
                )
            self.version = new_version

    def _resolve_update_fields(
        self, update_fields: list[str] | tuple[str, ...] | None
    ) -> set[str]:
        if update_fields is not None:
            return set(update_fields)
        return {
            f.attname
            for f in self._meta.concrete_fields
            if not f.primary_key and not f.auto_created
        }


__all__ = ["VersionedModel", "retry_on_conflict"]
