Skip to content

Commit 1857855

Browse files
committed
Refactor import_runner.py
Signed-off-by: Shivam Sandbhor <shivam.sandbhor@gmail.com>
1 parent 5991a67 commit 1857855

1 file changed

Lines changed: 180 additions & 61 deletions

File tree

vulnerabilities/import_runner.py

Lines changed: 180 additions & 61 deletions
Original file line numberDiff line numberDiff line change
@@ -24,20 +24,62 @@
2424
import dataclasses
2525
import datetime
2626
import logging
27+
from itertools import chain
2728
from typing import Dict
2829
from typing import List
2930
from typing import Set
3031
from typing import Tuple
32+
from typing import Optional
33+
from typing import Sequence
3134

3235
import packageurl
3336
from django.db import DataError
37+
from django.core import serializers
3438

3539
from vulnerabilities import models
3640
from vulnerabilities.data_source import Advisory, DataSource
3741
from vulnerabilities.data_source import PackageURL
3842

3943
logger = logging.getLogger(__name__)
4044

45+
# These _inserter classes are used to instantiate model objects.
46+
# Frozen dataclass store args required to store instantiate
47+
# model objects, this way model objects can be hashed indirectly which
48+
# is required in this implementation.
49+
50+
51+
@dataclasses.dataclass(frozen=True)
52+
class VulnerabilityReference_inserter:
53+
vulnerability: models.Vulnerability
54+
reference_id: Optional[str] = ''
55+
url: Optional[str] = ''
56+
57+
def __post_init__(self):
58+
if not any([self.reference_id, self.url]):
59+
raise TypeError(
60+
"VulnerabilityReference_inserter expects either reference_id or url")
61+
62+
def to_model_object(self):
63+
return models.VulnerabilityReference(**dataclasses.asdict(self))
64+
65+
# These _inserter classes are used to instantiate model objects.
66+
# Frozen dataclass store args required to store instantiate
67+
# model objects, this way model objects can be hashed indirectly which
68+
# is required in this implementation.
69+
70+
71+
@dataclasses.dataclass(frozen=True)
72+
class VulnerabilityImpact_inserter:
73+
vulnerability: models.Vulnerability
74+
is_vulnerable: bool
75+
package: models.Package
76+
77+
def to_model_object(self):
78+
if self.is_vulnerable:
79+
return models.ImpactedPackage(vulnerability=self.vulnerability, package=self.package)
80+
else:
81+
return models.ResolvedPackage(vulnerability=self.vulnerability, package=self.package)
82+
4183

