Coverage for src/backend/InvenTree/InvenTree/mixins.py: 95%
112 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 17:47 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 17:47 +0000
1"""Mixins for (API) views in the whole project."""
3from django.core.exceptions import FieldDoesNotExist
5from rest_framework import generics, mixins, status
6from rest_framework.response import Response
8import data_exporter.mixins
9import importer.mixins
10from InvenTree.fields import InvenTreeNotesField, OutputConfiguration
11from InvenTree.helpers import (
12 clean_markdown,
13 remove_non_printable_characters,
14 strip_html_tags,
15)
16from InvenTree.schema import schema_for_view_output_options
17from InvenTree.serializers import FilterableSerializerMixin
20class CleanMixin:
21 """Model mixin class which cleans inputs using nh3."""
23 # Define a list of field names which will *not* be cleaned
24 SAFE_FIELDS = []
26 def create(self, request, *args, **kwargs):
27 """Override to clean data before processing it."""
28 serializer = self.get_serializer(data=self.clean_data(request.data))
29 serializer.is_valid(raise_exception=True)
30 self.perform_create(serializer)
31 headers = self.get_success_headers(serializer.data)
32 return Response(
33 serializer.data, status=status.HTTP_201_CREATED, headers=headers
34 )
36 def update(self, request, *args, **kwargs):
37 """Override to clean data before processing it."""
38 partial = kwargs.pop('partial', False)
39 instance = self.get_object()
40 serializer = self.get_serializer(
41 instance, data=self.clean_data(request.data), partial=partial
42 )
43 serializer.is_valid(raise_exception=True)
44 self.perform_update(serializer)
46 if getattr(instance, '_prefetched_objects_cache', None):
47 # If 'prefetch_related' has been applied to a queryset, we need to
48 # forcibly invalidate the prefetch cache on the instance.
49 instance._prefetched_objects_cache = {}
51 return Response(serializer.data)
53 def clean_string(self, field: str, data: str) -> str:
54 """Clean / sanitize a single input string."""
55 cleaned = data
57 # By default, newline characters are removed
58 remove_newline = True
59 is_markdown = False
61 try:
62 if hasattr(self, 'serializer_class'): 62 ↛ 80line 62 didn't jump to line 80 because the condition on line 62 was always true
63 model = self.serializer_class.Meta.model
64 field_base = model._meta.get_field(field)
66 # The following field types allow newline characters
67 allow_newline = [(InvenTreeNotesField, True)]
69 for field_type in allow_newline:
70 if issubclass(type(field_base), field_type[0]):
71 remove_newline = False
72 is_markdown = field_type[1]
73 break
75 except AttributeError:
76 pass
77 except FieldDoesNotExist:
78 pass
80 cleaned = remove_non_printable_characters(
81 cleaned, remove_newline=remove_newline
82 )
84 cleaned = strip_html_tags(cleaned, field_name=field)
86 if is_markdown:
87 cleaned = clean_markdown(cleaned)
89 return cleaned
91 def clean_data(self, data: dict) -> dict:
92 """Clean / sanitize data.
94 This uses nh3 under the hood to disable certain html tags by
95 encoding them - this leads to script tags etc. to not work.
96 The results can be longer then the input; might make some character combinations
97 `ugly`. Prevents XSS on the server-level.
99 Args:
100 data (dict): Data that should be Sanitized.
102 Returns:
103 dict: Provided data Sanitized; still in the same order.
104 """
105 clean_data = {}
107 for k, v in data.items():
108 if k in self.SAFE_FIELDS: 108 ↛ 109line 108 didn't jump to line 109 because the condition on line 108 was never true
109 ret = v
110 elif isinstance(v, str):
111 ret = self.clean_string(k, v)
112 elif isinstance(v, dict):
113 ret = self.clean_data(v)
114 else:
115 ret = v
117 clean_data[k] = ret
119 return clean_data
122class ListAPI(generics.ListAPIView):
123 """View for list API."""
126class ListCreateAPI(CleanMixin, generics.ListCreateAPIView):
127 """View for list and create API."""
130class CreateAPI(CleanMixin, generics.CreateAPIView):
131 """View for create API."""
134class RetrieveAPI(generics.RetrieveAPIView):
135 """View for retrieve API."""
138class RetrieveUpdateAPI(CleanMixin, generics.RetrieveUpdateAPIView):
139 """View for retrieve and update API."""
142class CustomDestroyModelMixin:
143 """This mixin was created pass the kwargs from the API to the models."""
145 def destroy(self, request, *args, **kwargs):
146 """Custom destroy method to pass kwargs."""
147 instance = self.get_object()
148 self.perform_destroy(instance, **kwargs)
149 return Response(status=status.HTTP_204_NO_CONTENT)
151 def perform_destroy(self, instance, **kwargs):
152 """Custom destroy method to pass kwargs."""
153 instance.delete(**kwargs)
156class CustomRetrieveUpdateDestroyAPIView(
157 mixins.RetrieveModelMixin,
158 mixins.UpdateModelMixin,
159 CustomDestroyModelMixin,
160 generics.GenericAPIView,
161):
162 """This APIView was created pass the kwargs from the API to the models."""
164 def get(self, request, *args, **kwargs):
165 """Custom get method to pass kwargs."""
166 return self.retrieve(request, *args, **kwargs)
168 def put(self, request, *args, **kwargs):
169 """Custom put method to pass kwargs."""
170 return self.update(request, *args, **kwargs)
172 def patch(self, request, *args, **kwargs):
173 """Custom patch method to pass kwargs."""
174 return self.partial_update(request, *args, **kwargs)
176 def delete(self, request, *args, **kwargs):
177 """Custom delete method to pass kwargs."""
178 return self.destroy(request, *args, **kwargs)
181class CustomRetrieveUpdateDestroyAPI(CleanMixin, CustomRetrieveUpdateDestroyAPIView):
182 """This APIView was created pass the kwargs from the API to the models."""
185class RetrieveUpdateDestroyAPI(CleanMixin, generics.RetrieveUpdateDestroyAPIView):
186 """View for retrieve, update and destroy API."""
189class RetrieveDestroyAPI(generics.RetrieveDestroyAPIView):
190 """View for retrieve and destroy API."""
193class UpdateAPI(CleanMixin, generics.UpdateAPIView):
194 """View for update API."""
197class DataImportExportSerializerMixin(
198 data_exporter.mixins.DataExportSerializerMixin,
199 importer.mixins.DataImportSerializerMixin,
200):
201 """Mixin class for adding data import/export functionality to a DRF serializer."""
204class OutputOptionsMixin:
205 """Mixin to handle output options for API endpoints."""
207 output_options: OutputConfiguration = None
209 def __init_subclass__(cls, **kwargs):
210 """Automatically attaches OpenAPI schema parameters for its output options."""
211 super().__init_subclass__(**kwargs)
213 if getattr(cls, 'output_options', None) is not None:
214 schema_for_view_output_options(cls)
216 def get_serializer(self, *args, **kwargs):
217 """Return serializer instance with output options applied."""
218 request = getattr(self, 'request', None)
220 if self.output_options and request:
221 params = self.request.query_params
222 kwargs.update(self.output_options.format_params(params))
224 # Ensure the request is included in the serializer context
225 context = kwargs.get('context', {})
226 context['request'] = request
227 kwargs['context'] = context
229 return super().get_serializer(*args, **kwargs)
231 def get_queryset(self):
232 """Return the queryset with output options applied.
234 This automatically applies any prefetching defined against the optional fields.
235 """
236 queryset = super().get_queryset()
237 serializer = self.get_serializer()
239 if isinstance(serializer, FilterableSerializerMixin): 239 ↛ 242line 239 didn't jump to line 242 because the condition on line 239 was always true
240 queryset = serializer.prefetch_queryset(queryset)
242 return queryset
245class SerializerContextMixin:
246 """Mixin to add context to serializer."""
248 def get_serializer(self, *args, **kwargs):
249 """Add context to serializer."""
250 kwargs['context'] = self.get_serializer_context()
251 return super().get_serializer(*args, **kwargs)