Skip to content

Commit 225fcd9

Browse files
authored
Merge pull request #490 from Hritik14/velocity-test_upstream
Speed up upstream tests
2 parents 09839b9 + f6c1635 commit 225fcd9

5 files changed

Lines changed: 44 additions & 38 deletions

File tree

vulnerabilities/importers/elixir_security.py

Lines changed: 1 addition & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -49,16 +49,7 @@ def set_api(self, packages):
4949
asyncio.run(self.pkg_manager_api.load_api(packages))
5050

5151
def updated_advisories(self) -> Set[Advisory]:
52-
files = self._updated_files
53-
advisories = []
54-
for f in files:
55-
processed_data = self.process_file(f)
56-
if processed_data:
57-
advisories.append(processed_data)
58-
return self.batch_advisories(advisories)
59-
60-
def added_advisories(self) -> Set[Advisory]:
61-
files = self._added_files
52+
files = self._updated_files.union(self._added_files)
6253
advisories = []
6354
for f in files:
6455
processed_data = self.process_file(f)

vulnerabilities/importers/retiredotnet.py

Lines changed: 1 addition & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -43,16 +43,7 @@ def __enter__(self):
4343
)
4444

4545
def updated_advisories(self) -> Set[Advisory]:
46-
files = self._updated_files
47-
advisories = []
48-
for f in files:
49-
processed_data = self.process_file(f)
50-
if processed_data:
51-
advisories.append(processed_data)
52-
return self.batch_advisories(advisories)
53-
54-
def added_advisories(self) -> Set[Advisory]:
55-
files = self._added_files
46+
files = self._updated_files.union(self._added_files)
5647
advisories = []
5748
for f in files:
5849
processed_data = self.process_file(f)

vulnerabilities/importers/ruby.py

Lines changed: 1 addition & 10 deletions
Original file line numberDiff line numberDiff line change
@@ -54,16 +54,7 @@ def set_api(self, packages):
5454
asyncio.run(self.pkg_manager_api.load_api(packages))
5555

5656
def updated_advisories(self) -> Set[Advisory]:
57-
files = self._updated_files
58-
advisories = []
59-
for f in files:
60-
processed_data = self.process_file(f)
61-
if processed_data:
62-
advisories.append(processed_data)
63-
return self.batch_advisories(advisories)
64-
65-
def added_advisories(self) -> Set[Advisory]:
66-
files = self._added_files
57+
files = self._updated_files.union(self._added_files)
6758
advisories = []
6859
for f in files:
6960
processed_data = self.process_file(f)

vulnerabilities/importers/rust.py

Lines changed: 1 addition & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -62,11 +62,8 @@ def crates_api(self):
6262
def set_api(self, packages):
6363
asyncio.run(self.crates_api.load_api(packages))
6464

65-
def added_advisories(self) -> Set[Advisory]:
66-
return self._load_advisories(self._added_files)
67-
6865
def updated_advisories(self) -> Set[Advisory]:
69-
return self._load_advisories(self._updated_files)
66+
return self._load_advisories(self._updated_files.union(self._added_files))
7067

7168
def _load_advisories(self, files) -> Set[Advisory]:
7269
# per @tarcieri It will always be named RUSTSEC-0000-0000.md
Lines changed: 40 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,18 +1,54 @@
1+
import inspect
2+
from unittest.mock import patch
3+
14
import pytest
5+
26
from vulnerabilities import importers
7+
from vulnerabilities.data_source import Advisory
38
from vulnerabilities.importer_yielder import IMPORTER_REGISTRY
49

10+
MAX_ADVISORIES = 1
11+
12+
13+
class MaxAdvisoriesCreatedInterrupt(BaseException):
14+
# Inheriting BaseException is intentional because the function being tested might catch Exception
15+
pass
16+
517

618
@pytest.mark.webtest
719
@pytest.mark.parametrize(
820
("data_source", "config"),
921
((data["data_source"], data["data_source_cfg"]) for data in IMPORTER_REGISTRY),
1022
)
1123
def test_updated_advisories(data_source, config):
12-
1324
if not data_source == "GitHubAPIDataSource":
1425
data_src = getattr(importers, data_source)
15-
data_src = data_src(batch_size=1, config=config)
16-
with data_src:
17-
for i in data_src.updated_advisories():
26+
data_src = data_src(batch_size=MAX_ADVISORIES, config=config)
27+
advisory_counter = 0
28+
29+
def patched_advisory(*args, **kwargs):
30+
nonlocal advisory_counter
31+
32+
if advisory_counter >= MAX_ADVISORIES:
33+
raise MaxAdvisoriesCreatedInterrupt
34+
35+
advisory_counter += 1
36+
return Advisory(*args, **kwargs)
37+
38+
module = inspect.getmodule(data_src)
39+
module_members = [m[0] for m in inspect.getmembers(module)]
40+
advisory_class = f"{module.__name__}.Advisory"
41+
if "Advisory" not in module_members:
42+
advisory_class = "vulnerabilities.data_source.Advisory"
43+
44+
# Either
45+
# 1) Advisory class is successfully patched and MaxAdvisoriesCreatedInterrupt is thrown when
46+
# an importer tries to create an Advisory or
47+
# 2) Importer somehow bypasses the patch / handles BaseException internally, then
48+
# updated_advisories is required to return non zero advisories
49+
with patch(advisory_class, side_effect=patched_advisory):
50+
try:
51+
with data_src:
52+
assert len(list(data_src.updated_advisories())) > 0
53+
except MaxAdvisoriesCreatedInterrupt:
1854
pass

0 commit comments

Comments
 (0)