Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 6 additions & 1 deletion vulnerabilities/api_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -205,6 +205,10 @@ def get_affected_by_vulnerabilities(self, obj):
purl = None
if fixed_by_package:
purl = fixed_by_package.package_url

#exposing severities inside the affectedbyvulnerabilities
severities = VulnerabilitySeverityV2Serializer(vuln.severities.all(), many=True).data

# Get code fixed for a vulnerability
code_fixes = CodeFix.objects.filter(
affected_package_vulnerability__vulnerability=vuln
Expand All @@ -216,6 +220,7 @@ def get_affected_by_vulnerabilities(self, obj):

result[vuln.vulnerability_id] = {
"vulnerability_id": vuln.vulnerability_id,
"severities" : severities,
"fixed_by_packages": purl,
"code_fixes": code_fix_urls,
}
Expand Down Expand Up @@ -260,7 +265,7 @@ class PackageV2ViewSet(viewsets.ReadOnlyModelViewSet):
queryset = Package.objects.all().prefetch_related(
Prefetch(
"affected_by_vulnerabilities",
queryset=Vulnerability.objects.prefetch_related("fixed_by_packages"),
queryset=Vulnerability.objects.prefetch_related("fixed_by_packages","severities"),
to_attr="prefetched_affected_vulnerabilities",
)
)
Expand Down
123 changes: 110 additions & 13 deletions vulnerabilities/tests/test_api_v2.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
from vulnerabilities.models import ApiUser
from vulnerabilities.models import Package
from vulnerabilities.models import Vulnerability
from vulnerabilities.models import VulnerabilitySeverity
from vulnerabilities.models import VulnerabilityReference
from vulnerabilities.models import Weakness

Expand Down Expand Up @@ -210,13 +211,23 @@ def setUp(self):
self.client = APIClient(enforce_csrf_checks=True)
self.client.credentials(HTTP_AUTHORIZATION=self.auth)

#create vulnerability severities
self.severity = VulnerabilitySeverity.objects.create(
scoring_system="CVSSv3",
scoring_elements="",
url="https://example.com",
value="7.5"
)

self.vuln1.severities.add(self.severity)

def test_list_packages(self):
"""
Test listing packages without filters.
Should return a list of packages with their details and associated vulnerabilities.
"""
url = reverse("package-v2-list")
with self.assertNumQueries(32):
with self.assertNumQueries(33):
response = self.client.get(url, format="json")
self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertIn("results", response.data)
Expand All @@ -238,7 +249,7 @@ def test_filter_packages_by_purl(self):
Test filtering packages by one or more PURLs.
"""
url = reverse("package-v2-list")
with self.assertNumQueries(20):
with self.assertNumQueries(21):
response = self.client.get(url, {"purl": "pkg:pypi/django@3.2"}, format="json")
self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual(len(response.data["results"]["packages"]), 1)
Expand All @@ -249,7 +260,7 @@ def test_filter_packages_by_affected_vulnerability(self):
Test filtering packages by affected_by_vulnerability.
"""
url = reverse("package-v2-list")
with self.assertNumQueries(20):
with self.assertNumQueries(21):
response = self.client.get(
url, {"affected_by_vulnerability": "VCID-1234"}, format="json"
)
Expand All @@ -275,13 +286,22 @@ def test_package_serializer_fields(self):
# Fetch the package
package = Package.objects.get(package_url="pkg:pypi/django@3.2")

#retrieving the vulnerability
Vulnerability_with_severities = Vulnerability.objects.get(vulnerability_id="VCID-1234")
Vulnerability_without_severities = Vulnerability.objects.get(vulnerability_id="VCID-5678")

severity_created = self.severity
Vulnerability_with_severities.severities.add(severity_created)

package.affected_by_vulnerabilities.add(Vulnerability_with_severities, Vulnerability_without_severities)

# Ensure prefetched data is available for the serializer
package = (
Package.objects.filter(package_url="pkg:pypi/django@3.2")
.prefetch_related(
Prefetch(
"affected_by_vulnerabilities",
queryset=Vulnerability.objects.prefetch_related("fixed_by_packages"),
queryset=Vulnerability.objects.prefetch_related("fixed_by_packages","severities"),
to_attr="prefetched_affected_vulnerabilities",
)
)
Expand Down Expand Up @@ -311,8 +331,22 @@ def test_package_serializer_fields(self):
"VCID-1234": {
"code_fixes": [],
"vulnerability_id": "VCID-1234",
"severities": [
{
"url":"https://example.com",
"value": "7.5",
"scoring_system": "CVSSv3",
"scoring_elements": "",
}
],
"fixed_by_packages": None,
}
},
"VCID-5678": {
"code_fixes": [],
"vulnerability_id": "VCID-5678",
"severities": [],
"fixed_by_packages": "pkg:npm/lodash@4.17.20",
},
}
self.assertEqual(data["affected_by_vulnerabilities"], expected_affected_by_vulnerabilities)

