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

1"""Mixins for (API) views in the whole project.""" 

2 

3from django.core.exceptions import FieldDoesNotExist 

4 

5from rest_framework import generics, mixins, status 

6from rest_framework.response import Response 

7 

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 

18 

19 

20class CleanMixin: 

21 """Model mixin class which cleans inputs using nh3.""" 

22 

23 # Define a list of field names which will *not* be cleaned 

24 SAFE_FIELDS = [] 

25 

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 ) 

35 

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) 

45 

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 = {} 

50 

51 return Response(serializer.data) 

52 

53 def clean_string(self, field: str, data: str) -> str: 

54 """Clean / sanitize a single input string.""" 

55 cleaned = data 

56 

57 # By default, newline characters are removed 

58 remove_newline = True 

59 is_markdown = False 

60 

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) 

65 

66 # The following field types allow newline characters 

67 allow_newline = [(InvenTreeNotesField, True)] 

68 

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 

74 

75 except AttributeError: 

76 pass 

77 except FieldDoesNotExist: 

78 pass 

79 

80 cleaned = remove_non_printable_characters( 

81 cleaned, remove_newline=remove_newline 

82 ) 

83 

84 cleaned = strip_html_tags(cleaned, field_name=field) 

85 

86 if is_markdown: 

87 cleaned = clean_markdown(cleaned) 

88 

89 return cleaned 

90 

91 def clean_data(self, data: dict) -> dict: 

92 """Clean / sanitize data. 

93 

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. 

98 

99 Args: 

100 data (dict): Data that should be Sanitized. 

101 

102 Returns: 

103 dict: Provided data Sanitized; still in the same order. 

104 """ 

105 clean_data = {} 

106 

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 

116 

117 clean_data[k] = ret 

118 

119 return clean_data 

120 

121 

122class ListAPI(generics.ListAPIView): 

123 """View for list API.""" 

124 

125 

126class ListCreateAPI(CleanMixin, generics.ListCreateAPIView): 

127 """View for list and create API.""" 

128 

129 

130class CreateAPI(CleanMixin, generics.CreateAPIView): 

131 """View for create API.""" 

132 

133 

134class RetrieveAPI(generics.RetrieveAPIView): 

135 """View for retrieve API.""" 

136 

137 

138class RetrieveUpdateAPI(CleanMixin, generics.RetrieveUpdateAPIView): 

139 """View for retrieve and update API.""" 

140 

141 

142class CustomDestroyModelMixin: 

143 """This mixin was created pass the kwargs from the API to the models.""" 

144 

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) 

150 

151 def perform_destroy(self, instance, **kwargs): 

152 """Custom destroy method to pass kwargs.""" 

153 instance.delete(**kwargs) 

154 

155 

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.""" 

163 

164 def get(self, request, *args, **kwargs): 

165 """Custom get method to pass kwargs.""" 

166 return self.retrieve(request, *args, **kwargs) 

167 

168 def put(self, request, *args, **kwargs): 

169 """Custom put method to pass kwargs.""" 

170 return self.update(request, *args, **kwargs) 

171 

172 def patch(self, request, *args, **kwargs): 

173 """Custom patch method to pass kwargs.""" 

174 return self.partial_update(request, *args, **kwargs) 

175 

176 def delete(self, request, *args, **kwargs): 

177 """Custom delete method to pass kwargs.""" 

178 return self.destroy(request, *args, **kwargs) 

179 

180 

181class CustomRetrieveUpdateDestroyAPI(CleanMixin, CustomRetrieveUpdateDestroyAPIView): 

182 """This APIView was created pass the kwargs from the API to the models.""" 

183 

184 

185class RetrieveUpdateDestroyAPI(CleanMixin, generics.RetrieveUpdateDestroyAPIView): 

186 """View for retrieve, update and destroy API.""" 

187 

188 

189class RetrieveDestroyAPI(generics.RetrieveDestroyAPIView): 

190 """View for retrieve and destroy API.""" 

191 

192 

193class UpdateAPI(CleanMixin, generics.UpdateAPIView): 

194 """View for update API.""" 

195 

196 

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.""" 

202 

203 

204class OutputOptionsMixin: 

205 """Mixin to handle output options for API endpoints.""" 

206 

207 output_options: OutputConfiguration = None 

208 

209 def __init_subclass__(cls, **kwargs): 

210 """Automatically attaches OpenAPI schema parameters for its output options.""" 

211 super().__init_subclass__(**kwargs) 

212 

213 if getattr(cls, 'output_options', None) is not None: 

214 schema_for_view_output_options(cls) 

215 

216 def get_serializer(self, *args, **kwargs): 

217 """Return serializer instance with output options applied.""" 

218 request = getattr(self, 'request', None) 

219 

220 if self.output_options and request: 

221 params = self.request.query_params 

222 kwargs.update(self.output_options.format_params(params)) 

223 

224 # Ensure the request is included in the serializer context 

225 context = kwargs.get('context', {}) 

226 context['request'] = request 

227 kwargs['context'] = context 

228 

229 return super().get_serializer(*args, **kwargs) 

230 

231 def get_queryset(self): 

232 """Return the queryset with output options applied. 

233 

234 This automatically applies any prefetching defined against the optional fields. 

235 """ 

236 queryset = super().get_queryset() 

237 serializer = self.get_serializer() 

238 

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) 

241 

242 return queryset 

243 

244 

245class SerializerContextMixin: 

246 """Mixin to add context to serializer.""" 

247 

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)