Skip to content

Commit 400baec

Browse files
committed
Add prefetch in models
Signed-off-by: Tushar Goel <tushar.goel.dav@gmail.com>
1 parent 7e2a70a commit 400baec

2 files changed

Lines changed: 108 additions & 59 deletions

File tree

vulnerabilities/api.py

Lines changed: 53 additions & 54 deletions
Original file line numberDiff line numberDiff line change
@@ -1,12 +1,3 @@
1-
#
2-
# Copyright (c) nexB Inc. and others. All rights reserved.
3-
# VulnerableCode is a trademark of nexB Inc.
4-
# SPDX-License-Identifier: Apache-2.0
5-
# See http://www.apache.org/licenses/LICENSE-2.0 for the license text.
6-
# See https://github.com/nexB/vulnerablecode for support or download.
7-
# See https://aboutcode.org for more information about nexB OSS projects.
8-
#
9-
101
from urllib.parse import unquote
112

123
from django.db.models import Prefetch
@@ -24,14 +15,7 @@
2415
from rest_framework.throttling import AnonRateThrottle
2516
from rest_framework.throttling import UserRateThrottle
2617

27-
from vulnerabilities.models import Alias
28-
from vulnerabilities.models import Kev
29-
from vulnerabilities.models import Package
30-
from vulnerabilities.models import Vulnerability
31-
from vulnerabilities.models import VulnerabilityReference
32-
from vulnerabilities.models import VulnerabilitySeverity
33-
from vulnerabilities.models import Weakness
34-
from vulnerabilities.models import get_purl_query_lookups
18+
from vulnerabilities.models import Alias, Kev, Package, PackageRelatedVulnerability, Vulnerability, VulnerabilityReference, VulnerabilitySeverity, Weakness, get_purl_query_lookups
3519
from vulnerabilities.throttling import StaffUserRateThrottle
3620

3721

@@ -63,9 +47,6 @@ def get_fields(self):
6347
def get_resource_url(self, instance):
6448
"""
6549
Return the instance fully qualified URL including the schema and domain.
66-
67-
Usage:
68-
resource_url = serializers.SerializerMethodField()
6950
"""
7051
resource_url = instance.get_absolute_url()
7152

@@ -146,7 +127,7 @@ def to_representation(self, instance):
146127

147128
class Meta:
148129
model = Vulnerability
149-
fields = ["url", "vulnerability_id", "summary", "references", "fixed_packages", "aliases"]
130+
fields = ["url", "vulnerability_id", "summary", "fixed_packages", "references", "aliases"]
150131

151132

152133
class WeaknessSerializer(serializers.HyperlinkedModelSerializer):
@@ -224,25 +205,21 @@ def to_representation(self, instance):
224205
return data
225206

226207
next_non_vulnerable_version = serializers.SerializerMethodField("get_next_non_vulnerable")
208+
latest_non_vulnerable_version = serializers.SerializerMethodField("get_latest_non_vulnerable")
209+
purl = serializers.CharField(source="package_url")
210+
affected_by_vulnerabilities = serializers.SerializerMethodField("get_affected_vulnerabilities")
211+
fixing_vulnerabilities = serializers.SerializerMethodField("get_fixed_vulnerabilities")
227212

228213
def get_next_non_vulnerable(self, package):
229214
next_non_vulnerable = package.fixed_package_details.get("next_non_vulnerable", None)
230215
if next_non_vulnerable:
231216
return next_non_vulnerable.version
232217

233-
latest_non_vulnerable_version = serializers.SerializerMethodField("get_latest_non_vulnerable")
234-
235218
def get_latest_non_vulnerable(self, package):
236219
latest_non_vulnerable = package.fixed_package_details.get("latest_non_vulnerable", None)
237220
if latest_non_vulnerable:
238221
return latest_non_vulnerable.version
239222

