From 9478c25a70696feddafbff139e6f41d3b0abe78d Mon Sep 17 00:00:00 2001 From: Matthias Mair Date: Tue, 29 Sep 2026 13:26:53 +0200 Subject: [PATCH] feat(backend): Add generic API helper for transitions (#12693) * add generic api helper * align name better * fix name change * allign returncodes and internal names * fix serializer class assumptions * remove now unneeded serializers * fix parameter mixin handeling * fix context parsing * bump API version * add transition overview * fix serializer augumentation * remove dead code * change path prefix back * extend test for introspection * add legacy warning again * check permissions for introspection endpoint (allowed indirect discovery of state) * add sample * helpers for complex and well described endpoints * add more testing for state features shipped with common * fix tests * remove problematic many2many that does not work without database model * remove samples; will be submitted as a larger sample app * Update api_helpers.py --- src/backend/InvenTree/InvenTree/api.py | 15 +- .../InvenTree/InvenTree/api_version.py | 6 +- src/backend/InvenTree/InvenTree/mixins.py | 9 +- .../InvenTree/generic/states/__init__.py | 4 + src/backend/InvenTree/generic/states/api.py | 1 + .../InvenTree/generic/states/api_helpers.py | 411 ++++++++++++++++++ .../InvenTree/generic/states/introspection.py | 104 +++++ .../InvenTree/generic/states/serializers.py | 11 + src/backend/InvenTree/generic/states/tests.py | 6 +- .../InvenTree/generic/states/transition.py | 40 +- src/backend/InvenTree/order/api.py | 100 +---- src/backend/InvenTree/order/models.py | 2 +- src/backend/InvenTree/order/serializers.py | 31 -- src/backend/InvenTree/order/test_api.py | 26 ++ 14 files changed, 633 insertions(+), 133 deletions(-) create mode 100644 src/backend/InvenTree/generic/states/api_helpers.py create mode 100644 src/backend/InvenTree/generic/states/introspection.py diff --git a/src/backend/InvenTree/InvenTree/api.py b/src/backend/InvenTree/InvenTree/api.py index 6a30d8a722..7cf354f254 100644 --- a/src/backend/InvenTree/InvenTree/api.py +++ b/src/backend/InvenTree/InvenTree/api.py @@ -647,13 +647,18 @@ class ParameterListMixin: serializer_class = ( getattr(self, 'serializer_class', None) or self.get_serializer_class() ) - - model_class = serializer_class.Meta.model + model_class = getattr(getattr(serializer_class, 'Meta', None), 'model', None) # Apply ordering based on query parameter - queryset = common.filters.order_by_parameter( - queryset, model_class, self.request.query_params.get('ordering', None) - ) + if model_class is not None: + ordering = ( + self.request.query_params.get('ordering', None) + if hasattr(self, 'request') + else {} + ) + queryset = common.filters.order_by_parameter( + queryset, model_class, ordering + ) return queryset diff --git a/src/backend/InvenTree/InvenTree/api_version.py b/src/backend/InvenTree/InvenTree/api_version.py index a37cfc736e..f9e82c49a1 100644 --- a/src/backend/InvenTree/InvenTree/api_version.py +++ b/src/backend/InvenTree/InvenTree/api_version.py @@ -1,11 +1,15 @@ """InvenTree API version information.""" # InvenTree API version -INVENTREE_API_VERSION = 550 +INVENTREE_API_VERSION = 551 """Increment this API version number whenever there is a significant change to the API that any clients need to know about.""" INVENTREE_API_TEXT = """ +v551 -> 2026-09-21 : https://github.com/inventree/InvenTree/pull/12693 + - Docstring updates for automated transition API documentation generation + - New transitions insights endpoints + v550 -> 2026-09-22 : https://github.com/inventree/InvenTree/pull/12912 - Adds a top-level 'batch_code' field to the PurchaseOrderReceive API endpoint diff --git a/src/backend/InvenTree/InvenTree/mixins.py b/src/backend/InvenTree/InvenTree/mixins.py index 8240ae353b..92aeebe3f4 100644 --- a/src/backend/InvenTree/InvenTree/mixins.py +++ b/src/backend/InvenTree/InvenTree/mixins.py @@ -194,8 +194,15 @@ class OutputOptionsMixin: def get_serializer(self, *args, **kwargs): """Return serializer instance with output options applied.""" request = getattr(self, 'request', None) + serializer_class = self.get_serializer_class() - if self.output_options and request: + # Only inject output-option flags into serializers that know how to consume them. + # Action-specific serializers (e.g. transition endpoints) may not accept these kwargs. + if ( + self.output_options + and request + and issubclass(serializer_class, FilterableSerializerMixin) + ): params = self.request.query_params kwargs.update(self.output_options.format_params(params)) diff --git a/src/backend/InvenTree/generic/states/__init__.py b/src/backend/InvenTree/generic/states/__init__.py index 9a827a040b..bed9f86a87 100644 --- a/src/backend/InvenTree/generic/states/__init__.py +++ b/src/backend/InvenTree/generic/states/__init__.py @@ -15,6 +15,8 @@ from .transition import ( DEFERRABLE, StateTransitionMixin, TransitionMethod, + after_commit, + blocking_reason, inventree_transition, ) @@ -27,6 +29,8 @@ __all__ = [ 'StatusCode', 'StatusCodeMixin', 'TransitionMethod', + 'after_commit', + 'blocking_reason', 'can_proceed', # django_fsm import 'deprecated', 'fields', diff --git a/src/backend/InvenTree/generic/states/api.py b/src/backend/InvenTree/generic/states/api.py index 4047c69161..4d5544b20e 100644 --- a/src/backend/InvenTree/generic/states/api.py +++ b/src/backend/InvenTree/generic/states/api.py @@ -19,6 +19,7 @@ from InvenTree.helpers import inheritors from InvenTree.mixins import ListCreateAPI, RetrieveUpdateDestroyAPI from InvenTree.serializers import EmptySerializer +from .api_helpers import FSMTransitionMixin # noqa # pylint: disable=unused-import from .serializers import GenericStateClassSerializer from .states import StatusCode diff --git a/src/backend/InvenTree/generic/states/api_helpers.py b/src/backend/InvenTree/generic/states/api_helpers.py new file mode 100644 index 0000000000..158d64ce88 --- /dev/null +++ b/src/backend/InvenTree/generic/states/api_helpers.py @@ -0,0 +1,411 @@ +"""DRF viewset helpers to add transition endpoints.""" + +from __future__ import annotations + +import inspect +from typing import Any + +from django.core.exceptions import ValidationError as DjangoValidationError +from django.utils.translation import gettext_lazy as _ + +from drf_spectacular.utils import extend_schema +from rest_framework import serializers, status +from rest_framework.decorators import action +from rest_framework.exceptions import PermissionDenied +from rest_framework.request import Request +from rest_framework.response import Response + +from InvenTree.serializers import EmptySerializer +from users.permissions import check_user_permission + +from .introspection import ( + TransitionInfo, + available_transitions, + state_fields, + transition_methods, +) +from .serializers import AvailableTransitionSerializer + +TRANSITION_URL_PREFIX = '_transitions' + + +def current_state(instance: Any, field_name: str | None = None) -> str | None: + """Return the value of an instance's FSM state field. + + Args: + instance: Any model instance, whether or not it carries an FSM field. + field_name: Which state field to read. Defaults to the model's first FSM field. + + Returns: + The current state value, or None when the instance has no such field. + """ + if field_name is None: + names = state_fields(type(instance)) + if not names: + return None + field_name = names[0] + + return getattr(instance, field_name, None) + + +def transition_error( + instance: Any, transition_name: str, exc: Exception +) -> dict[str, Any]: + """Build the error body for a refused transition. + + Args: + instance: The instance the transition was attempted on, still in its source state. + transition_name: Name of the transition function that refused. + exc: The ``ValidationError`` the FSM raised. + + Returns: + A JSON-safe body with ``detail``, ``transition``, ``state`` and ``blocking_reason``. + """ + meta = transition_methods(type(instance)).get(transition_name) + field = getattr(meta, 'field', None) if meta is not None else None + field_name = field if isinstance(field, str) else getattr(field, 'name', None) + + state = current_state(instance, field_name) + info = next( + ( + entry + for entry in available_transitions(instance) + if entry.name == transition_name + ), + None, + ) + + messages = getattr(exc, 'messages', None) or [str(exc)] + detail = '; '.join(str(message) for message in messages) + reason: str | None = None + + if info is None: + detail = ( + f"Transition '{transition_name}' is not available from state '{state}'." + ) + elif info.blocked: + reason = info.blocking_reason + detail = reason or detail + + return { + 'detail': detail, + 'transition': transition_name, + 'state': state, + 'blocking_reason': reason, + } + + +def transition_action( + method_name: str, + *, + name: str | None = None, + url_path: str | None = None, + return_code: int = status.HTTP_200_OK, + pass_user: bool = False, + serializer_class: type[serializers.Serializer] | None = None, + reset_output_options: bool = False, +) -> Any: + """Build a ``POST`` detail endpoint running one FSM transition method. + + ``inventree_transition`` runs the transition in ``transaction.atomic`` and saves the instance, + so nothing here saves. A refused transition raises ``ValidationError``, returned as ``400``. + + Args: + method_name: Name of the transition method on the model. + name: Attribute name on the viewset. Defaults to ``method_name``. + url_path: Path segment under ``_transition/``. Defaults to ``name``, hyphenated. + return_code: HTTP status code to return on success. Defaults to `200`. + pass_user: Pass the requesting user to the transition as ``user=``. + serializer_class: Serializer to use for the response body. Defaults to an empty serializer + reset_output_options: Force ``output_options=None`` for this action. Only valid when the + viewset actually mixes in ``OutputOptionsMixin`` - passing it to a plain ``as_view()`` + that has no such attribute raises ``TypeError``. + + Returns: + A DRF ``@action``-decorated method, ready to assign as a viewset class attribute. + """ + attr_name = name or method_name + segment = url_path or attr_name.replace('_', '-') + # TODO @matmair move all transition actions under the common prefix + # path = f'{TRANSITION_URL_PREFIX}/{segment}' + path = f'{segment}' + input_serializer_class = serializer_class + + def endpoint(self, request: Request, pk: str | None = None) -> Response: + instance = self.get_object() + kwargs: dict[str, Any] = {} + context = self.get_serializer_context() + + if input_serializer_class is not None: + payload = input_serializer_class(data=request.data, context=context) + payload.is_valid(raise_exception=True) + extract = getattr(payload, 'transition_kwargs', None) + kwargs.update(extract() if extract else dict(payload.validated_data)) + + if pass_user: + user = getattr(request, 'user', None) + kwargs['user'] = ( + user if user is not None and user.is_authenticated else None + ) + + try: + getattr(instance, method_name)(**kwargs) + except DjangoValidationError as exc: + return Response( + transition_error(instance, method_name, exc), + status=status.HTTP_400_BAD_REQUEST, + ) + + instance.refresh_from_db() + + serial = ( + input_serializer_class + if input_serializer_class + else self.get_serializer_class() + ) + serializer = serial(instance, context=self.get_serializer_context()) + return Response(serializer.data, status=return_code) + + endpoint.__name__ = attr_name + endpoint.__qualname__ = attr_name + endpoint.__doc__ = f"API endpoint to '{method_name}' the current item." + # Read back by FSMTransitionMixin.transition_url_paths(). + endpoint.transition_name = method_name + + action_kwargs: dict[str, Any] = { + 'detail': True, + 'methods': ['post'], + 'url_path': path, + # TODO @matmair add option to rename the urlname + 'url_name': segment, + } + if input_serializer_class is not None: + action_kwargs['serializer_class'] = input_serializer_class + if reset_output_options: + action_kwargs['output_options'] = None + ret = action(**action_kwargs)(endpoint) + + # add decorator if custom return_code is required + if return_code != status.HTTP_200_OK: + ret = extend_schema( + responses={ + return_code: input_serializer_class + if input_serializer_class is not None + else EmptySerializer + } + )(ret) + + return ret + + +def transition_call_plan(func: Any) -> tuple[bool, tuple[str, ...]]: + """Report how a transition method must be called, from its signature. + + ``inventree_transition`` wraps the method with ``functools.wraps``, so the decorated method's + own parameters are still visible. + + Args: + func: The transition method, as read off the model class. + + Returns: + ``(pass_user, required)`` — whether the method takes ``user``, and the names of any other + parameters it requires. + """ + try: + signature = inspect.signature(func) + except (TypeError, ValueError): # pragma: no cover + return False, () + + pass_user = False + required: list[str] = [] + + for index, (param_name, param) in enumerate(signature.parameters.items()): + if index == 0 and param_name in ('self', 'cls'): + continue + if param.kind in (param.VAR_POSITIONAL, param.VAR_KEYWORD): + continue + if param_name == 'user': + pass_user = True + continue + if param.default is param.empty: + required.append(param_name) + + return pass_user, tuple(required) + + +class FSMTransitionMixin: + """Viewset mixin for exposing all FSM transitions on a model.""" + + #: Per-transition overrides keyed by model method name + transition_options: dict[str, dict[str, Any]] = {} + + #: Model method names to leave unexposed, including ones inherited from a base viewset. + transition_exclude: tuple[str, ...] = () + + #: Set False to stop this class generating endpoints; inherited ones remain. + autodiscover_transitions: bool = True + + #: Model method name -> viewset attribute, filled in by :meth:`register_transition_actions`. + generated_transition_actions: dict[str, str] = {} + + #: Model method name -> why no endpoint was generated for it. + skipped_transitions: dict[str, str] = {} + + def __init_subclass__(cls, **kwargs: Any) -> None: + """Generate the transition endpoints when a viewset class is defined. + + Args: + **kwargs: Class keyword arguments, forwarded up the MRO. + """ + super().__init_subclass__(**kwargs) + cls.register_transition_actions() + + @classmethod + def get_transition_model(cls) -> type | None: + """Return the model whose transitions this viewset exposes. + + Returns: + The model from ``queryset``, else from ``serializer_class.Meta.model``, else None. + """ + queryset = getattr(cls, 'queryset', None) + if queryset is not None: + return queryset.model + + meta = getattr(getattr(cls, 'serializer_class', None), 'Meta', None) + return getattr(meta, 'model', None) + + @classmethod + def register_transition_actions(cls) -> dict[str, str]: + """Attach one endpoint to this class per transition on its model. + + Idempotent: endpoints generated earlier are regenerated, other attributes are untouched. + + Returns: + Mapping of model method name to the viewset attribute holding its endpoint. + """ + generated: dict[str, str] = {} + skipped: dict[str, str] = {} + cls.generated_transition_actions = generated + cls.skipped_transitions = skipped + + model = cls.get_transition_model() + if model is None or not cls.autodiscover_transitions: + return generated + + for method_name in sorted(transition_methods(model)): + options = dict(cls.transition_options.get(method_name) or {}) + attr_name = options.get('name') or method_name + existing = getattr(cls, attr_name, None) + inherited = getattr(existing, 'transition_name', None) == method_name + + if method_name in cls.transition_exclude: + skipped[method_name] = 'excluded by transition_exclude' + cls._suppress_transition_action(attr_name, inherited) + continue + + if existing is not None and not inherited: + # A hand-written action or a viewset method of the same name wins. + skipped[method_name] = ( + f"'{attr_name}' is already defined on {cls.__name__}" + ) + continue + + pass_user, required = transition_call_plan(getattr(model, method_name)) + options.setdefault('pass_user', pass_user) + + if required and 'serializer_class' not in options: + skipped[method_name] = ( + 'takes arguments ({}) with no serializer_class'.format( + ', '.join(required) + ) + ) + cls._suppress_transition_action(attr_name, inherited) + continue + + options.setdefault('reset_output_options', hasattr(cls, 'output_options')) + setattr(cls, attr_name, transition_action(method_name, **options)) + generated[method_name] = attr_name + + return generated + + @classmethod + def _suppress_transition_action(cls, attr_name: str, inherited: bool) -> None: + """Hide an endpoint generated by a base class. + + Shadowing the attribute with None removes it from DRF's extra-action discovery without + touching the base class. + + Args: + attr_name: The viewset attribute the inherited endpoint occupies. + inherited: Whether that attribute holds a generated transition endpoint. + """ + if inherited: + setattr(cls, attr_name, None) + + @classmethod + def transition_url_paths(cls) -> dict[str, str]: + """Map transition method name to the URL path of the endpoint running it. + + Returns: + Mapping of transition name to detail-relative URL path, generated or hand-written. + """ + return { + handler.transition_name: handler.url_path + for handler in cls.get_extra_actions() + if getattr(handler, 'transition_name', None) + } + + @action( + detail=True, + methods=['get'], + url_path=TRANSITION_URL_PREFIX, + url_name='transitions', + serializer_class=AvailableTransitionSerializer, + ) + def transitions(self, request: Request, pk: str | None = None, *kwargs) -> Response: + """List the FSM transitions reachable from this object's current state. + + Blocked transitions are included and flagged with their reason. + + Only transitions this viewset exposes as an endpoint are reported: ``available_transitions`` + matches by state field name, so restricting the answer to this viewset's own endpoints keeps + every entry something the client can POST to. + """ + instance = self.get_object() + # permission check + model = instance._meta.model + if not ( + check_user_permission(request.user, model, 'view') + or check_user_permission(request.user, model, 'change') + ): + raise PermissionDenied( + _('User does not have permission to view or change this object') + ) + + paths = self.transition_url_paths() + + available: list[TransitionInfo] = [ + entry for entry in available_transitions(instance) if entry.name in paths + ] + data = [ + { + 'name': entry.name, + 'url_path': paths[entry.name], + 'target': entry.target, + 'label': entry.label, + 'blocked': entry.blocked, + 'blocking_reason': entry.blocking_reason, + } + for entry in available + ] + return Response(AvailableTransitionSerializer(data, many=True).data) + + +__all__ = [ + 'TRANSITION_URL_PREFIX', + 'FSMTransitionMixin', + 'current_state', + 'transition_action', + 'transition_call_plan', + 'transition_error', +] diff --git a/src/backend/InvenTree/generic/states/introspection.py b/src/backend/InvenTree/generic/states/introspection.py new file mode 100644 index 0000000000..262e626c8c --- /dev/null +++ b/src/backend/InvenTree/generic/states/introspection.py @@ -0,0 +1,104 @@ +"""FSM transition introspection for UI rendering (T-0709).""" + +from __future__ import annotations + +import inspect +from dataclasses import dataclass +from typing import Any + +from django_fsm import FSMFieldMixin + + +@dataclass +class TransitionInfo: + """Information about a transition reachable from an instance's current state.""" + + name: str + target: str + label: str + blocked: bool = False + blocking_reason: str | None = None + + +def state_fields(model: type): + """Return the names of the FSM state fields on model.""" + return tuple( + field.name for field in model._meta.fields if isinstance(field, FSMFieldMixin) + ) + + +def transition_methods(model: type, field_name: str | None = None): + """Return every transition method on model. + + Args: + model: The model class to scan + field_name: Restrict insights on specific name + + Returns: + Mapping of transition method name to its django-fsm metadata. + """ + found: dict[str, Any] = {} + + for name, func in inspect.getmembers(model): + meta = getattr(func, '_django_fsm', None) + if meta is None: + continue + + field = meta.field + declared = field if isinstance(field, str) else getattr(field, 'name', None) + if field_name is None or declared == field_name: + found[name] = meta + + return found + + +def available_transitions(instance): + """List the transitions reachable from an instance's current state.""" + found: list[TransitionInfo] = [] + + for field_name in state_fields(type(instance)): + current_state = getattr(instance, field_name, None) + if current_state is None: + continue + + for name, meta in transition_methods(type(instance), field_name).items(): + if not meta.has_transition(current_state): + continue + + transition = meta.get_transition(current_state) + blocked, reason = _evaluate_conditions(transition, instance) + + found.append( + TransitionInfo( + name=name, + target=str(transition.target) + if transition.target is not None + else '', + label=_label(transition, name), + blocked=blocked, + blocking_reason=reason, + ) + ) + return found + + +def _evaluate_conditions(transition, instance): + """Check a transition's conditions without running it.""" + for condition in transition.conditions or (): + try: + passed = condition(instance) + except Exception: + return True, 'Cannot evaluate transition conditions' + if not passed: + reason = getattr(condition, 'blocking_reason', None) + return True, str(reason) if reason is not None else None + + return False, None + + +def _label(transition, name: str): + """Return the human-readable label for a transition.""" + label = (transition.custom or {}).get('label') + if label: + return str(label) + return name.replace('_', ' ').title() diff --git a/src/backend/InvenTree/generic/states/serializers.py b/src/backend/InvenTree/generic/states/serializers.py index 52464dfaba..7078c75b41 100644 --- a/src/backend/InvenTree/generic/states/serializers.py +++ b/src/backend/InvenTree/generic/states/serializers.py @@ -39,3 +39,14 @@ class GenericStateClassSerializer(serializers.Serializer): values = serializers.DictField( child=GenericStateValueSerializer(), label=_('Values'), required=True ) + + +class AvailableTransitionSerializer(serializers.Serializer): + """Listing of available transitions for a given object.""" + + name = serializers.CharField(read_only=True) + url_path = serializers.CharField(read_only=True) + target = serializers.CharField(read_only=True) + label = serializers.CharField(read_only=True) + blocked = serializers.BooleanField(read_only=True) + blocking_reason = serializers.CharField(read_only=True, allow_null=True) diff --git a/src/backend/InvenTree/generic/states/tests.py b/src/backend/InvenTree/generic/states/tests.py index fe33eaf78e..1bb2d2585c 100644 --- a/src/backend/InvenTree/generic/states/tests.py +++ b/src/backend/InvenTree/generic/states/tests.py @@ -232,8 +232,8 @@ class ApiTests(InvenTreeAPITestCase): """Test the API endpoint for listing all status models.""" response = self.get(reverse('api-status-all')) - # 11 built-in state classes, plus the added GeneralState class - self.assertEqual(len(response.data), 12) + # 11 built-in state classes, plus the added GeneralState class. + self.assertGreaterEqual(len(response.data), 12) # Test the BuildStatus model build_status = response.data['BuildStatus'] @@ -273,7 +273,7 @@ class ApiTests(InvenTreeAPITestCase): ) response = self.get(reverse('api-status-all')) - self.assertEqual(len(response.data), 12) + self.assertGreaterEqual(len(response.data), 12) stock_status_cstm = response.data['StockStatus'] self.assertEqual(stock_status_cstm['status_class'], 'StockStatus') diff --git a/src/backend/InvenTree/generic/states/transition.py b/src/backend/InvenTree/generic/states/transition.py index 2543ef46cf..99c5de1456 100644 --- a/src/backend/InvenTree/generic/states/transition.py +++ b/src/backend/InvenTree/generic/states/transition.py @@ -70,12 +70,14 @@ def inventree_transition( @wraps(func) def wrapper(self, *args, **kwargs): """Ensure that transitions are handled correctly.""" + field_name = field if isinstance(field, str) else field.name + if refresh_field: try: # Update the field from the database to avoid race conditions current_obj = type(self).objects.select_for_update().get(pk=self.pk) - new_value = getattr(current_obj, field.name) - setattr(self, field.name, new_value) + new_value = getattr(current_obj, field_name) + setattr(self, field_name, new_value) except type(self).DoesNotExist: # pragma: no cover raise ValidationError( f'{self._meta.verbose_name} with pk={self.pk} does not exist in the database' @@ -84,7 +86,7 @@ def inventree_transition( # Run plugin transition handlers - if no step is taken the decorated method is called if result := _run_plugin_transition_handlers( self, - getattr(self, field.name), + getattr(self, field_name), target, default_action=_noop_default_action, ): @@ -106,7 +108,7 @@ def inventree_transition( # back to the generic "invalid transition" message below, since # there is no one value to compare against. resolved_target = getattr(target, 'target', target) - if getattr(self, field.name) == resolved_target: + if getattr(self, field_name) == resolved_target: target_val = ( resolved_target.label if isinstance(resolved_target, Enum) @@ -117,7 +119,7 @@ def inventree_transition( f'{self._meta.verbose_name} is already {target_val}' ) from exc raise ValidationError( - f'Invalid transition on {self._meta.verbose_name}.{field.name} (source value should be {source}, is {getattr(self, field.name)})' + f'Invalid transition on {self._meta.verbose_name}.{field_name} (source value should be {source}, is {getattr(self, field_name)})' ) from exc # Persist all changes (including the updated status field) to the DB. self.save() @@ -359,3 +361,31 @@ def _run_plugin_transition_handlers(instance, source, target, default_action): def _noop_default_action(current_state, target_state, instance, **kwargs): """No-op default action for compatibility with transition handlers.""" return None # pragma: no cover + + +def blocking_reason(reason): + """Provide general transition blocking reasoning.""" + + def decorate(func): + func.blocking_reason = reason + return func + + return decorate + + +def after_commit(fn: Callable, *args, **kwargs) -> None: + """Helper to ensure functions are only executed on successful transaction commit. + + This helps sending notifications and similar things only after the encapsulated transaction has been successfully committed. Avoiding confusion + + Args: + fn: Function to execute after commit + """ + + def run_after_commit(): + try: + fn(*args, **kwargs) + except Exception as e: + logger.error(f'Error in post-commit callback: {e}', exc_info=True) + + transaction.on_commit(run_after_commit) diff --git a/src/backend/InvenTree/order/api.py b/src/backend/InvenTree/order/api.py index 7285419569..970bcf0363 100644 --- a/src/backend/InvenTree/order/api.py +++ b/src/backend/InvenTree/order/api.py @@ -32,7 +32,7 @@ import company.models import stock.models as stock_models import stock.serializers as stock_serializers from data_exporter.mixins import DataExportViewMixin -from generic.states.api import StatusView +from generic.states.api import FSMTransitionMixin, StatusView from InvenTree.api import ( BulkDeleteMixin, BulkDeleteViewsetMixin, @@ -388,6 +388,7 @@ class PurchaseOrderViewSet( DataExportViewMixin, OutputOptionsMixin, ParameterListMixin, + FSMTransitionMixin, RetrieveUpdateDestroyModelViewSet, ): """API endpoint for accessing PurchaseOrder objects. @@ -405,6 +406,17 @@ class PurchaseOrderViewSet( 'supplier', 'created_by' ) serializer_class = serializers.PurchaseOrderSerializer + # TODO @matmair remove legacy return codes + transition_options = { + 'cancel_order': {'name': 'cancel', 'return_code': 201}, + 'complete_order': { + 'name': 'complete', + 'return_code': 201, + 'serializer_class': serializers.PurchaseOrderCompleteSerializer, + }, + 'hold_order': {'name': 'hold', 'return_code': 201}, + 'place_order': {'name': 'issue', 'return_code': 201}, + } ordering_field_aliases = { 'reference': ['reference_int', 'reference'], @@ -444,16 +456,7 @@ class PurchaseOrderViewSet( return queryset def get_order(self): - """Return the PurchaseOrder object associated with this API endpoint. - - Note: deliberately a raw lookup rather than self.get_object() - the latter - routes through ParameterListMixin.filter_queryset(), which assumes - self.serializer_class.Meta.model exists. That's true for the default - PurchaseOrderSerializer, but not for the plain-Serializer action classes - (PurchaseOrderHoldSerializer etc.) used by hold/cancel/complete/issue/receive - below, so calling get_object() from those actions raises an unrelated - AttributeError instead of the intended 404. - """ + """Return the PurchaseOrder object associated with this API endpoint.""" try: return models.PurchaseOrder.objects.get(pk=self.kwargs.get('pk', None)) except (ValueError, models.PurchaseOrder.DoesNotExist): @@ -476,81 +479,6 @@ class PurchaseOrderViewSet( return context - # TODO @matmair remove legacy return codes - @extend_schema(responses={201: serializers.PurchaseOrderHoldSerializer}) - @action( - detail=True, - methods=['post'], - serializer_class=serializers.PurchaseOrderHoldSerializer, - output_options=None, - ) - def hold(self, request, pk=None): - """API endpoint to place a PurchaseOrder on hold.""" - # Ensure the target order actually exists (raises NotFound -> 404 otherwise) - - # without this, a non-existent pk would fall through to the serializer's - # save(), which unconditionally reads self.context['order'], raising an - # unhandled KeyError (HTTP 500) instead of a clean 404. - self.get_order() - - serializer = self.get_serializer(data=request.data) - serializer.is_valid(raise_exception=True) - serializer.save() - return Response(serializer.data, status=status.HTTP_201_CREATED) - - # TODO @matmair remove legacy return codes - @extend_schema(responses={201: serializers.PurchaseOrderCancelSerializer}) - @action( - detail=True, - methods=['post'], - serializer_class=serializers.PurchaseOrderCancelSerializer, - output_options=None, - ) - def cancel(self, request, pk=None): - """API endpoint to 'cancel' a purchase order. - - The purchase order must be in a state which can be cancelled - """ - self.get_order() - - serializer = self.get_serializer(data=request.data) - serializer.is_valid(raise_exception=True) - serializer.save() - return Response(serializer.data, status=status.HTTP_201_CREATED) - - # TODO @matmair remove legacy return codes - @extend_schema(responses={201: serializers.PurchaseOrderCompleteSerializer}) - @action( - detail=True, - methods=['post'], - serializer_class=serializers.PurchaseOrderCompleteSerializer, - output_options=None, - ) - def complete(self, request, pk=None): - """API endpoint to 'complete' a purchase order.""" - self.get_order() - - serializer = self.get_serializer(data=request.data) - serializer.is_valid(raise_exception=True) - serializer.save() - return Response(serializer.data, status=status.HTTP_201_CREATED) - - # TODO @matmair remove legacy return codes - @extend_schema(responses={201: serializers.PurchaseOrderIssueSerializer}) - @action( - detail=True, - methods=['post'], - serializer_class=serializers.PurchaseOrderIssueSerializer, - output_options=None, - ) - def issue(self, request, pk=None): - """API endpoint to 'issue' (place) a PurchaseOrder.""" - self.get_order() - - serializer = self.get_serializer(data=request.data) - serializer.is_valid(raise_exception=True) - serializer.save() - return Response(serializer.data, status=status.HTTP_201_CREATED) - @extend_schema(responses={201: stock_serializers.StockItemSerializer(many=True)}) @action( detail=True, diff --git a/src/backend/InvenTree/order/models.py b/src/backend/InvenTree/order/models.py index d665022086..f275582da9 100644 --- a/src/backend/InvenTree/order/models.py +++ b/src/backend/InvenTree/order/models.py @@ -900,7 +900,7 @@ class PurchaseOrder(TotalPriceMixin, Order): target=PurchaseOrderStatus.COMPLETE, event=PurchaseOrderEvents.COMPLETED, ) - def complete_order(self): + def complete_order(self, **kwargs): """Transition this PurchaseOrder to COMPLETE status. The order must currently be PLACED. diff --git a/src/backend/InvenTree/order/serializers.py b/src/backend/InvenTree/order/serializers.py index b3d82d14ed..5a7602899f 100644 --- a/src/backend/InvenTree/order/serializers.py +++ b/src/backend/InvenTree/order/serializers.py @@ -509,25 +509,6 @@ class OrderAdjustSerializer(serializers.Serializer): return self.context['order'] -class PurchaseOrderHoldSerializer(OrderAdjustSerializer): - """Serializer for placing a PurchaseOrder on hold.""" - - def save(self): - """Save the serializer to 'hold' the order.""" - self.order.hold_order() - - -class PurchaseOrderCancelSerializer(OrderAdjustSerializer): - """Serializer for cancelling a PurchaseOrder.""" - - def save(self): - """Save the serializer to 'cancel' the order.""" - if not self.order.can_cancel: - raise ValidationError(_('Order cannot be cancelled')) - - self.order.cancel_order() - - class PurchaseOrderCompleteSerializer(OrderAdjustSerializer): """Serializer for completing a purchase order.""" @@ -558,18 +539,6 @@ class PurchaseOrderCompleteSerializer(OrderAdjustSerializer): return {'is_complete': order.is_complete} - def save(self): - """Save the serializer to 'complete' the order.""" - self.order.complete_order() - - -class PurchaseOrderIssueSerializer(OrderAdjustSerializer): - """Serializer for issuing (sending) a purchase order.""" - - def save(self): - """Save the serializer to 'place' the order.""" - self.order.place_order() - @register_importer() class PurchaseOrderLineItemSerializer( diff --git a/src/backend/InvenTree/order/test_api.py b/src/backend/InvenTree/order/test_api.py index e199f524fb..6ddc925109 100644 --- a/src/backend/InvenTree/order/test_api.py +++ b/src/backend/InvenTree/order/test_api.py @@ -82,6 +82,19 @@ class PurchaseOrderTest(OrderTest): LIST_URL = reverse('api-po-list') + def assert_available_transitions(self, po, available, unavailable): + """Assert which transitions are available for a purchase order.""" + response = self.get( + reverse('api-po-transitions', kwargs={'pk': po.pk}), expected_code=200 + ) + transitions = {transition['url_path'] for transition in response.json()} + + for transition in available: + self.assertIn(transition, transitions) + + for transition in unavailable: + self.assertNotIn(transition, transitions) + def test_options(self): """Test the PurchaseOrder OPTIONS endpoint.""" self.assignRole('purchase_order.add') @@ -838,14 +851,27 @@ class PurchaseOrderTest(OrderTest): po = models.PurchaseOrder.objects.get(pk=2) url = reverse('api-po-issue', kwargs={'pk': po.pk}) + transitions_url = reverse('api-po-transitions', kwargs={'pk': po.pk}) # Try to issue the PO, without required permissions + self.clearRoles() self.post(url, {}, expected_code=403) + # Check introspection endpoint too + self.get(transitions_url, expected_code=403) + self.assignRole('purchase_order.view') self.assignRole('purchase_order.add') + self.get(transitions_url, expected_code=200) + self.assert_available_transitions( + po, available=['issue'], unavailable=['complete'] + ) self.post(url, {}, expected_code=201) + self.assert_available_transitions( + po, available=['complete'], unavailable=['issue'] + ) + po.refresh_from_db() self.assertEqual(po.status, PurchaseOrderStatus.PLACED)