Skip to content

Commit 3d2d6a3

Browse files
committed
Make more style corrections in GitHubAPI importer
Signed-off-by: Shivam Sandbhor <shivam.sandbhor@gmail.com>
1 parent ead9e97 commit 3d2d6a3

2 files changed

Lines changed: 59 additions & 84 deletions

File tree

vulnerabilities/importers/github.py

Lines changed: 58 additions & 83 deletions
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
# Copyright (c) 2017 nexB Inc. and others. All rights reserved.
1+
# Copyright (c) nexB Inc. and others. All rights reserved.
22
# http://nexb.com and https://github.com/nexB/vulnerablecode/
33
# The VulnerableCode software is licensed under the Apache License version 2.0.
44
# Data generated with VulnerableCode require an acknowledgment.
@@ -44,7 +44,7 @@
4444
# second '%s' is interesting, it will have the value '' for the first request,
4545
# since we don't have any value for endCursor at the beginning
4646
# for all the subsequent requests it will have value 'after: "{endCursor}""
47-
query = '''
47+
query = """
4848
query MyQuery {
4949
securityVulnerabilities(first: 100, ecosystem: %s, %s) {
5050
edges {
@@ -68,7 +68,7 @@
6868
}
6969
}
7070
}
71-
'''
71+
"""
7272

7373

7474
class GitHubTokenError(Exception):
@@ -88,9 +88,9 @@ class GitHubAPIDataSource(DataSource):
8888
def __init__(self, *args, **kwargs):
8989
super().__init__(*args, **kwargs)
9090
try:
91-
self.gh_token = os.environ['GH_TOKEN']
91+
self.gh_token = os.environ["GH_TOKEN"]
9292
except KeyError:
93-
raise GitHubTokenError('Envirnomental variable GH_TOKEN is missing')
93+
raise GitHubTokenError("Envirnomental variable GH_TOKEN is missing")
9494

9595
def __enter__(self):
9696
self.advisories = self.fetch()
@@ -99,66 +99,53 @@ def updated_advisories(self) -> Set[Advisory]:
9999
return self.batch_advisories(self.process_response())
100100

101101
def fetch(self) -> Mapping[str, List[Mapping]]:
102-
103-
headers = {'Authorization': 'token ' + self.gh_token}
102+
headers = {"Authorization": "token " + self.gh_token}
104103
api_data = {}
105104
for ecosystem in self.config.ecosystems:
106105

107106
api_data[ecosystem] = []
108-
end_cursor_exp = ''
107+
end_cursor_exp = ""
109108

110109
while True:
111110

112-
query_json = {'query': query % (ecosystem, end_cursor_exp)}
113-
resp = requests.post(
114-
self.config.endpoint, headers=headers, json=query_json
115-
).json()
111+
query_json = {"query": query % (ecosystem, end_cursor_exp)}
112+
resp = requests.post(self.config.endpoint, headers=headers, json=query_json).json()
116113

117-
if resp.get('message') == 'Bad credentials':
118-
raise GitHubTokenError('Invalid GitHub token')
114+
if resp.get("message") == "Bad credentials":
115+
raise GitHubTokenError("Invalid GitHub token")
119116

120-
end_cursor = resp['data']['securityVulnerabilities']['pageInfo'][
121-
'endCursor'
122-
]
123-
end_cursor_exp = 'after: {}'.format('"{}"'.format(end_cursor))
117+
end_cursor = resp["data"]["securityVulnerabilities"]["pageInfo"]["endCursor"]
118+
end_cursor_exp = "after: {}".format('"{}"'.format(end_cursor))
124119
api_data[ecosystem].append(resp)
125-
print(resp)
126120

127-
if not resp['data']['securityVulnerabilities']['pageInfo'][
128-
'hasNextPage'
129-
]:
121+
if not resp["data"]["securityVulnerabilities"]["pageInfo"]["hasNextPage"]:
130122
break
131123
return api_data
132124

133125
def set_version_api(self, ecosystem: str) -> None:
134-
135-
if ecosystem == 'MAVEN':
136-
self.version_api = MavenVersionAPI()
137-
138-
elif ecosystem == 'NUGET':
139-
self.version_api = NugetVersionAPI()
140-
141-
elif ecosystem == 'COMPOSER':
142-
self.version_api = ComposerVersionAPI()
126+
versioners = {
127+
"MAVEN": MavenVersionAPI,
128+
"NUGET": NugetVersionAPI,
129+
"COMPOSER": ComposerVersionAPI,
130+
}
131+
versioner = versioners.get(ecosystem)
132+
if versioner:
133+
self.version_api = versioner()
143134

144135
@staticmethod
145-
def process_name(
146-
ecosystem: str, pkg_name: str
147-
) -> Optional[Tuple[Optional[str], str]]:
148-
149-
if ecosystem == 'MAVEN':
150-
151-
artifact_comps = pkg_name.split(':')
136+
def process_name(ecosystem: str, pkg_name: str) -> Optional[Tuple[Optional[str], str]]:
137+
if ecosystem == "MAVEN":
138+
artifact_comps = pkg_name.split(":")
152139
if len(artifact_comps) != 2:
153140
return
154141
ns, name = artifact_comps
155142
return ns, name
156143

157-
if ecosystem == 'NUGET':
144+
if ecosystem == "NUGET":
158145
return None, pkg_name
159146

160-
if ecosystem == 'COMPOSER':
161-
vendor, name = pkg_name.split('/')
147+
if ecosystem == "COMPOSER":
148+
vendor, name = pkg_name.split("/")
162149
return vendor, name
163150

164151
def process_response(self) -> List[Advisory]:
@@ -167,42 +154,38 @@ def process_response(self) -> List[Advisory]:
167154
self.set_version_api(ecosystem)
168155
pkg_type = ecosystem.lower()
169156
for resp_page in self.advisories[ecosystem]:
170-
for adv in resp_page['data']['securityVulnerabilities']['edges']:
171-
name = adv['node']['package']['name']
157+
for adv in resp_page["data"]["securityVulnerabilities"]["edges"]:
158+
name = adv["node"]["package"]["name"]
172159

173160
if self.process_name(ecosystem, name):
174161
ns, pkg_name = self.process_name(ecosystem, name)
175162
else:
176163
continue
177-
aff_range = adv['node']['vulnerableVersionRange']
164+
aff_range = adv["node"]["vulnerableVersionRange"]
178165
self.version_api.load_to_api(name)
179166
aff_vers, unaff_vers = self.categorize_versions(
180167
aff_range, self.version_api.get(name)
181168
)
182169

183170
affected_purls = {
184-
PackageURL(
185-
name=pkg_name, namespace=ns, version=version, type=pkg_type
186-
)
171+
PackageURL(name=pkg_name, namespace=ns, version=version, type=pkg_type)
187172
for version in aff_vers
188173
}
189174

190175
unaffected_purls = {
191-
PackageURL(
192-
name=pkg_name, namespace=ns, version=version, type=pkg_type
193-
)
176+
PackageURL(name=pkg_name, namespace=ns, version=version, type=pkg_type)
194177
for version in unaff_vers
195178
}
196179

197180
cve_ids = set()
198181
ref_ids = set()
199-
vuln_desc = adv['node']['advisory']['summary']
182+
vuln_desc = adv["node"]["advisory"]["summary"]
200183

201-
for vuln in adv['node']['advisory']['identifiers']:
202-
if vuln['type'] == 'CVE':
203-
cve_ids.add(vuln['value'])
184+
for vuln in adv["node"]["advisory"]["identifiers"]:
185+
if vuln["type"] == "CVE":
186+
cve_ids.add(vuln["value"])
204187
else:
205-
ref_ids.add(vuln['value'])
188+
ref_ids.add(vuln["value"])
206189
for cve_id in cve_ids:
207190
adv_list.append(
208191
Advisory(
@@ -216,13 +199,9 @@ def process_response(self) -> List[Advisory]:
216199
return adv_list
217200

218201
@staticmethod
219-
def categorize_versions(
220-
version_range: str, all_versions: Set[str]
221-
) -> Tuple[Set[str], Set[str]]:
202+
def categorize_versions(version_range: str, all_versions: Set[str]) -> Tuple[Set[str], Set[str]]: # nopep8
222203
version_range = RangeSpecifier(version_range)
223-
affected_versions = {
224-
version for version in all_versions if version in version_range
225-
}
204+
affected_versions = {version for version in all_versions if version in version_range}
226205
return (affected_versions, all_versions - affected_versions)
227206

228207

@@ -234,39 +213,34 @@ def get(self, pkg_name: str) -> Set[str]:
234213
return self.cache.get(pkg_name, set())
235214

236215
def load_to_api(self, pkg_name: str) -> None:
237-
238216
if pkg_name in self.cache:
239217
return
240218

241-
artifact_comps = pkg_name.split(':')
219+
artifact_comps = pkg_name.split(":")
242220
endpoint = self.artifact_url(artifact_comps)
243221
resp = requests.get(endpoint).content
244222

245223
try:
246-
247-
xml_resp = ET.ElementTree(ET.fromstring(resp.decode('utf-8')))
224+
xml_resp = ET.ElementTree(ET.fromstring(resp.decode("utf-8")))
248225
self.cache[pkg_name] = self.extract_versions(xml_resp)
249-
250226
except ET.ParseError:
251227
self.cache[pkg_name] = set()
252228

253229
@staticmethod
254230
def artifact_url(artifact_comps: List[str]) -> str:
255-
256-
base_url = 'https://repo.maven.apache.org/maven2/{}'
231+
base_url = "https://repo.maven.apache.org/maven2/{}"
257232
group_id, artifact_id = artifact_comps
258-
group_url = group_id.replace('.', '/')
259-
suffix = group_url + '/' + artifact_id + '/' + 'maven-metadata.xml'
233+
group_url = group_id.replace(".", "/")
234+
suffix = group_url + "/" + artifact_id + "/" + "maven-metadata.xml"
260235
endpoint = base_url.format(suffix)
261236

262237
return endpoint
263238

264239
@staticmethod
265240
def extract_versions(xml_response: ET.ElementTree) -> Set[str]:
266-
267241
all_versions = set()
268242
for child in xml_response.getroot().iter():
269-
if child.tag == 'version':
243+
if child.tag == "version":
270244
all_versions.add(child.text)
271245

272246
return all_versions
@@ -295,18 +269,18 @@ def load_to_api(self, pkg_name: str) -> None:
295269

296270
@staticmethod
297271
def nuget_url(pkg_name: str) -> str:
298-
base_url = 'https://api.nuget.org/v3/registration5-semver1/{}/index.json'
272+
base_url = "https://api.nuget.org/v3/registration5-semver1/{}/index.json"
299273
return base_url.format(pkg_name.lower())
300274

301275
@staticmethod
302-
def extract_versions(json_resp: dict) -> Set[str]:
276+
def extract_versions(resp: dict) -> Set[str]:
303277
all_versions = set()
304278
try:
305-
for entry in json_resp['items'][0]['items']:
306-
all_versions.add(entry['catalogEntry']['version'])
279+
for entry in resp["items"][0]["items"]:
280+
all_versions.add(entry["catalogEntry"]["version"])
307281
# json response for YamlDotNet.Signed triggers this exception
308282
except KeyError:
309-
return all_versions
283+
pass
310284

311285
return all_versions
312286

@@ -321,22 +295,23 @@ def get(self, pkg_name: str) -> Set[str]:
321295
def load_to_api(self, pkg_name: str) -> None:
322296
if pkg_name in self.cache:
323297
return
298+
324299
endpoint = self.composer_url(pkg_name)
325300
json_resp = requests.get(endpoint).json()
326301
self.cache[pkg_name] = self.extract_versions(json_resp, pkg_name)
327302

328303
@staticmethod
329304
def composer_url(pkg_name: str) -> str:
330-
vendor, name = pkg_name.split('/')
331-
return f'https://repo.packagist.org/p/{vendor}/{name}.json'
305+
vendor, name = pkg_name.split("/")
306+
return f"https://repo.packagist.org/p/{vendor}/{name}.json"
332307

333308
@staticmethod
334-
def extract_versions(json_resp: dict, pkg_name: str) -> Set[str]:
335-
all_versions = json_resp['packages'][pkg_name].keys()
309+
def extract_versions(resp: dict, pkg_name: str) -> Set[str]:
310+
all_versions = resp["packages"][pkg_name].keys()
336311
# This filter ensures, that all_versions contains only released versions
337-
all_versions = set(filter(lambda x: 'dev' not in x, all_versions))
312+
all_versions = set(filter(lambda x: "dev" not in x, all_versions))
338313
# See https://github.com/composer/composer/blob/44a4429978d1b3c6223277b875762b2930e83e8c/doc/articles/versions.md#tags # nopep8
339314
# for explanation of removing 'v'
340-
all_versions = set(map(lambda x: x.replace('v', ''), all_versions))
315+
all_versions = set(map(lambda x: x.replace("v", ""), all_versions))
341316

342317
return all_versions

vulnerabilities/tests/test_github.py

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,4 @@
1-
# Copyright (c) 2017 nexB Inc. and others. All rights reserved.
1+
# Copyright (c) nexB Inc. and others. All rights reserved.
22
# http://nexb.com and https://github.com/nexB/vulnerablecode/
33
# The VulnerableCode software is licensed under the Apache License version 2.0.
44
# Data generated with VulnerableCode require an acknowledgment.

0 commit comments

Comments
 (0)