240-
purl = serializers.CharField(source="package_url")
241-
242-
affected_by_vulnerabilities = serializers.SerializerMethodField("get_affected_vulnerabilities")
243-
244-
fixing_vulnerabilities = serializers.SerializerMethodField("get_fixed_vulnerabilities")
245-
246223
def get_fixed_packages(self, package):
247224
"""
248225
Return a queryset of all packages that fix a vulnerability with
@@ -263,20 +240,30 @@ def get_vulnerabilities_for_a_package(self, package, fix) -> dict:
263240
Return vulnerabilities that affect the `package` if given `fix` flag is False,
264241
otherwise return vulnerabilities fixed by the `package`.
265242
"""
266-
fixed_packages = self.get_fixed_packages(package=package)
267-
qs = package.vulnerabilities.filter(packagerelatedvulnerability__fix=fix)
268-
qs = qs.prefetch_related(
269-
Prefetch(
270-
"packages",
271-
queryset=fixed_packages,
272-
to_attr="filtered_fixed_packages",
273-
)
274-
)
275-
return VulnSerializerRefsAndSummary(
276-
instance=qs,
277-
many=True,
278-
context={"request": self.context["request"]},
279-
).data
243+
# Retrieve all related PackageRelatedVulnerability objects for this package
244+
related_vulnerabilities = PackageRelatedVulnerability.objects.filter(
245+
package_id=package.id,
246+
fix=fix
247+
).select_related('vulnerability')
248+
249+
# Cache the vulnerabilities to avoid duplicated queries
250+
vulnerabilities = {
251+
prv.vulnerability: prv.vulnerability for prv in related_vulnerabilities
252+
}
253+
254+
# Now, map the vulnerabilities with their fixed packages
255+
vulnerabilities_with_fixed_packages = []
256+
for vulnerability in vulnerabilities.values():
257+
vuln_data = VulnSerializerRefsAndSummary(
258+
instance=vulnerability,
259+
context={"request": self.context["request"]},
260+
).data
261+
262+
vuln_data["fixed_packages"] = [pkg.purl for pkg in vulnerability.fixed_packages_for_vuln(package=package)]
263+
vulnerabilities_with_fixed_packages.append(vuln_data)
264+
265+
return vulnerabilities_with_fixed_packages
266+
280267

281268
def get_fixed_vulnerabilities(self, package) -> dict:
282269
"""
@@ -292,15 +279,15 @@ def get_affected_vulnerabilities(self, package) -> dict:
292279
excluded_purls = []
293280
package_vulnerabilities = self.get_vulnerabilities_for_a_package(package=package, fix=False)
294281

295-
for vuln in package_vulnerabilities:
296-
for pkg in vuln["fixed_packages"]:
297-
real_purl = PackageURL.from_string(pkg["purl"])
298-
if package.version_class(real_purl.version) <= package.current_version:
299-
excluded_purls.append(pkg)
282+
# for vuln in package_vulnerabilities:
283+
# for pkg in vuln["fixed_packages"]:
284+
# real_purl = PackageURL.from_string(pkg["purl"])
285+
# if package.version_class(real_purl.version) <= package.current_version:
286+
# excluded_purls.append(pkg["purl"])
300287

301-
vuln["fixed_packages"] = [
302-
pkg for pkg in vuln["fixed_packages"] if pkg not in excluded_purls
303-
]
288+
# vuln["fixed_packages"] = [
289+
# pkg for pkg in vuln["fixed_packages"] if pkg["purl"] not in excluded_purls
290+
# ]
304291

305292
return package_vulnerabilities
306293

@@ -377,7 +364,21 @@ class PackageViewSet(viewsets.ReadOnlyModelViewSet):
377364
Lookup for vulnerable packages by Package URL.
378365
"""
379366

