Coverage for src/backend/InvenTree/InvenTree/schema.py: 89%
168 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"""Schema processing functions for cleaning up generated schema."""
3from itertools import chain
4from typing import Any, Optional
6from django.conf import settings
8from drf_spectacular.contrib.django_oauth_toolkit import DjangoOAuthToolkitScheme
9from drf_spectacular.drainage import warn
10from drf_spectacular.openapi import AutoSchema
11from drf_spectacular.plumbing import ComponentRegistry
12from drf_spectacular.types import OpenApiTypes
13from drf_spectacular.utils import (
14 OpenApiParameter,
15 _SchemaType,
16 extend_schema,
17 extend_schema_view,
18)
19from rest_framework.pagination import LimitOffsetPagination
21from InvenTree.permissions import OASTokenMixin
22from users.oauth2_scopes import oauth2_scopes
25class ExtendedOAuth2Scheme(DjangoOAuthToolkitScheme):
26 """Extend drf-spectacular to allow customizing the schema to match the actual API behavior."""
28 target_class = 'users.authentication.ExtendedOAuth2Authentication'
30 def get_security_requirement(self, auto_schema):
31 """Get the security requirement for the current view."""
32 ret = super().get_security_requirement(auto_schema)
33 if ret: 33 ↛ 34line 33 didn't jump to line 34 because the condition on line 33 was never true
34 return ret
36 # If no security requirement is found, try if the view uses our OASTokenMixin
37 for permission in auto_schema.view.get_permissions():
38 if isinstance(permission, OASTokenMixin):
39 alt_scopes = permission.get_required_alternate_scopes(
40 auto_schema.view.request, auto_schema.view
41 )
42 alt_scopes = alt_scopes.get(auto_schema.method, [])
43 return [{self.name: group} for group in alt_scopes]
46class ExtendedAutoSchema(AutoSchema):
47 """Extend drf-spectacular to allow customizing the schema to match the actual API behavior."""
49 def is_bulk_action(self, ref: str) -> bool:
50 """Check the class of the current view for the bulk mixins."""
51 return ref in [c.__name__ for c in type(self.view).__mro__]
53 def get_operation_id(self) -> str:
54 """Custom path handling overrides, falling back to default behavior."""
55 result_id = super().get_operation_id()
57 # rename bulk actions to deconflict with single action operation_id
58 if (
59 (self.method == 'DELETE' and self.is_bulk_action('BulkDeleteMixin'))
60 or (
61 self.method == 'DELETE'
62 and self.is_bulk_action('BulkDeleteViewsetMixin')
63 and self.view.action == 'bulk_delete'
64 )
65 or (
66 (self.method == 'PUT' or self.method == 'PATCH')
67 and self.is_bulk_action('BulkUpdateMixin')
68 )
69 ):
70 action = self.method_mapping[self.method.lower()]
71 result_id = result_id.replace(action, 'bulk_' + action)
73 return result_id
75 def get_operation(
76 self,
77 path: str,
78 path_regex: str,
79 path_prefix: str,
80 method: str,
81 registry: ComponentRegistry,
82 ) -> Optional[_SchemaType]:
83 """Custom operation handling, falling back to default behavior."""
84 operation = super().get_operation(
85 path, path_regex, path_prefix, method, registry
86 )
87 if operation is None:
88 return None
90 # drf-spectacular doesn't support a body on DELETE endpoints because the semantics are not well-defined and
91 # OpenAPI recommends against it. This allows us to generate a schema that follows existing behavior.
92 if (self.method == 'DELETE' and self.is_bulk_action('BulkDeleteMixin')) or (
93 self.method == 'DELETE'
94 and getattr(self.view, 'action', None) == 'bulk_delete'
95 and self.is_bulk_action('BulkDeleteViewsetMixin')
96 ):
97 original_method = self.method
98 self.method = 'PUT'
99 request_body = self._get_request_body()
100 request_body['required'] = True
101 operation['requestBody'] = request_body
102 self.method = original_method
104 parameters = operation.get('parameters', [])
106 # If pagination limit is not set (default state) then all results will return unpaginated. This doesn't match
107 # what the schema defines to be the expected result. This forces limit to be present, producing the expected
108 # type.
109 pagination_class = getattr(self.view, 'pagination_class', None)
110 if pagination_class and pagination_class == LimitOffsetPagination:
111 for parameter in parameters:
112 if parameter['name'] == 'limit':
113 parameter['required'] = True
115 # Add valid order selections to the ordering field description.
116 ordering_fields = getattr(self.view, 'ordering_fields', None)
117 if ordering_fields is not None:
118 for parameter in parameters:
119 if parameter['name'] == 'ordering':
120 schema_order = []
121 for field in ordering_fields:
122 schema_order.append(field)
123 schema_order.append('-' + field)
124 parameter['schema']['enum'] = schema_order
126 # Add valid search fields to the search description.
127 search_fields = getattr(self.view, 'search_fields', None)
129 if search_fields is not None:
130 # Ensure consistent ordering of search fields
131 search_fields = sorted(search_fields)
132 for parameter in parameters:
133 if parameter['name'] == 'search':
134 parameter['description'] = (
135 f'{parameter["description"]} Searched fields: {", ".join(search_fields)}.'
136 )
138 # Change return to array type, simply annotating this return type attempts to paginate, which doesn't work for
139 # a create method and removing the pagination also affects the list method
140 if self.method == 'POST' and type(self.view).__name__ == 'StockList':
141 schema = operation['responses']['201']['content']['application/json'][
142 'schema'
143 ]
144 schema['type'] = 'array'
145 schema['items'] = {'$ref': schema['$ref']}
146 del schema['$ref']
148 # Add vendor extensions for custom behavior
149 operation.update(self.get_inventree_extensions())
151 return operation
153 def get_inventree_extensions(self):
154 """Add InvenTree specific extensions to the schema."""
155 from rest_framework.generics import RetrieveAPIView
156 from rest_framework.mixins import RetrieveModelMixin, UpdateModelMixin
158 from data_exporter.mixins import DataExportViewMixin
159 from InvenTree.api import BulkOperationMixin
160 from InvenTree.mixins import CleanMixin
162 lvl = settings.SCHEMA_VENDOREXTENSION_LEVEL
163 """Level of detail for InvenTree extensions."""
165 if lvl == 0: 165 ↛ 168line 165 didn't jump to line 168 because the condition on line 165 was always true
166 return {}
168 mro = self.view.__class__.__mro__
170 data = {}
171 if lvl >= 1:
172 data['x-inventree-meta'] = {
173 'version': '1.0',
174 'is_detail': any(
175 a in mro
176 for a in [RetrieveModelMixin, UpdateModelMixin, RetrieveAPIView]
177 ),
178 'is_bulk': BulkOperationMixin in mro,
179 'is_cleaned': CleanMixin in mro,
180 'is_filtered': hasattr(self.view, 'output_options'),
181 'is_exported': DataExportViewMixin in mro,
182 }
183 if lvl >= 2:
184 data['x-inventree-components'] = [str(a) for a in mro]
185 try:
186 qs = self.view.get_queryset()
187 qs = qs.model if qs is not None and hasattr(qs, 'model') else None
188 except Exception:
189 qs = None
191 data['x-inventree-model'] = {
192 'scope': 'core',
193 'model': str(qs.__name__) if qs else None,
194 'app': str(qs._meta.app_label) if qs else None,
195 }
197 return data
200def postprocess_schema_enums(result, generator, **kwargs):
201 """Override call to drf-spectacular's enum postprocessor to filter out specific warnings."""
202 from drf_spectacular import drainage
204 # Monkey-patch the warn function temporarily
205 original_warn = drainage.warn
207 def custom_warn(msg: str, delayed: Any = None) -> None:
208 """Custom patch to ignore some drf-spectacular warnings.
210 - Some warnings are unavoidable due to the way that InvenTree implements generic relationships (via ContentType).
211 - The cleanest way to handle this appears to be to override the 'warn' function from drf-spectacular.
213 Ref: https://github.com/inventree/InvenTree/pull/10699
214 """
215 ignore_patterns = [
216 'enum naming encountered a non-optimally resolvable collision for fields named "model_type"'
217 ]
219 if any(pattern in msg for pattern in ignore_patterns): 219 ↛ 222line 219 didn't jump to line 222 because the condition on line 219 was always true
220 return
222 original_warn(msg, delayed)
224 # Replace the warn function with our custom version
225 drainage.warn = custom_warn
227 import drf_spectacular.hooks
229 result = drf_spectacular.hooks.postprocess_schema_enums(result, generator, **kwargs)
231 # Restore the original warn function
232 drainage.warn = original_warn
234 return result
237def postprocess_required_nullable(result, generator, request, public):
238 """Un-require nullable fields.
240 Read-only values are all marked as required by spectacular, but InvenTree doesn't always include them in the
241 response. This removes them from the required list to allow responses lacking read-only nullable fields to validate
242 against the schema.
243 """
244 # Process schema section
245 schemas = result.get('components', {}).get('schemas', {})
246 for schema in schemas.values():
247 required_fields = schema.get('required', [])
248 properties = schema.get('properties', {})
250 # copy list to allow removing from it while iterating
251 for field in list(required_fields):
252 field_dict = properties.get(field, {})
253 if field_dict.get('readOnly') and field_dict.get('nullable'):
254 required_fields.remove(field)
255 if 'required' in schema and len(required_fields) == 0:
256 schema.pop('required')
258 return result
261def postprocess_print_stats(result, generator, request, public):
262 """Prints statistics against schema."""
263 rlt_dict = {}
264 for path in result['paths']:
265 for method in result['paths'][path]:
266 sec = result['paths'][path][method].get('security', [])
267 scopes = list(filter(None, (item.get('oauth2') for item in sec)))
268 rlt_dict[f'{path}:{method}'] = {
269 'method': method,
270 'oauth': list(chain(*scopes)),
271 'sec': sec is None,
272 }
274 # Get paths without oauth2
275 no_oauth2 = [
276 path for path, details in rlt_dict.items() if not any(details['oauth'])
277 ]
278 no_oauth2_wa = [path for path in no_oauth2 if not path.startswith('/api/auth/v1/')]
279 # Get paths without security
280 no_security = [path for path, details in rlt_dict.items() if details['sec']]
281 # Get path counts per scope
282 scopes = {}
283 for path, details in rlt_dict.items():
284 if details['oauth']:
285 for scope in list(details['oauth']):
286 if scope not in scopes:
287 scopes[scope] = []
288 scopes[scope].append(path)
289 # Sort scopes by keys
290 scopes = dict(sorted(scopes.items()))
292 # Print statistics
293 print('\nSchema statistics:')
294 print(f'Paths without oauth2: {len(no_oauth2)}')
295 print(f'Paths without oauth2 (without allauth): {len(no_oauth2_wa)}')
296 print(f'Paths without security: {len(no_security)}\n')
297 print('Scope stats:')
298 for scope, paths in scopes.items():
299 print(f' {scope}: {len(paths)}')
300 print()
302 # Check for unknown scopes
303 for scope, paths in scopes.items():
304 if scope not in oauth2_scopes: 304 ↛ 305line 304 didn't jump to line 305 because the condition on line 304 was never true
305 warn(f'unknown scope `{scope}` in {len(paths)} paths')
307 # Raise error if the paths missing scopes are not specifically excluded from oauth2
308 wrong_url = [
309 path for path in no_oauth2_wa if path not in settings.OAUTH2_CHECK_EXCLUDED
310 ]
311 if len(wrong_url) > 0: 311 ↛ 312line 311 didn't jump to line 312 because the condition on line 311 was never true
312 warn(
313 f'Found {len(wrong_url)} paths without oauth2 that are not excluded:\n{", ".join(wrong_url)}. '
314 '\n\nPlease check the schema and add oauth2 scopes where necessary.'
315 )
317 return result
320def schema_for_view_output_options(view_class):
321 """A class decorator that automatically generates schema parameters for a view.
323 It works by introspecting the `output_options` attribute on the view itself.
324 This decorator reads the `output_options` attribute from the view class,
325 extracts the `OPTIONS` list from it, and creates an OpenApiParameter for each option.
326 """
327 output_config_class = view_class.output_options
329 parameters = []
330 for option in output_config_class.OPTIONS:
331 param = OpenApiParameter(
332 name=option.flag,
333 type=OpenApiTypes.BOOL,
334 location=OpenApiParameter.QUERY,
335 description=option.description,
336 default=option.default,
337 )
338 parameters.append(param)
340 extended_view = extend_schema_view(get=extend_schema(parameters=parameters))(
341 view_class
342 )
343 return extended_view
346def exclude_from_schema(klass: type[Any], alternative_path: str) -> type[Any]:
347 """Decorator to exclude a view from the OpenAPI schema.
349 This is used to hide legacy endpoints from the schema, while still retaining them for backwards compatibility.
350 """
352 class LegacyView(klass):
353 """Dummy doc."""
355 LegacyView.__name__ = klass.__name__ + ' - Legacy'
356 LegacyView.__doc__ = f'This is a legacy endpoint, retained for backwards compatibility. Consider migrating to the new endpoint under {alternative_path}.'
358 # Exclude all default operations from the schema
359 for operation in ['get', 'post', 'put', 'patch', 'delete']:
360 if hasattr(klass, operation):
361 LegacyView = extend_schema_view(**{operation: extend_schema(exclude=True)})(
362 LegacyView
363 )
364 return LegacyView