Skip to content

Commit d2dc98d

Browse files
authored
Merge pull request #291 from sbs2001/improve_github_importer
Improve GitHub importer
2 parents 5e55167 + 5a9e2b6 commit d2dc98d

4 files changed

Lines changed: 66 additions & 30 deletions

File tree

vulnerabilities/importer_yielder.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -179,7 +179,7 @@
179179
'data_source': 'GitHubAPIDataSource',
180180
'data_source_cfg': {
181181
'endpoint': 'https://api.github.com/graphql',
182-
'ecosystems': ['MAVEN', 'NUGET', 'COMPOSER']
182+
'ecosystems': ['MAVEN', 'NUGET', 'COMPOSER', 'PIP', 'RUBYGEMS']
183183
}
184184
},
185185
{

vulnerabilities/importers/github.py

Lines changed: 34 additions & 29 deletions
Original file line numberDiff line numberDiff line change
@@ -42,6 +42,8 @@
4242
from vulnerabilities.package_managers import MavenVersionAPI
4343
from vulnerabilities.package_managers import NugetVersionAPI
4444
from vulnerabilities.package_managers import ComposerVersionAPI
45+
from vulnerabilities.package_managers import PypiVersionAPI
46+
from vulnerabilities.package_managers import RubyVersionAPI
4547

4648
# set of all possible values of first '%s' = {'MAVEN','COMPOSER', 'NUGET'}
4749
# second '%s' is interesting, it will have the value '' for the first request,
@@ -132,6 +134,8 @@ def set_version_api(self, ecosystem: str) -> None:
132134
"MAVEN": MavenVersionAPI,
133135
"NUGET": NugetVersionAPI,
134136
"COMPOSER": ComposerVersionAPI,
137+
"PIP": PypiVersionAPI,
138+
"RUBYGEMS": RubyVersionAPI,
135139
}
136140
versioner = versioners.get(ecosystem)
137141
if versioner:
@@ -147,13 +151,13 @@ def process_name(ecosystem: str, pkg_name: str) -> Optional[Tuple[Optional[str],
147151
ns, name = artifact_comps
148152
return ns, name
149153

150-
if ecosystem == "NUGET":
151-
return None, pkg_name
152-
153154
if ecosystem == "COMPOSER":
154155
vendor, name = pkg_name.split("/")
155156
return vendor, name
156157

158+
if ecosystem == "NUGET" or ecosystem == "PIP" or ecosystem == "RUBYGEMS":
159+
return None, pkg_name
160+
157161
def collect_packages(self, ecosystem):
158162
packages = set()
159163
for page in self.advisories[ecosystem]:
@@ -165,30 +169,29 @@ def process_response(self) -> List[Advisory]:
165169
adv_list = []
166170
for ecosystem in self.advisories:
167171
self.set_version_api(ecosystem)
168-
pkg_type = ecosystem.lower()
172+
pkg_type = self.version_api.package_type
169173
for resp_page in self.advisories[ecosystem]:
170174
for adv in resp_page["data"]["securityVulnerabilities"]["edges"]:
171175
name = adv["node"]["package"]["name"]
172176

173177
if self.process_name(ecosystem, name):
174178
ns, pkg_name = self.process_name(ecosystem, name)
179+
aff_range = adv["node"]["vulnerableVersionRange"]
180+
aff_vers, unaff_vers = self.categorize_versions(
181+
aff_range, self.version_api.get(name)
182+
)
183+
affected_purls = {
184+
PackageURL(name=pkg_name, namespace=ns, version=version, type=pkg_type)
185+
for version in aff_vers
186+
}
187+
188+
unaffected_purls = {
189+
PackageURL(name=pkg_name, namespace=ns, version=version, type=pkg_type)
190+
for version in unaff_vers
191+
}
175192
else:
176-
continue
177-
aff_range = adv["node"]["vulnerableVersionRange"]
178-
aff_vers, unaff_vers = self.categorize_versions(
179-
aff_range, self.version_api.get(name)
180-
)
181-
affected_purls = {
182-
PackageURL(name=pkg_name, namespace=ns,
183-
version=version, type=pkg_type)
184-
for version in aff_vers
185-
}
186-
187-
unaffected_purls = {
188-
PackageURL(name=pkg_name, namespace=ns,
189-
version=version, type=pkg_type)
190-
for version in unaff_vers
191-
}
193+
affected_purls = set()
194+
unaffected_purls = set()
192195

193196
cve_ids = set()
194197
vuln_references = []
@@ -199,12 +202,13 @@ def process_response(self) -> List[Advisory]:
199202
cve_ids.add(vuln["value"])
200203

201204
elif vuln["type"] == "GHSA":
202-
ghsa = vuln['value']
203-
vuln_references.append(Reference(
204-
reference_id=ghsa,
205-
url="https://github.com/advisories/{}".format(
206-
ghsa)
207-
))
205+
ghsa = vuln["value"]
206+
vuln_references.append(
207+
Reference(
208+
reference_id=ghsa,
209+
url="https://github.com/advisories/{}".format(ghsa),
210+
)
211+
)
208212

209213
for cve_id in cve_ids:
210214
adv_list.append(
@@ -219,8 +223,9 @@ def process_response(self) -> List[Advisory]:
219223
return adv_list
220224

221225
@staticmethod
222-
def categorize_versions(version_range: str, all_versions: Set[str]) -> Tuple[Set[str], Set[str]]: # nopep8
226+
def categorize_versions(
227+
version_range: str, all_versions: Set[str]
228+
) -> Tuple[Set[str], Set[str]]: # nopep8
223229
version_range = RangeSpecifier(version_range)
224-
affected_versions = {
225-
version for version in all_versions if version in version_range}
230+
affected_versions = {version for version in all_versions if version in version_range}
226231
return (affected_versions, all_versions - affected_versions)

vulnerabilities/package_managers.py

Lines changed: 30 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -41,6 +41,9 @@ def get(self, package_name: str) -> Set[str]:
4141

4242

4343
class LaunchpadVersionAPI(VersionAPI):
44+
45+
package_type = "deb"
46+
4447
async def load_api(self, pkg_set):
4548
async with ClientSession(raise_for_status=True) as session:
4649
await asyncio.gather(
@@ -75,6 +78,9 @@ async def set_api(self, pkg, session):
7578

7679

7780
class PypiVersionAPI(VersionAPI):
81+
82+
package_type = "pypi"
83+
7884
async def load_api(self, pkg_set):
7985
async with ClientSession(raise_for_status=True) as session:
8086
await asyncio.gather(
@@ -96,6 +102,9 @@ async def fetch(self, pkg, session):
96102

97103

98104
class CratesVersionAPI(VersionAPI):
105+
106+
package_type = "cargo"
107+
99108
async def load_api(self, pkg_set):
100109
async with ClientSession(raise_for_status=True) as session:
101110
await asyncio.gather(
@@ -114,6 +123,9 @@ async def fetch(self, pkg, session):
114123

115124

116125
class RubyVersionAPI(VersionAPI):
126+
127+
package_type = "gem"
128+
117129
async def load_api(self, pkg_set):
118130
async with ClientSession(raise_for_status=True) as session:
119131
await asyncio.gather(
@@ -135,6 +147,9 @@ async def fetch(self, pkg, session):
135147

136148

137149
class NpmVersionAPI(VersionAPI):
150+
151+
package_type = "npm"
152+
138153
async def load_api(self, pkg_set):
139154
async with ClientSession(raise_for_status=True) as session:
140155
await asyncio.gather(
@@ -156,6 +171,9 @@ async def fetch(self, pkg, session):
156171

157172

158173
class DebianVersionAPI(VersionAPI):
174+
175+
package_type = "deb"
176+
159177
async def load_api(self, pkg_set):
160178
# Need to set the headers, because the Debian API upgrades
161179
# the connection to HTTP 2.0
@@ -189,6 +207,9 @@ async def set_api(self, pkg, session, retry_count=5):
189207

190208

191209
class MavenVersionAPI(VersionAPI):
210+
211+
package_type = "maven"
212+
192213
async def load_api(self, pkg_set):
193214
async with ClientSession(raise_for_status=True) as session:
194215
await asyncio.gather(
@@ -242,6 +263,9 @@ def extract_versions(xml_response: ET.ElementTree) -> Set[str]:
242263

243264

244265
class NugetVersionAPI(VersionAPI):
266+
267+
package_type = "nuget"
268+
245269
async def load_api(self, pkg_set):
246270
async with ClientSession(raise_for_status=True) as session:
247271
await asyncio.gather(
@@ -274,6 +298,9 @@ def extract_versions(resp: dict) -> Set[str]:
274298

275299

276300
class ComposerVersionAPI(VersionAPI):
301+
302+
package_type = "composer"
303+
277304
async def load_api(self, pkg_set):
278305
async with ClientSession(raise_for_status=True) as session:
279306
await asyncio.gather(
@@ -304,6 +331,9 @@ def extract_versions(resp: dict, pkg_name: str) -> Set[str]:
304331

305332

306333
class GitHubTagsAPI(VersionAPI):
334+
335+
package_type = "github"
336+
307337
async def load_api(self, repo_set):
308338
async with ClientSession(raise_for_status=True) as session:
309339
await asyncio.gather(

vulnerabilities/tests/test_github.py

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -303,6 +303,7 @@ def test_process_response(self):
303303
]
304304

305305
mock_version_api = MagicMock()
306+
mock_version_api.package_type = "maven"
306307
mock_version_api.get = lambda x: {'1.2.0', '9.0.2'}
307308
with patch('vulnerabilities.importers.github.MavenVersionAPI', return_value=mock_version_api): # nopep8
308309
with patch('vulnerabilities.importers.github.GitHubAPIDataSource.set_api'):

0 commit comments

Comments
 (0)