4284
class ImportRunner:
4385
"""
@@ -55,6 +97,7 @@ class ImportRunner:
5597
(the data source should know this instead).
5698
- All update and select operations must use indexed columns.
5799
"""
100+
58101
def __init__(self, importer: models.Importer, batch_size: int):
59102
self.importer = importer
60103
self.batch_size = batch_size
@@ -71,19 +114,54 @@ def run(self, cutoff_date: datetime.datetime = None) -> None:
71114
from all Linux distributions that package this kernel version.
72115
"""
73116
logger.debug(f'Starting import for {self.importer.name}.')
74-
data_source = self.importer.make_data_source(self.batch_size, cutoff_date=cutoff_date)
75-
117+
data_source = self.importer.make_data_source(
118+
self.batch_size, cutoff_date=cutoff_date)
76119
with data_source:
77120
_process_added_advisories(data_source)
78121
_process_updated_advisories(data_source)
79-
80-
self.importer.last_run = datetime.datetime.now(tz=datetime.timezone.utc)
122+
self.importer.last_run = datetime.datetime.now(
123+
tz=datetime.timezone.utc)
81124
self.importer.data_source_cfg = dataclasses.asdict(data_source.config)
82125
self.importer.save()
83126

84127
logger.debug(f'Successfully finished import for {self.importer.name}.')
85128

86129

130+
def _process_updated_advisories(data_source: DataSource) -> None:
131+
132+
bulk_create_vuln_refs = set()
133+
bulk_create_ip = set()
134+
bulk_create_rp = set()
135+
136+
for batch in data_source.updated_advisories():
137+
for advisory in batch:
138+
139+
vuln, vuln_created, references = _create_vulnerability_and_references(
140+
advisory)
141+
bulk_create_vuln_refs.update(references)
142+
inew_refs = _create_pkg_vuln_refs(
143+
vuln, vuln_created, advisory.impacted_package_urls, is_vulnerable=True)
144+
rnew_refs = _create_pkg_vuln_refs(
145+
vuln, vuln_created, advisory.resolved_package_urls, is_vulnerable=False)
146+
147+
bulk_create_ip.update(inew_refs)
148+
bulk_create_rp.update(rnew_refs)
149+
150+
models.VulnerabilityReference.objects.bulk_create(
151+
[i.to_model_object() for i in bulk_create_vuln_refs]
152+
)
153+
154+
# FIXME : Detect intersection of bulk_create_ip and bulk_create_rp while
155+
# VulnerabilityImpact_inserter's is_vulnerable flag, those will be conflicts
156+
157+
models.ImpactedPackage.objects.bulk_create(
158+
[i.to_model_object() for i in bulk_create_ip]
159+
)
160+
models.ResolvedPackage.objects.bulk_create(
161+
[i.to_model_object() for i in bulk_create_rp]
162+
)
163+
164+
87165
def _process_added_advisories(data_source: DataSource) -> None:
88166
for batch in data_source.added_advisories():
89167
try:
@@ -92,54 +170,93 @@ def _process_added_advisories(data_source: DataSource) -> None:
92170

93171
vulnerabilities = _insert_vulnerabilities_and_references(batch)
94172

95-
_bulk_insert_impacted_and_resolved_packages(batch, vulnerabilities, impacted, resolved)
173+
_bulk_insert_impacted_and_resolved_packages(
174+
batch, vulnerabilities, impacted, resolved)
96175
except (DataError, RuntimeError) as e:
97176
# FIXME This exception might happen when the max. length of a DB column is exceeded.
98177
# Skipping an entire batch because one version number might be too long is obviously a
99178
# terrible way to handle this case.
100179
logger.exception(e)
101180

102181

103-
def _process_updated_advisories(data_source: DataSource) -> None:
104-
"""
105-
TODO: Make efficient; Current implementation needs way too many DB queries.
106-
"""
107-
for batch in data_source.updated_advisories():
108-
for advisory in batch:
109-
vuln, _ = _get_or_create_vulnerability(advisory)
182+
def _create_vulnerability_and_references(advisory: Advisory):
183+
vuln, vuln_created = _get_or_create_vulnerability(advisory)
184+
vuln_references = set()
185+
186+
if vuln_created:
187+
# This means vulnerability didn't previously exist in the db, so bulk create
188+
# is used without any hesitation
189+
for id_ in set(advisory.reference_ids):
190+
vuln_references.add(VulnerabilityReference_inserter(
191+
vulnerability=vuln, reference_id=id_))
110192

111-
for id_ in advisory.reference_ids:
112-
models.VulnerabilityReference.objects.get_or_create(
113-
vulnerability=vuln, reference_id=id_)
193+
for url in set(advisory.reference_urls):
194+
vuln_references.add(VulnerabilityReference_inserter(
195+
vulnerability=vuln, url=url))
114196

115-
for url in advisory.reference_urls:
116-
models.VulnerabilityReference.objects.get_or_create(vulnerability=vuln, url=url)
197+
else:
198+
vuln_refs_qs = models.VulnerabilityReference.objects.filter(
199+
vulnerability=vuln)
200+
vuln_ids = {ref.reference_id for ref in vuln_refs_qs}
201+
vuln_urls = {ref.url for ref in vuln_refs_qs}
117202

118-
for ipkg_url in advisory.impacted_package_urls:
119-
pkg, created = _get_or_create_package(ipkg_url)
203+
for id_ in advisory.reference_ids:
204+
# Add the item preventing duplicates pass to through.
205+
if id_ not in vuln_ids:
206+
vuln_ids.add(id_)
207+
vuln_references.add(VulnerabilityReference_inserter(
208+
vulnerability=vuln, reference_id=id_))
120209

121-
# FIXME Does not work yet due to cascading deletes.
122-
# if not created:
123-
# qs = models.ResolvedPackage.objects.filter(
124-
# vulnerability_id=vuln.id, package_id=pkg.id)
125-
# if qs:
126-
# qs[0].delete()
210+
for url in advisory.reference_urls:
211+
# Add the item preventing duplicates pass to through.
212+
if url not in vuln_urls:
213+
vuln_urls.add(url)
214+
vuln_references.add(VulnerabilityReference_inserter(
215+
vulnerability=vuln, url=url))
127216

128-
models.ImpactedPackage.objects.get_or_create(
129-
vulnerability_id=vuln.id, package_id=pkg.id)
217+
return vuln, vuln_created, vuln_references
130218

131-
for rpkg_url in advisory.resolved_package_urls:
132-
pkg, created = _get_or_create_package(rpkg_url)
133219

134-
# FIXME Does not work yet due to cascading deletes.
135-
# if not created:
136-
# qs = models.ImpactedPackage.objects.filter(
137-
# vulnerability_id=vuln.id, package_id=pkg.id)
138-
# if qs:
139-
# qs[0].delete()
220+
def _create_pkg_vuln_refs(vuln: models.Vulnerability, vuln_created: bool, purls: Sequence[PackageURL], is_vulnerable: bool): # nopep8
221+
new_refs, updated_refs = set(), set()
222+
for purl in purls:
223+
pkg, pkg_created = _get_or_create_package(purl)
224+
vuln_pkg_ref = VulnerabilityImpact_inserter(
225+
package=pkg, vulnerability=vuln, is_vulnerable=is_vulnerable)
140226

141-
models.ResolvedPackage.objects.get_or_create(
142-
vulnerability_id=vuln.id, package_id=pkg.id)
227+
if pkg_created or vuln_created:
228+
new_refs.add(vuln_pkg_ref)
229+
230+
else:
231+
232+
# Check for conflicts
233+
if is_vulnerable:
234+
op_qs = models.ResolvedPackage.objects.filter(
235+
vulnerability=vuln, package=pkg)
236+
qs = models.ImpactedPackage.objects.filter(
237+
vulnerability=vuln, package=pkg)
238+
239+
else:
240+
op_qs = models.ImpactedPackage.objects.filter(
241+
vulnerability=vuln, package=pkg)
242+
qs = models.ResolvedPackage.objects.filter(
243+
vulnerability=vuln, package=pkg)
244+
245+
if op_qs:
246+
conflicts = [i for i in qs]
247+
conflicts.append(vuln_pkg_ref.to_model_object())
248+
handle_conflicts(conflicts)
249+
qs.delete()
250+
251+
elif not qs:
252+
new_refs.add(vuln_pkg_ref)
253+
254+
return new_refs
255+
256+
257+
def handle_conflicts(conflicts):
258+
conflicts = serializers.serialize('json', [i for i in conflicts])
259+
models.ImportProblem.objects.create(conflicting_model=conflicts)
143260

144261

145262
def _get_or_create_vulnerability(advisory: Advisory) -> Tuple[models.Vulnerability, bool]:
@@ -169,35 +286,34 @@ def _get_or_create_package(p: PackageURL) -> Tuple[models.Package, bool]:
169286
}
170287

171288
if p.namespace:
172-
query_kwargs['namespace'] = packageurl.normalize_namespace(p.namespace, p.type, encode=True)
289+
query_kwargs['namespace'] = packageurl.normalize_namespace(
290+
p.namespace, p.type, encode=True)
173291

174292
if p.qualifiers:
175-
query_kwargs['qualifiers'] = packageurl.normalize_qualifiers(p.qualifiers, encode=False)
293+
query_kwargs['qualifiers'] = packageurl.normalize_qualifiers(
294+
p.qualifiers, encode=False)
176295

177296
if p.subpath:
178-
query_kwargs['subpath'] = packageurl.normalize_subpath(p.subpath, encode=True)
297+
query_kwargs['subpath'] = packageurl.normalize_subpath(
298+
p.subpath, encode=True)
179299

180300
return models.Package.objects.get_or_create(**query_kwargs)
181301

182302

183303
def _bulk_insert_packages(
184-
impacted: Set[PackageURL],
185-
resolved: Set[PackageURL]
304+
impacted: List[PackageURL],
305+
resolved: List[PackageURL]
186306
) -> Tuple[Dict[PackageURL, models.Package], Dict[PackageURL, models.Package]]:
187307

188-
packages = [_package_url_to_package(p) for p in impacted.union(resolved)]
189-
packages = models.Package.objects.bulk_create(packages)
190-
191-
impacted_packages, resolved_packages = {}, {}
308+
impacted_packages = models.Package.objects.bulk_create(
309+
[_package_url_to_package(p) for p in impacted])
310+
resolved_packages = models.Package.objects.bulk_create(
311+
[_package_url_to_package(p) for p in resolved])
192312

193-
for pkg in packages:
194-
# Unfortunately, PackageURLMixin.package_url returns a string, not a PackageURL
195-
purl = PackageURL.from_string(pkg.package_url)
196-
197-
if purl in impacted:
198-
impacted_packages[purl] = pkg
199-
elif purl in resolved:
200-
resolved_packages[purl] = pkg
313+
impacted_packages = dict(
314+
zip(impacted, [pkg for pkg in impacted_packages]))
315+
resolved_packages = dict(
316+
zip(resolved, [pkg for pkg in resolved_packages]))
201317

202318
return impacted_packages, resolved_packages
203319

@@ -257,15 +373,17 @@ def _insert_vulnerabilities_and_references(batch: Set[Advisory]) -> Set[models.V
257373
vuln: models.Vulnerability
258374

259375
if advisory.cve_id:
260-
vuln, created = models.Vulnerability.objects.get_or_create(cve_id=advisory.cve_id)
376+
vuln, created = models.Vulnerability.objects.get_or_create(
377+
cve_id=advisory.cve_id)
261378
if created and advisory.summary:
262379
vuln.summary = advisory.summary
263380
vuln.save()
264381
else:
265382
# FIXME
266383
# There is no way to check whether a vulnerability without a CVE ID already exists in
267384
# the database.
268-
vuln = models.Vulnerability.objects.create(summary=advisory.summary)
385+
vuln = models.Vulnerability.objects.create(
386+
summary=advisory.summary)
269387

270388
vulnerabilities.add(vuln)
271389

@@ -274,7 +392,8 @@ def _insert_vulnerabilities_and_references(batch: Set[Advisory]) -> Set[models.V
274392
vulnerability=vuln, reference_id=id_)
275393

276394
for url in advisory.reference_urls:
277-
models.VulnerabilityReference.objects.get_or_create(vulnerability=vuln, url=url)
395+
models.VulnerabilityReference.objects.get_or_create(
396+
vulnerability=vuln, url=url)
278397

279398
return vulnerabilities
280399

@@ -294,12 +413,12 @@ def _advisory_to_vulnerability(
294413
raise RuntimeError(f'No Vulnerability model object found for this Advisory: {advisory.summary}')
295414

296415

297-
def _collect_package_urls(batch: Set[Advisory]) -> Tuple[Set[PackageURL], Set[PackageURL]]:
298-
impacted, resolved = set(), set()
416+
def _collect_package_urls(batch: Set[Advisory]) -> Tuple[List[PackageURL], List[PackageURL]]:
417+
impacted, resolved = [], []
299418

300419
for advisory in batch:
301-
impacted.update(advisory.impacted_package_urls)
302-
resolved.update(advisory.resolved_package_urls)
420+
impacted.extend(advisory.impacted_package_urls)
421+
resolved.extend(advisory.resolved_package_urls)
303422

304423
return impacted, resolved
305424

0 commit comments

Comments
 (0)