Coverage for utilities/filters.py: 76%

109 statements  

« 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 

9 

10from .forms.fields import BigIntegerField 

11 

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) 

29 

30 

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 

38 

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 ] 

47 

48 def run_validators(self, value): 

49 for v in value: 

50 super().run_validators(v) 

51 

52 def validate(self, value): 

53 for v in value: 

54 super().validate(v) 

55 

56 return type(f'MultiValue{field_class.__name__}', (NewField,), dict()) 

57 

58 

59# 

60# Filters 

61# 

62 

63@extend_schema_field(OpenApiTypes.STR) 

64class MultiValueCharFilter(django_filters.MultipleChoiceFilter): 

65 field_class = multivalue_field_factory(forms.CharField) 

66 

67 

68@extend_schema_field(OpenApiTypes.DATE) 

69class MultiValueDateFilter(django_filters.MultipleChoiceFilter): 

70 field_class = multivalue_field_factory(forms.DateField) 

71 

72 

73@extend_schema_field(OpenApiTypes.DATETIME) 

74class MultiValueDateTimeFilter(django_filters.MultipleChoiceFilter): 

75 field_class = multivalue_field_factory(forms.DateTimeField) 

76 

77 

78@extend_schema_field(OpenApiTypes.INT32) 

79class MultiValueNumberFilter(django_filters.MultipleChoiceFilter): 

80 field_class = multivalue_field_factory(forms.IntegerField) 

81 

82 

83@extend_schema_field(OpenApiTypes.INT64) 

84class MultiValueBigNumberFilter(MultiValueNumberFilter): 

85 field_class = multivalue_field_factory(BigIntegerField) 

86 

87 

88@extend_schema_field(OpenApiTypes.DECIMAL) 

89class MultiValueDecimalFilter(django_filters.MultipleChoiceFilter): 

90 field_class = multivalue_field_factory(forms.DecimalField) 

91 

92 

93@extend_schema_field(OpenApiTypes.TIME) 

94class MultiValueTimeFilter(django_filters.MultipleChoiceFilter): 

95 field_class = multivalue_field_factory(forms.TimeField) 

96 

97 

98@extend_schema_field(OpenApiTypes.STR) 

99class MultiValueArrayFilter(django_filters.MultipleChoiceFilter): 

100 field_class = multivalue_field_factory(forms.CharField) 

101 

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) 

105 

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) 

111 

112 

113@extend_schema_field(OpenApiTypes.STR) 

114class MultiValueMACAddressFilter(django_filters.MultipleChoiceFilter): 

115 field_class = multivalue_field_factory(forms.CharField) 

116 

117 def filter(self, qs, value): 

118 try: 

119 return super().filter(qs, value) 

120 except ValidationError: 

121 return qs.none() 

122 

123 

124@extend_schema_field(OpenApiTypes.STR) 

125class MultiValueWWNFilter(django_filters.MultipleChoiceFilter): 

126 field_class = multivalue_field_factory(forms.CharField) 

127 

128 

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) 

139 

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) 

143 

144 

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 

154 

155 

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) 

164 

165 

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 

173 

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 ) 

184 

185 

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 

193 

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 

202 

203 return qs.filter( 

204 **{ 

205 f'{self.field_name}__in': content_types, 

206 } 

207 )