|
9 | 9 |
|
10 | 10 | from urllib.parse import unquote |
11 | 11 |
|
| 12 | +import rest_framework_filters |
12 | 13 | from django.db.models import Prefetch |
13 | | -from django.db.models import Q |
14 | 14 | from django_filters import rest_framework as filters |
15 | 15 | from packageurl import PackageURL |
16 | 16 | from rest_framework import serializers |
@@ -159,7 +159,7 @@ class PackageFilterSet(filters.FilterSet): |
159 | 159 |
|
160 | 160 | class Meta: |
161 | 161 | model = Package |
162 | | - fields = ["name", "type", "version", "subpath", "purl"] |
| 162 | + fields = ["name", "type", "version", "subpath", "purl", "packagerelatedvulnerability__fix"] |
163 | 163 |
|
164 | 164 | def filter_purl(self, queryset, name, value): |
165 | 165 | purl = unquote(value) |
@@ -217,33 +217,55 @@ def bulk_search(self, request): |
217 | 217 | return Response(response) |
218 | 218 |
|
219 | 219 |
|
| 220 | +class PackageFilter(filters.FilterSet): |
| 221 | + class Meta: |
| 222 | + model = Package |
| 223 | + fields = { |
| 224 | + "name": ["exact", "in", "startswith"], |
| 225 | + "namespace": ["exact", "in", "startswith"], |
| 226 | + "type": ["exact", "in", "startswith"], |
| 227 | + } |
| 228 | + |
| 229 | + |
220 | 230 | class VulnerabilityFilterSet(filters.FilterSet): |
| 231 | + |
| 232 | + package = rest_framework_filters.RelatedFilter( |
| 233 | + PackageFilter, field_name="packages", queryset=Package.objects.all() |
| 234 | + ) |
| 235 | + |
221 | 236 | class Meta: |
222 | 237 | model = Vulnerability |
223 | 238 | fields = ["vulnerability_id"] |
224 | 239 |
|
225 | 240 |
|
226 | 241 | class VulnerabilityViewSet(viewsets.ReadOnlyModelViewSet): |
| 242 | + def get_fixed_packages_qs(self): |
| 243 | + """ |
| 244 | + Filter the packages that fixes a vulnerability |
| 245 | + on fields like name, namespace and type. |
| 246 | + """ |
| 247 | + package_filter_data = {"packagerelatedvulnerability__fix": True} |
| 248 | + |
| 249 | + query_params = self.request.query_params |
| 250 | + for field_name in ["name", "namespace", "type"]: |
| 251 | + value = query_params.get(field_name) |
| 252 | + if value: |
| 253 | + package_filter_data[field_name] = value |
| 254 | + |
| 255 | + return PackageFilterSet(package_filter_data).qs |
| 256 | + |
227 | 257 | def get_queryset(self): |
228 | | - params = self.request.query_params |
229 | | - query = Q() |
230 | | - name = params.get("name") |
231 | | - if name: |
232 | | - query &= Q(name=name) |
233 | | - namespace = params.get("namespace") |
234 | | - if namespace: |
235 | | - query &= Q(namespace=namespace) |
236 | | - type = params.get("type") |
237 | | - if type: |
238 | | - query &= Q(type=type) |
239 | | - queryset = Vulnerability.objects.prefetch_related( |
| 258 | + """ |
| 259 | + Assign filtered packages queryset from `get_fixed_packages_qs` |
| 260 | + to a custom attribute `filtered_fixed_packages` |
| 261 | + """ |
| 262 | + return Vulnerability.objects.prefetch_related( |
240 | 263 | Prefetch( |
241 | 264 | "packages", |
242 | | - queryset=Package.objects.filter(query, packagerelatedvulnerability__fix=True), |
| 265 | + queryset=self.get_fixed_packages_qs(), |
243 | 266 | to_attr="filtered_fixed_packages", |
244 | 267 | ) |
245 | 268 | ) |
246 | | - return queryset |
247 | 269 |
|
248 | 270 | serializer_class = VulnerabilitySerializer |
249 | 271 | paginate_by = 50 |
|
0 commit comments