Coverage for utilities/testing/filtersets.py: 0%

77 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-10 18:35 +0000

1from datetime import UTC, datetime 

2from itertools import chain 

3 

4import django_filters 

5from django.contrib.contenttypes.fields import GenericForeignKey, GenericRelation 

6from django.contrib.contenttypes.models import ContentType 

7from django.db.models import ForeignKey, ManyToManyField, ManyToManyRel, ManyToOneRel, OneToOneRel 

8from django.utils.module_loading import import_string 

9from taggit.managers import TaggableManager 

10 

11from extras.filters import TagFilter 

12from netbox.models.ltree import LtreeModel 

13from utilities.filters import MultiValueContentTypeFilter, TreeNodeMultipleChoiceFilter 

14 

15__all__ = ( 

16 'BaseFilterSetTestMixin', 

17 'ChangeLoggedFilterSetTestMixin', 

18) 

19 

20EXEMPT_MODEL_FIELDS = ( 

21 'comments', 

22 'custom_field_data', 

23 'path', # ltree, trigger-maintained 

24 'sort_path', # ltree, trigger-maintained 

25) 

26 

27 

28class BaseFilterSetTestMixin: 

29 queryset = None 

30 filterset = None 

31 ignore_fields = tuple() 

32 filter_name_map = {} 

33 

34 def get_m2m_filter_name(self, field): 

35 """ 

36 Given a ManyToManyField, determine the correct name for its corresponding Filter. Individual test 

37 cases may override this method to prescribe deviations for specific fields. 

38 """ 

39 related_model_name = field.related_model._meta.verbose_name 

40 return related_model_name.lower().replace(' ', '_') 

41 

42 def get_filters_for_model_field(self, field): 

43 """ 

44 Given a model field, return an iterable of (name, class) for each filter that should be defined on 

45 the model's FilterSet class. If the appropriate filter class cannot be determined, it will be None. 

46 

47 filter_name_map provides a mechanism for developers to provide an actual field name for the 

48 filter that is being resolved, given the field's actual name. 

49 """ 

50 # If an alias is not present in filter_name_map, then use field.name 

51 filter_name = self.filter_name_map.get(field.name, field.name) 

52 

53 # ForeignKey & OneToOneField 

54 if issubclass(field.__class__, ForeignKey) or type(field) is OneToOneRel: 

55 

56 # Relationships to ContentType (used as part of a GFK) do not need a filter 

57 if field.related_model is ContentType: 

58 return [(None, None)] 

59 

60 # ForeignKey to an ltree-backed hierarchical model 

61 if issubclass(field.related_model, LtreeModel) and field.model is not field.related_model: 

62 return [(f'{filter_name}_id', TreeNodeMultipleChoiceFilter)] 

63 

64 return [(f'{filter_name}_id', django_filters.ModelMultipleChoiceFilter)] 

65 

66 # Many-to-many relationships (forward & backward) 

67 if type(field) in (ManyToManyField, ManyToManyRel): 

68 filter_name = self.get_m2m_filter_name(field) 

69 filter_name = self.filter_name_map.get(filter_name, filter_name) 

70 

71 # ManyToManyFields to ContentType need two filters: 'app.model' & PK 

72 if field.related_model is ContentType: 

73 # Standardize on object_type for filter name even though it's technically a ContentType 

74 filter_name = 'object_type' 

75 return [ 

76 (filter_name, MultiValueContentTypeFilter), 

77 (f'{filter_name}_id', django_filters.ModelMultipleChoiceFilter), 

78 ] 

79 

80 return [(f'{filter_name}_id', django_filters.ModelMultipleChoiceFilter)] 

81 

82 # Tag manager 

83 if isinstance(field, TaggableManager): 

84 return [('tag', TagFilter)] 

85 

86 # Unable to determine the correct filter class 

87 return [(filter_name, None)] 

88 

89 def test_id(self): 

90 """ 

91 Test filtering for two PKs from a set of >2 objects. 

92 """ 

93 params = {'id': self.queryset.values_list('pk', flat=True)[:2]} 

94 self.assertGreater(self.queryset.count(), 2) 

95 self.assertEqual(self.filterset(params, self.queryset).qs.count(), 2) 

96 

97 def test_missing_filters(self): 

98 """ 

99 Check for any model fields which do not have the required filter(s) defined. 

100 """ 

101 app_label = self.__class__.__module__.split('.')[0] 

102 model = self.queryset.model 

103 model_name = model.__name__ 

104 

105 # Import the FilterSet class & sanity check it 

106 filterset = import_string(f'{app_label}.filtersets.{model_name}FilterSet') 

107 self.assertEqual(model, filterset.Meta.model, "FilterSet model does not match!") 

108 

109 filters = filterset.get_filters() 

110 

111 # Check for missing filters 

112 for model_field in model._meta.get_fields(): 

113 

114 # Skip private fields 

115 if model_field.name.startswith('_'): 

116 continue 

117 

118 # Skip ignored fields 

119 if model_field.name in chain(self.ignore_fields, EXEMPT_MODEL_FIELDS): 

120 continue 

121 

122 # Skip reverse ForeignKey relationships 

123 if type(model_field) is ManyToOneRel: 

124 continue 

125 

126 # Skip generic relationships 

127 if type(model_field) in (GenericForeignKey, GenericRelation): 

128 continue 

129 

130 for filter_name, filter_class in self.get_filters_for_model_field(model_field): 

131 

132 if filter_name is None: 

133 # Field is exempt 

134 continue 

135 

136 # Check that the filter is defined 

137 self.assertIn( 

138 filter_name, 

139 filters.keys(), 

140 f'No filter defined for {filter_name} ({model_field.name})!' 

141 ) 

142 

143 # Check that the filter class is correct 

144 filter = filters[filter_name] 

145 if filter_class is not None: 

146 self.assertIsInstance( 

147 filter, 

148 filter_class, 

149 f"Invalid filter class {type(filter)} for {filter_name} (should be {filter_class})!" 

150 ) 

151 

152 

153class ChangeLoggedFilterSetTestMixin(BaseFilterSetTestMixin): 

154 

155 def test_created(self): 

156 pk_list = self.queryset.values_list('pk', flat=True)[:2] 

157 self.queryset.filter(pk__in=pk_list).update(created=datetime(2021, 1, 1, 0, 0, 0, tzinfo=UTC)) 

158 params = {'created': ['2021-01-01T00:00:00']} 

159 self.assertEqual(self.filterset(params, self.queryset).qs.count(), 2) 

160 

161 def test_last_updated(self): 

162 pk_list = self.queryset.values_list('pk', flat=True)[:2] 

163 self.queryset.filter(pk__in=pk_list).update(last_updated=datetime(2021, 1, 2, 0, 0, 0, tzinfo=UTC)) 

164 params = {'last_updated': ['2021-01-02T00:00:00']} 

165 self.assertEqual(self.filterset(params, self.queryset).qs.count(), 2)