380-
queryset = Package.objects.all()
367+
queryset = Package.objects.all().prefetch_related(
368+
Prefetch(
369+
'vulnerabilities',
370+
queryset=Vulnerability.objects.prefetch_related(
371+
Prefetch(
372+
'references',
373+
queryset=VulnerabilityReference.objects.all()
374+
),
375+
'aliases',
376+
'weaknesses',
377+
)
378+
),
379+
'packagerelatedvulnerability_set',
380+
)
381+
# queryset = Package.objects.all()
381382
serializer_class = PackageSerializer
382383
filter_backends = (filters.DjangoFilterBackend,)
383384
filterset_class = PackageFilterSet
@@ -436,8 +437,6 @@ def bulk_search(self, request):
436437
PackageSerializer(query, many=True, context={"request": request}).data
437438
)
438439

439-
# using order by and distinct because there will be
440-
# many fully qualified purl for a single plain purl
441440
vulnerable_purls = query.vulnerable().only("plain_package_url")
442441
vulnerable_purls = [str(package.plain_package_url) for package in vulnerable_purls]
443442
return Response(data=vulnerable_purls)

vulnerabilities/models.py

Lines changed: 55 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -219,6 +219,16 @@ def affected_packages(self):
219219
Return a queryset of packages that are affected by this vulnerability.
220220
"""
221221
return self.packages.affected()
222+
223+
def fixed_packages_for_vuln(self, package):
224+
"""
225+
Return a queryset of packages that are fixing this vulnerability.
226+
"""
227+
packages = []
228+
for pkg in package.fixed_packages.all():
229+
if pkg.vulnerabilities.filter(id=self.id).exists():
230+
packages.append(pkg)
231+
return packages
222232

223233
# legacy aliases
224234
vulnerable_packages = affected_packages
@@ -465,9 +475,12 @@ def fixing_packages(self, package, with_qualifiers_and_subpath=True):
465475
``package``.
466476
"""
467477

468-
return self.match_purl(
469-
purl=package.purl,
470-
with_qualifiers_and_subpath=with_qualifiers_and_subpath,
478+
return self.filter(
479+
name=package.name,
480+
namespace=package.namespace,
481+
type=package.type,
482+
qualifiers=package.qualifiers,
483+
subpath=package.subpath,
471484
).fixing()
472485

473486
def search(self, query: str = None):
@@ -616,7 +629,31 @@ def fixing(self):
616629
"""
617630
Return a queryset of vulnerabilities fixed by this package.
618631
"""
619-
return self.vulnerabilities.filter(packagerelatedvulnerability__fix=True)
632+
return self.vulnerabilities.filter(packagerelatedvulnerability__fix=True).prefetch_related(
633+
Prefetch(
634+
"references",
635+
queryset=VulnerabilityReference.objects.all(),
636+
),
637+
"aliases",
638+
"weaknesses",
639+
Prefetch(
640+
"packages",
641+
queryset=Package.objects.prefetch_related(
642+
Prefetch(
643+
"vulnerabilities",
644+
queryset=Vulnerability.objects.prefetch_related(
645+
Prefetch(
646+
"references",
647+
queryset=VulnerabilityReference.objects.all(),
648+
),
649+
"aliases",
650+
"weaknesses",
651+
),
652+
),
653+
"packagerelatedvulnerability_set",
654+
).distinct(),
655+
)
656+
)
620657

621658
# legacy aliases
622659
resolved_to = fixing
@@ -626,7 +663,20 @@ def fixed_packages(self):
626663
"""
627664
Return a queryset of packages that are fixed.
628665
"""
629-
return Package.objects.fixing_packages(package=self).distinct()
666+
return Package.objects.fixing_packages(package=self).prefetch_related(
667+
Prefetch(
668+
'vulnerabilities',
669+
queryset=Vulnerability.objects.prefetch_related(
670+
Prefetch(
671+
'references',
672+
queryset=VulnerabilityReference.objects.all()
673+
),
674+
'aliases',
675+
'weaknesses',
676+
)
677+
),
678+
'packagerelatedvulnerability_set',
679+
).distinct()
630680

631681
@property
632682
def is_vulnerable(self) -> bool:

0 commit comments

Comments
 (0)