Expand Down Expand Up @@ -375,30 +409,61 @@ def test_get_affected_by_vulnerabilities(self):
"""
Test the get_affected_by_vulnerabilities method in the serializer.
"""
serializer = PackageV2Serializer()
package = Package.objects.get(package_url="pkg:pypi/django@3.2")

vulnerability_with_severities, vulnerability_without_severities = list(Vulnerability.objects.filter(vulnerability_id__in=["VCID-1234", "VCID-5678"]
).prefetch_related("fixed_by_packages", "severities"))
package.affected_by_vulnerabilities.add(vulnerability_with_severities, vulnerability_without_severities)

package = (
Package.objects.filter(package_url="pkg:pypi/django@3.2")
.prefetch_related(
Prefetch(
"affected_by_vulnerabilities",
queryset=Vulnerability.objects.prefetch_related("fixed_by_packages"),
queryset=Vulnerability.objects.prefetch_related("fixed_by_packages","severities"),
to_attr="prefetched_affected_vulnerabilities",
)
)
.first()
)

serializer = PackageV2Serializer()
severity_created = self.severity
vulnerability_with_severities.severities.add(severity_created)

vulnerabilities = serializer.get_affected_by_vulnerabilities(package)


for vuln_data in vulnerabilities.values():
for severity in vuln_data["severities"]:
self.assertIn("url", severity)
self.assertIn("value", severity)
self.assertIn("scoring_system", severity)
self.assertIn("scoring_elements", severity)

self.assertEqual(
vulnerabilities,
{
"VCID-1234": {
"code_fixes": [],
"vulnerability_id": "VCID-1234",
"severities": [
{
"url":"https://example.com",
"value": "7.5",
"scoring_system": "CVSSv3",
"scoring_elements": "",
}
],
"fixed_by_packages": None,
}
},
"VCID-5678": {
"code_fixes": [],
"vulnerability_id": "VCID-5678",
"severities": [], #No severity for this vulnerability
"fixed_by_packages": "pkg:npm/lodash@4.17.20",
},
)
},)

def test_get_fixing_vulnerabilities(self):
"""
Expand Down Expand Up @@ -601,7 +666,25 @@ def test_lookup_with_valid_purl(self):
"""
url = reverse("package-v2-lookup")
data = {"purl": "pkg:pypi/django@3.2"}
with self.assertNumQueries(13):
package = Package.objects.filter(package_url="pkg:pypi/django@3.2").prefetch_related(
Prefetch(
"affected_by_vulnerabilities",
queryset=Vulnerability.objects.prefetch_related(Prefetch("fixed_by_packages",queryset=Package.objects.all()),
Prefetch("severities",queryset=VulnerabilitySeverity.objects.all())),
to_attr="prefetched_affected_vulnerabilities",
)
).first()

vulnerability_with_severities, vulnerability_without_severities = list(Vulnerability.objects.filter(vulnerability_id__in=["VCID-1234", "VCID-5678"]
).prefetch_related("fixed_by_packages", "severities"))

