Coverage for utilities/filters.py: 76%
109 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 18:35 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-10 18:35 +0000
1import django_filters
2from django import forms
3from django.conf import settings
4from django.contrib.contenttypes.models import ContentType
5from django.core.exceptions import ValidationError
6from django_filters.constants import EMPTY_VALUES
7from drf_spectacular.types import OpenApiTypes
8from drf_spectacular.utils import extend_schema_field
10from .forms.fields import BigIntegerField
12__all__ = (
13 'ContentTypeFilter',
14 'MultiValueArrayFilter',
15 'MultiValueBigNumberFilter',
16 'MultiValueCharFilter',
17 'MultiValueContentTypeFilter',
18 'MultiValueDateFilter',
19 'MultiValueDateTimeFilter',
20 'MultiValueDecimalFilter',
21 'MultiValueMACAddressFilter',
22 'MultiValueNumberFilter',
23 'MultiValueTimeFilter',
24 'MultiValueWWNFilter',
25 'NullableCharFieldFilter',
26 'NumericArrayFilter',
27 'TreeNodeMultipleChoiceFilter',
28)
31def multivalue_field_factory(field_class):
32 """
33 Given a form field class, return a subclass capable of accepting multiple values. This allows us to OR on multiple
34 filter values while maintaining the field's built-in validation. Example: GET /api/dcim/devices/?name=foo&name=bar
35 """
36 class NewField(field_class):
37 widget = forms.SelectMultiple
39 def to_python(self, value):
40 if not value:
41 return []
42 field = field_class()
43 return [
44 # Only append non-empty values (this avoids e.g. trying to cast '' as an integer)
45 field.to_python(v) for v in value if v
46 ]
48 def run_validators(self, value):
49 for v in value:
50 super().run_validators(v)
52 def validate(self, value):
53 for v in value:
54 super().validate(v)
56 return type(f'MultiValue{field_class.__name__}', (NewField,), dict())
59#
60# Filters
61#
63@extend_schema_field(OpenApiTypes.STR)
64class MultiValueCharFilter(django_filters.MultipleChoiceFilter):
65 field_class = multivalue_field_factory(forms.CharField)
68@extend_schema_field(OpenApiTypes.DATE)
69class MultiValueDateFilter(django_filters.MultipleChoiceFilter):
70 field_class = multivalue_field_factory(forms.DateField)
73@extend_schema_field(OpenApiTypes.DATETIME)
74class MultiValueDateTimeFilter(django_filters.MultipleChoiceFilter):
75 field_class = multivalue_field_factory(forms.DateTimeField)
78@extend_schema_field(OpenApiTypes.INT32)
79class MultiValueNumberFilter(django_filters.MultipleChoiceFilter):
80 field_class = multivalue_field_factory(forms.IntegerField)
83@extend_schema_field(OpenApiTypes.INT64)
84class MultiValueBigNumberFilter(MultiValueNumberFilter):
85 field_class = multivalue_field_factory(BigIntegerField)
88@extend_schema_field(OpenApiTypes.DECIMAL)
89class MultiValueDecimalFilter(django_filters.MultipleChoiceFilter):
90 field_class = multivalue_field_factory(forms.DecimalField)
93@extend_schema_field(OpenApiTypes.TIME)
94class MultiValueTimeFilter(django_filters.MultipleChoiceFilter):
95 field_class = multivalue_field_factory(forms.TimeField)
98@extend_schema_field(OpenApiTypes.STR)
99class MultiValueArrayFilter(django_filters.MultipleChoiceFilter):
100 field_class = multivalue_field_factory(forms.CharField)
102 def __init__(self, *args, lookup_expr='contains', **kwargs):
103 # Set default lookup_expr to 'contains'
104 super().__init__(*args, lookup_expr=lookup_expr, **kwargs)
106 def get_filter_predicate(self, v):
107 # If filtering for null values, ignore lookup_expr
108 if v is None:
109 return {self.field_name: None}
110 return super().get_filter_predicate(v)
113@extend_schema_field(OpenApiTypes.STR)
114class MultiValueMACAddressFilter(django_filters.MultipleChoiceFilter):
115 field_class = multivalue_field_factory(forms.CharField)
117 def filter(self, qs, value):
118 try:
119 return super().filter(qs, value)
120 except ValidationError:
121 return qs.none()
124@extend_schema_field(OpenApiTypes.STR)
125class MultiValueWWNFilter(django_filters.MultipleChoiceFilter):
126 field_class = multivalue_field_factory(forms.CharField)
129@extend_schema_field(OpenApiTypes.STR)
130class TreeNodeMultipleChoiceFilter(django_filters.ModelMultipleChoiceFilter):
131 """
132 Filters for a set of Models, including all descendant models within a Tree. Example: [<Region: R1>,<Region: R2>]
133 """
134 def get_filter_predicate(self, v):
135 # Null value filtering
136 if v is None:
137 return {f"{self.field_name}__isnull": True}
138 return super().get_filter_predicate(v)
140 def filter(self, qs, value):
141 value = [node.get_descendants(include_self=True) if not isinstance(node, str) else node for node in value]
142 return super().filter(qs, value)
145class NullableCharFieldFilter(django_filters.CharFilter):
146 """
147 Allow matching on null field values by passing a special string used to signify NULL.
148 """
149 def filter(self, qs, value):
150 if value != settings.FILTERS_NULL_CHOICE_VALUE:
151 return super().filter(qs, value)
152 qs = self.get_method(qs)(**{'{}__isnull'.format(self.field_name): True})
153 return qs.distinct() if self.distinct else qs
156class NumericArrayFilter(django_filters.NumberFilter):
157 """
158 Filter based on the presence of an integer within an ArrayField.
159 """
160 def filter(self, qs, value):
161 if value: 161 ↛ 162line 161 didn't jump to line 162 because the condition on line 161 was never true
162 value = [value]
163 return super().filter(qs, value)
166class ContentTypeFilter(django_filters.CharFilter):
167 """
168 Allow specifying a ContentType by <app_label>.<model> (e.g. "dcim.site").
169 """
170 def filter(self, qs, value):
171 if value in EMPTY_VALUES:
172 return qs
174 try:
175 app_label, model = value.lower().split('.')
176 content_type = ContentType.objects.get_by_natural_key(app_label, model)
177 except (ValueError, ContentType.DoesNotExist):
178 return qs.none()
179 return qs.filter(
180 **{
181 f'{self.field_name}': content_type,
182 }
183 )
186class MultiValueContentTypeFilter(MultiValueCharFilter):
187 """
188 A multi-value version of ContentTypeFilter.
189 """
190 def filter(self, qs, value):
191 if value in EMPTY_VALUES:
192 return qs
194 content_types = []
195 for key in value:
196 try:
197 app_label, model = key.lower().split('.')
198 ct = ContentType.objects.get_by_natural_key(app_label, model)
199 content_types.append(ct)
200 except (ValueError, ContentType.DoesNotExist):
201 continue
203 return qs.filter(
204 **{
205 f'{self.field_name}__in': content_types,
206 }
207 )