Skip to content
Merged
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
11 changes: 1 addition & 10 deletions vulnerabilities/importers/elixir_security.py
Original file line number Diff line number Diff line change
Expand Up @@ -49,16 +49,7 @@ def set_api(self, packages):
asyncio.run(self.pkg_manager_api.load_api(packages))

def updated_advisories(self) -> Set[Advisory]:
files = self._updated_files
advisories = []
for f in files:
processed_data = self.process_file(f)
if processed_data:
advisories.append(processed_data)
return self.batch_advisories(advisories)

def added_advisories(self) -> Set[Advisory]:
files = self._added_files
files = self._updated_files.union(self._added_files)
advisories = []
for f in files:
processed_data = self.process_file(f)
Expand Down
11 changes: 1 addition & 10 deletions vulnerabilities/importers/retiredotnet.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,16 +43,7 @@ def __enter__(self):
)

def updated_advisories(self) -> Set[Advisory]:
files = self._updated_files
advisories = []
for f in files:
processed_data = self.process_file(f)
if processed_data:
advisories.append(processed_data)
return self.batch_advisories(advisories)

def added_advisories(self) -> Set[Advisory]:
files = self._added_files
files = self._updated_files.union(self._added_files)
advisories = []
for f in files:
processed_data = self.process_file(f)
Expand Down
11 changes: 1 addition & 10 deletions vulnerabilities/importers/ruby.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,16 +54,7 @@ def set_api(self, packages):
asyncio.run(self.pkg_manager_api.load_api(packages))

def updated_advisories(self) -> Set[Advisory]:
files = self._updated_files
advisories = []
for f in files:
processed_data = self.process_file(f)
if processed_data:
advisories.append(processed_data)
return self.batch_advisories(advisories)

def added_advisories(self) -> Set[Advisory]:
files = self._added_files
files = self._updated_files.union(self._added_files)
advisories = []
for f in files:
processed_data = self.process_file(f)
Expand Down
5 changes: 1 addition & 4 deletions vulnerabilities/importers/rust.py
Original file line number Diff line number Diff line change
Expand Up @@ -62,11 +62,8 @@ def crates_api(self):
def set_api(self, packages):
asyncio.run(self.crates_api.load_api(packages))

def added_advisories(self) -> Set[Advisory]:
return self._load_advisories(self._added_files)

def updated_advisories(self) -> Set[Advisory]:
return self._load_advisories(self._updated_files)
return self._load_advisories(self._updated_files.union(self._added_files))

def _load_advisories(self, files) -> Set[Advisory]:
# per @tarcieri It will always be named RUSTSEC-0000-0000.md
Expand Down
44 changes: 40 additions & 4 deletions vulnerabilities/tests/test_upstream.py
Original file line number Diff line number Diff line change
@@ -1,18 +1,54 @@
import inspect
from unittest.mock import patch

import pytest

from vulnerabilities import importers
from vulnerabilities.data_source import Advisory
from vulnerabilities.importer_yielder import IMPORTER_REGISTRY

MAX_ADVISORIES = 1


class MaxAdvisoriesCreatedInterrupt(BaseException):
# Inheriting BaseException is intentional because the function being tested might catch Exception
pass


@pytest.mark.webtest
@pytest.mark.parametrize(
("data_source", "config"),
((data["data_source"], data["data_source_cfg"]) for data in IMPORTER_REGISTRY),
)
def test_updated_advisories(data_source, config):

if not data_source == "GitHubAPIDataSource":
data_src = getattr(importers, data_source)
data_src = data_src(batch_size=1, config=config)
with data_src:
for i in data_src.updated_advisories():
data_src = data_src(batch_size=MAX_ADVISORIES, config=config)
advisory_counter = 0

def patched_advisory(*args, **kwargs):
nonlocal advisory_counter

if advisory_counter >= MAX_ADVISORIES:
raise MaxAdvisoriesCreatedInterrupt

advisory_counter += 1
return Advisory(*args, **kwargs)

module = inspect.getmodule(data_src)
module_members = [m[0] for m in inspect.getmembers(module)]
advisory_class = f"{module.__name__}.Advisory"
if "Advisory" not in module_members:
advisory_class = "vulnerabilities.data_source.Advisory"

# Either
# 1) Advisory class is successfully patched and MaxAdvisoriesCreatedInterrupt is thrown when
# an importer tries to create an Advisory or
# 2) Importer somehow bypasses the patch / handles BaseException internally, then
# updated_advisories is required to return non zero advisories
with patch(advisory_class, side_effect=patched_advisory):
try:
with data_src:
assert len(list(data_src.updated_advisories())) > 0
except MaxAdvisoriesCreatedInterrupt:
pass