diff --git a/vulnerabilities/importers/elixir_security.py b/vulnerabilities/importers/elixir_security.py index 46ef933a8..5f352beca 100644 --- a/vulnerabilities/importers/elixir_security.py +++ b/vulnerabilities/importers/elixir_security.py @@ -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) diff --git a/vulnerabilities/importers/retiredotnet.py b/vulnerabilities/importers/retiredotnet.py index 25e08b04d..7096c9ba6 100644 --- a/vulnerabilities/importers/retiredotnet.py +++ b/vulnerabilities/importers/retiredotnet.py @@ -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) diff --git a/vulnerabilities/importers/ruby.py b/vulnerabilities/importers/ruby.py index ad1e7dc59..883c7c6b3 100644 --- a/vulnerabilities/importers/ruby.py +++ b/vulnerabilities/importers/ruby.py @@ -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) diff --git a/vulnerabilities/importers/rust.py b/vulnerabilities/importers/rust.py index 32d5f07fb..2bcc10f22 100644 --- a/vulnerabilities/importers/rust.py +++ b/vulnerabilities/importers/rust.py @@ -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 diff --git a/vulnerabilities/tests/test_upstream.py b/vulnerabilities/tests/test_upstream.py index a875f570e..c3e7cf392 100644 --- a/vulnerabilities/tests/test_upstream.py +++ b/vulnerabilities/tests/test_upstream.py @@ -1,7 +1,19 @@ +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( @@ -9,10 +21,34 @@ ((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