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
This commit is contained in:
Matthias Mair authored and GitHub committed 2026-09-29 21:26:53 +10:00
1 parent afab76835f
commit 9478c25a70
14 files changed
+631 -131

No files matched your search

+8 -3
View File
@@ -647,12 +647,17 @@ class ParameterListMixin:
serializer_class = ( serializer_class = (
getattr(self, 'serializer_class', None) or self.get_serializer_class() getattr(self, 'serializer_class', None) or self.get_serializer_class()
) )
model_class = getattr(getattr(serializer_class, 'Meta', None), 'model', None)
model_class = serializer_class.Meta.model
# Apply ordering based on query parameter # Apply ordering based on query parameter
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 = common.filters.order_by_parameter(
queryset, model_class, self.request.query_params.get('ordering', None) queryset, model_class, ordering
) )
return queryset return queryset
@@ -1,11 +1,15 @@
"""InvenTree API version information.""" """InvenTree API version information."""
# InvenTree API version # 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.""" """Increment this API version number whenever there is a significant change to the API that any clients need to know about."""
INVENTREE_API_TEXT = """ 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 v550 -> 2026-09-22 : https://github.com/inventree/InvenTree/pull/12912
- Adds a top-level 'batch_code' field to the PurchaseOrderReceive API endpoint - Adds a top-level 'batch_code' field to the PurchaseOrderReceive API endpoint
+8 -1
View File
@@ -194,8 +194,15 @@ class OutputOptionsMixin:
def get_serializer(self, *args, **kwargs): def get_serializer(self, *args, **kwargs):
"""Return serializer instance with output options applied.""" """Return serializer instance with output options applied."""
request = getattr(self, 'request', None) 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 params = self.request.query_params
kwargs.update(self.output_options.format_params(params)) kwargs.update(self.output_options.format_params(params))
@@ -15,6 +15,8 @@ from .transition import (
DEFERRABLE, DEFERRABLE,
StateTransitionMixin, StateTransitionMixin,
TransitionMethod, TransitionMethod,
after_commit,
blocking_reason,
inventree_transition, inventree_transition,
) )
@@ -27,6 +29,8 @@ __all__ = [
'StatusCode', 'StatusCode',
'StatusCodeMixin', 'StatusCodeMixin',
'TransitionMethod', 'TransitionMethod',
'after_commit',
'blocking_reason',
'can_proceed', # django_fsm import 'can_proceed', # django_fsm import
'deprecated', 'deprecated',
'fields', 'fields',
@@ -19,6 +19,7 @@ from InvenTree.helpers import inheritors
from InvenTree.mixins import ListCreateAPI, RetrieveUpdateDestroyAPI from InvenTree.mixins import ListCreateAPI, RetrieveUpdateDestroyAPI
from InvenTree.serializers import EmptySerializer from InvenTree.serializers import EmptySerializer
from .api_helpers import FSMTransitionMixin # noqa # pylint: disable=unused-import
from .serializers import GenericStateClassSerializer from .serializers import GenericStateClassSerializer
from .states import StatusCode from .states import StatusCode
@@ -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',
]
@@ -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()
@@ -39,3 +39,14 @@ class GenericStateClassSerializer(serializers.Serializer):
values = serializers.DictField( values = serializers.DictField(
child=GenericStateValueSerializer(), label=_('Values'), required=True 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)
@@ -232,8 +232,8 @@ class ApiTests(InvenTreeAPITestCase):
"""Test the API endpoint for listing all status models.""" """Test the API endpoint for listing all status models."""
response = self.get(reverse('api-status-all')) response = self.get(reverse('api-status-all'))
# 11 built-in state classes, plus the added GeneralState class # 11 built-in state classes, plus the added GeneralState class.
self.assertEqual(len(response.data), 12) self.assertGreaterEqual(len(response.data), 12)
# Test the BuildStatus model # Test the BuildStatus model
build_status = response.data['BuildStatus'] build_status = response.data['BuildStatus']
@@ -273,7 +273,7 @@ class ApiTests(InvenTreeAPITestCase):
) )
response = self.get(reverse('api-status-all')) 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'] stock_status_cstm = response.data['StockStatus']
self.assertEqual(stock_status_cstm['status_class'], 'StockStatus') self.assertEqual(stock_status_cstm['status_class'], 'StockStatus')
@@ -70,12 +70,14 @@ def inventree_transition(
@wraps(func) @wraps(func)
def wrapper(self, *args, **kwargs): def wrapper(self, *args, **kwargs):
"""Ensure that transitions are handled correctly.""" """Ensure that transitions are handled correctly."""
field_name = field if isinstance(field, str) else field.name
if refresh_field: if refresh_field:
try: try:
# Update the field from the database to avoid race conditions # Update the field from the database to avoid race conditions
current_obj = type(self).objects.select_for_update().get(pk=self.pk) current_obj = type(self).objects.select_for_update().get(pk=self.pk)
new_value = getattr(current_obj, field.name) new_value = getattr(current_obj, field_name)
setattr(self, field.name, new_value) setattr(self, field_name, new_value)
except type(self).DoesNotExist: # pragma: no cover except type(self).DoesNotExist: # pragma: no cover
raise ValidationError( raise ValidationError(
f'{self._meta.verbose_name} with pk={self.pk} does not exist in the database' 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 # Run plugin transition handlers - if no step is taken the decorated method is called
if result := _run_plugin_transition_handlers( if result := _run_plugin_transition_handlers(
self, self,
getattr(self, field.name), getattr(self, field_name),
target, target,
default_action=_noop_default_action, default_action=_noop_default_action,
): ):
@@ -106,7 +108,7 @@ def inventree_transition(
# back to the generic "invalid transition" message below, since # back to the generic "invalid transition" message below, since
# there is no one value to compare against. # there is no one value to compare against.
resolved_target = getattr(target, 'target', target) resolved_target = getattr(target, 'target', target)
if getattr(self, field.name) == resolved_target: if getattr(self, field_name) == resolved_target:
target_val = ( target_val = (
resolved_target.label resolved_target.label
if isinstance(resolved_target, Enum) if isinstance(resolved_target, Enum)
@@ -117,7 +119,7 @@ def inventree_transition(
f'{self._meta.verbose_name} is already {target_val}' f'{self._meta.verbose_name} is already {target_val}'
) from exc ) from exc
raise ValidationError( 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 ) from exc
# Persist all changes (including the updated status field) to the DB. # Persist all changes (including the updated status field) to the DB.
self.save() 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): def _noop_default_action(current_state, target_state, instance, **kwargs):
"""No-op default action for compatibility with transition handlers.""" """No-op default action for compatibility with transition handlers."""
return None # pragma: no cover 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)
+14 -86
View File
@@ -32,7 +32,7 @@ import company.models
import stock.models as stock_models import stock.models as stock_models
import stock.serializers as stock_serializers import stock.serializers as stock_serializers
from data_exporter.mixins import DataExportViewMixin from data_exporter.mixins import DataExportViewMixin
from generic.states.api import StatusView from generic.states.api import FSMTransitionMixin, StatusView
from InvenTree.api import ( from InvenTree.api import (
BulkDeleteMixin, BulkDeleteMixin,
BulkDeleteViewsetMixin, BulkDeleteViewsetMixin,
@@ -388,6 +388,7 @@ class PurchaseOrderViewSet(
DataExportViewMixin, DataExportViewMixin,
OutputOptionsMixin, OutputOptionsMixin,
ParameterListMixin, ParameterListMixin,
FSMTransitionMixin,
RetrieveUpdateDestroyModelViewSet, RetrieveUpdateDestroyModelViewSet,
): ):
"""API endpoint for accessing PurchaseOrder objects. """API endpoint for accessing PurchaseOrder objects.
@@ -405,6 +406,17 @@ class PurchaseOrderViewSet(
'supplier', 'created_by' 'supplier', 'created_by'
) )
serializer_class = serializers.PurchaseOrderSerializer 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 = { ordering_field_aliases = {
'reference': ['reference_int', 'reference'], 'reference': ['reference_int', 'reference'],
@@ -444,16 +456,7 @@ class PurchaseOrderViewSet(
return queryset return queryset
def get_order(self): def get_order(self):
"""Return the PurchaseOrder object associated with this API endpoint. """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.
"""
try: try:
return models.PurchaseOrder.objects.get(pk=self.kwargs.get('pk', None)) return models.PurchaseOrder.objects.get(pk=self.kwargs.get('pk', None))
except (ValueError, models.PurchaseOrder.DoesNotExist): except (ValueError, models.PurchaseOrder.DoesNotExist):
@@ -476,81 +479,6 @@ class PurchaseOrderViewSet(
return context 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)}) @extend_schema(responses={201: stock_serializers.StockItemSerializer(many=True)})
@action( @action(
detail=True, detail=True,
+1 -1
View File
@@ -900,7 +900,7 @@ class PurchaseOrder(TotalPriceMixin, Order):
target=PurchaseOrderStatus.COMPLETE, target=PurchaseOrderStatus.COMPLETE,
event=PurchaseOrderEvents.COMPLETED, event=PurchaseOrderEvents.COMPLETED,
) )
def complete_order(self): def complete_order(self, **kwargs):
"""Transition this PurchaseOrder to COMPLETE status. """Transition this PurchaseOrder to COMPLETE status.
The order must currently be PLACED. The order must currently be PLACED.
@@ -509,25 +509,6 @@ class OrderAdjustSerializer(serializers.Serializer):
return self.context['order'] 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): class PurchaseOrderCompleteSerializer(OrderAdjustSerializer):
"""Serializer for completing a purchase order.""" """Serializer for completing a purchase order."""
@@ -558,18 +539,6 @@ class PurchaseOrderCompleteSerializer(OrderAdjustSerializer):
return {'is_complete': order.is_complete} 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() @register_importer()
class PurchaseOrderLineItemSerializer( class PurchaseOrderLineItemSerializer(
+26
View File
@@ -82,6 +82,19 @@ class PurchaseOrderTest(OrderTest):
LIST_URL = reverse('api-po-list') 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): def test_options(self):
"""Test the PurchaseOrder OPTIONS endpoint.""" """Test the PurchaseOrder OPTIONS endpoint."""
self.assignRole('purchase_order.add') self.assignRole('purchase_order.add')
@@ -838,14 +851,27 @@ class PurchaseOrderTest(OrderTest):
po = models.PurchaseOrder.objects.get(pk=2) po = models.PurchaseOrder.objects.get(pk=2)
url = reverse('api-po-issue', kwargs={'pk': po.pk}) 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 # Try to issue the PO, without required permissions
self.clearRoles()
self.post(url, {}, expected_code=403) 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.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.post(url, {}, expected_code=201)
self.assert_available_transitions(
po, available=['complete'], unavailable=['issue']
)
po.refresh_from_db() po.refresh_from_db()
self.assertEqual(po.status, PurchaseOrderStatus.PLACED) self.assertEqual(po.status, PurchaseOrderStatus.PLACED)