severity_created = self.severity
vulnerability_with_severities.severities.add(severity_created)

package.affected_by_vulnerabilities.add(vulnerability_with_severities, vulnerability_without_severities)


with self.assertNumQueries(15):
response = self.client.post(url, data, format="json")
self.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertEqual(1, len(response.data))
Expand All @@ -617,10 +700,24 @@ def test_lookup_with_valid_purl(self):
"VCID-1234": {
"code_fixes": [],
"vulnerability_id": "VCID-1234",
"severities": [
{
"url":"https://example.com",
"value": "7.5",
"scoring_system": "CVSSv3",
"scoring_elements": '',
}
],
"fixed_by_packages": None,
}
},
)
},
"VCID-5678": {
"code_fixes": [],
"vulnerability_id": "VCID-5678",
"severities": [],
"fixed_by_packages": "pkg:npm/lodash@4.17.20",
},
})

self.assertEqual(response.data[0]["fixing_vulnerabilities"], [])

def test_lookup_with_invalid_purl(self):
Expand Down
45 changes: 41 additions & 4 deletions vulnerabilities/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,7 @@
import logging
import os
import re
import time
import urllib.request
from collections import defaultdict
from functools import total_ordering
Expand Down Expand Up @@ -47,6 +48,42 @@
find_all_cve = cve_regex.findall
cwe_regex = r"CWE-\d+"

# store the last request time for each domain
last_request_times = {}


def polite_requests(url, method="GET", headers=None, data=None, delay=1, max_retries=3):
"""
Make an API request while enforcing politeness (delays, retries, logging).

- Enforces a delay between requests to the same API.
- Retries if a request fails due to rate limits (429 Too Many Requests).
- Logs requests for debugging.
"""
global last_request_times
domain = url.split("/")[2]
last_time = last_request_times.get(domain, 0)
elapsed_time = time.time() - last_time

# enforce a delay before making a requests
if elapsed_time < delay:
time.sleep(delay - elapsed_time)

for attempt in range(max_retries):
try:
response = requests.request(method, url, headers=headers, data=data)
if response.status_code == 429:
retry_after = int(response.headers.get("Retry-After", delay))
logging.warning(f"Rate limited! Retrying after {retry_after} seconds.")
time.sleep(retry_after)
continue # retry again
last_request_times[domain] = time.time()
return response
except requests.exceptions.RequestException as e:
logging.error(f"Request failed: {e}")
time.sleep(2**attempt)
raise Exception(f"Failed to fetch data from {url!r} after {max_retries} attempts.")


@dataclasses.dataclass(order=True, frozen=True)
class AffectedPackage:
Expand All @@ -70,7 +107,7 @@ def load_toml(path):


def fetch_yaml(url):
response = requests.get(url)
response = polite_requests(url)
return saneyaml.load(response.content)


Expand Down Expand Up @@ -265,7 +302,7 @@ def _get_gh_response(gh_token, graphql_query):
"""
endpoint = "https://api.github.com/graphql"
headers = {"Authorization": f"bearer {gh_token}"}
return requests.post(endpoint, headers=headers, json=graphql_query).json()
return polite_requests(endpoint, headers=headers, json=graphql_query).json()


def dedupe(original: List) -> List:
Expand Down Expand Up @@ -365,7 +402,7 @@ def fetch_response(url):
"""
Fetch and return `response` from the `url`
"""
response = requests.get(url)
response = polite_requests(url)
if response.status_code == HTTPStatus.OK:
return response
raise Exception(f"Failed to fetch data from {url!r} with status code: {response.status_code!r}")
Expand All @@ -388,7 +425,7 @@ def plain_purl(purl):


def fetch_and_read_from_csv(url):
response = urllib.request.urlopen(url)
response = polite_requests(url)
lines = [l.decode("utf-8") for l in response.readlines()]
return csv.reader(lines)

Expand Down