Skip to content

Commit b10d0b4

Browse files
committed
npm importer - improver migration
Signed-off-by: ziad <ziadhany2016@gmail.com>
1 parent c94ed57 commit b10d0b4

5 files changed

Lines changed: 242 additions & 179 deletions

File tree

vulnerabilities/importer.py

Lines changed: 25 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -312,10 +312,28 @@ def advisory_data(self) -> Iterable[AdvisoryData]:
312312
raise NotImplementedError
313313

314314

315+
@dataclasses.dataclass
316+
class GitConfig:
317+
repository_url: str
318+
branch: Optional[str] = None
319+
create_working_directory: bool = True
320+
remove_working_directory: bool = True
321+
working_directory: Optional[str] = ""
322+
last_run_date: Optional[str] = None
323+
cutoff_date: Optional[str] = None
324+
325+
315326
# TODO: Needs rewrite
316327
class GitImporter(Importer):
317-
def validate_configuration(self) -> None:
328+
def __init__(self, config, cutoff_timestamp):
329+
super().__init__()
330+
self.config = config
331+
self.cutoff_timestamp = cutoff_timestamp
318332

333+
self._ensure_working_directory()
334+
self._ensure_repository()
335+
336+
def validate_configuration(self) -> None:
319337
if not self.config.create_working_directory and self.config.working_directory is None:
320338
self.error(
321339
'"create_working_directory" is not set but "working_directory" is set to '
@@ -336,10 +354,6 @@ def validate_configuration(self) -> None:
336354
"the default, which calls tempfile.mkdtemp()"
337355
)
338356

339-
def __enter__(self):
340-
self._ensure_working_directory()
341-
self._ensure_repository()
342-
343357
def __exit__(self, exc_type, exc_val, exc_tb):
344358
if self.config.remove_working_directory:
345359
shutil.rmtree(self.config.working_directory)
@@ -353,7 +367,6 @@ def file_changes(
353367
"""
354368
Returns all added and modified files since last_run_date or cutoff_date (whichever is more
355369
recent).
356-
357370
:param subdir: filter by files in this directory
358371
:param recursive: whether to include files in subdirectories
359372
:param file_ext: filter files by this extension
@@ -477,6 +490,12 @@ def _update_from_remote(self, remote, branch) -> None:
477490
branch.set_reference(remote.refs[branch.name])
478491
self._repo.head.reset(index=True, working_tree=True)
479492

493+
def advisory_data(self):
494+
raise NotImplementedError
495+
496+
def error(self, param):
497+
pass
498+
480499

481500
def _include_file(
482501
path: str,

vulnerabilities/importers/__init__.py

Lines changed: 2 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -12,6 +12,7 @@
1212
from vulnerabilities.importers import github
1313
from vulnerabilities.importers import gitlab
1414
from vulnerabilities.importers import nginx
15+
from vulnerabilities.importers import npm
1516
from vulnerabilities.importers import nvd
1617
from vulnerabilities.importers import openssl
1718
from vulnerabilities.importers import pysec
@@ -27,6 +28,7 @@
2728
pysec.PyPIImporter,
2829
debian.DebianImporter,
2930
gitlab.GitLabAPIImporter,
31+
npm.NpmImporter,
3032
]
3133

3234
IMPORTERS_REGISTRY = {x.qualified_name: x for x in IMPORTERS_REGISTRY}

vulnerabilities/importers/npm.py

Lines changed: 81 additions & 73 deletions
Original file line numberDiff line numberDiff line change
@@ -8,108 +8,96 @@
88
#
99

1010
# Author: Navonil Das (@NavonilDas)
11-
12-
import asyncio
11+
import logging
12+
from typing import Iterable
1313
from typing import List
14+
from typing import Optional
1415
from typing import Set
1516
from typing import Tuple
1617
from urllib.parse import quote
1718

1819
import pytz
1920
from dateutil.parser import parse
2021
from packageurl import PackageURL
21-
from univers.version_range import VersionRange
22+
from univers.version_range import NpmVersionRange
23+
from univers.versions import InvalidVersion
2224
from univers.versions import SemverVersion
2325

2426
from vulnerabilities.importer import AdvisoryData
27+
from vulnerabilities.importer import GitConfig
2528
from vulnerabilities.importer import GitImporter
2629
from vulnerabilities.importer import Reference
2730
from vulnerabilities.package_managers import NpmVersionAPI
31+
from vulnerabilities.package_managers import PackageVersion
2832
from vulnerabilities.utils import load_json
2933
from vulnerabilities.utils import nearest_patched_package
3034

3135
NPM_URL = "https://registry.npmjs.org{}"
36+
logger = logging.getLogger(__name__)
3237

3338

3439
class NpmImporter(GitImporter):
35-
def __enter__(self):
36-
super(NpmImporter, self).__enter__()
37-
if not getattr(self, "_added_files", None):
38-
self._added_files, self._updated_files = self.file_changes(
39-
recursive=True, file_ext="json", subdir="./vuln/npm"
40-
)
41-
42-
self._versions = NpmVersionAPI()
43-
self.set_api(self.collect_packages())
44-
45-
def updated_advisories(self) -> Set[AdvisoryData]:
46-
files = self._updated_files.union(self._added_files)
47-
advisories = []
48-
for f in files:
49-
processed_data = self.process_file(f)
50-
if processed_data:
51-
advisories.extend(processed_data)
52-
return self.batch_advisories(advisories)
53-
54-
def set_api(self, packages):
55-
asyncio.run(self._versions.load_api(packages))
56-
57-
def collect_packages(self):
58-
packages = set()
59-
files = self._updated_files.union(self._added_files)
60-
for f in files:
61-
data = load_json(f)
62-
packages.add(data["module_name"].strip())
63-
64-
return packages
65-
66-
@property
67-
def versions(self): # quick hack to make it patchable
68-
return self._versions
69-
70-
def process_file(self, file) -> List[AdvisoryData]:
40+
license_url = "https://github.com/nodejs/security-wg/blob/main/LICENSE.md"
41+
spdx_license_expression = "MIT"
42+
config = GitConfig(
43+
repository_url="https://github.com/nodejs/security-wg.git",
44+
working_directory="npm",
45+
branch="main",
46+
)
47+
cutoff_timestamp = 1
48+
49+
def __init__(self):
50+
super().__init__(config=self.config, cutoff_timestamp=self.cutoff_timestamp)
51+
self._added_files, self._updated_files = self.file_changes(
52+
recursive=True, file_ext="json", subdir="vuln/npm"
53+
)
54+
self.pkg_manager_api = NpmVersionAPI()
7155

72-
record = load_json(file)
73-
advisories = []
56+
def parse_advisory_data(self, record) -> Optional[AdvisoryData]:
7457
package_name = record["module_name"].strip()
7558

7659
publish_date = parse(record["updated_at"])
7760
publish_date = publish_date.replace(tzinfo=pytz.UTC)
7861

79-
all_versions = self.versions.get(package_name, until=publish_date).valid_versions
62+
all_versions = self.pkg_manager_api.fetch(package_name)
8063
aff_range = record.get("vulnerable_versions")
8164
if not aff_range:
8265
aff_range = ""
8366
fixed_range = record.get("patched_versions")
8467
if not fixed_range:
8568
fixed_range = ""
8669

87-
if aff_range == "*" or fixed_range == "*":
88-
return []
70+
# if aff_range == "*" or fixed_range == "*":
71+
# return None
8972

9073
impacted_versions, resolved_versions = categorize_versions(
9174
all_versions, aff_range, fixed_range
9275
)
9376

9477
impacted_purls = _versions_to_purls(package_name, impacted_versions)
9578
resolved_purls = _versions_to_purls(package_name, resolved_versions)
79+
9680
vuln_reference = [
9781
Reference(
9882
url=NPM_URL.format(f'/-/npm/v1/advisories/{record["id"]}'),
9983
reference_id=record["id"],
10084
)
10185
]
86+
cve_id = record.get("cves") or []
87+
88+
return AdvisoryData(
89+
aliases=cve_id,
90+
summary=record.get("overview", ""),
91+
affected_packages=nearest_patched_package(impacted_purls, resolved_purls),
92+
references=vuln_reference,
93+
date_published=publish_date,
94+
)
10295

103-
for cve_id in record.get("cves") or [""]:
104-
advisories.append(
105-
AdvisoryData(
106-
summary=record.get("overview", ""),
107-
vulnerability_id=cve_id,
108-
affected_packages=nearest_patched_package(impacted_purls, resolved_purls),
109-
references=vuln_reference,
110-
)
111-
)
112-
return advisories
96+
def advisory_data(self) -> Iterable[AdvisoryData]:
97+
files = self._updated_files.union(self._added_files)
98+
for file in files:
99+
record = load_json(file)
100+
yield self.parse_advisory_data(record)
113101

114102

115103
def _versions_to_purls(package_name, versions):
@@ -150,10 +138,10 @@ def normalize_ranges(version_range_string):
150138

151139

152140
def categorize_versions(
153-
all_versions: Set[str],
141+
all_versions: Iterable[PackageVersion],
154142
affected_version_range: str,
155143
fixed_version_range: str,
156-
) -> Tuple[Set[str], Set[str]]:
144+
) -> Tuple[Set[SemverVersion], Set[SemverVersion]]:
157145
"""
158146
Seperate list of affected versions and unaffected versions from all versions
159147
using the ranges specified.
@@ -168,31 +156,51 @@ def categorize_versions(
168156
fix_spec = []
169157

170158
if affected_version_range:
171-
aff_specs = normalize_ranges(affected_version_range)
172-
aff_spec = [
173-
VersionRange.from_scheme_version_spec_string("semver", spec)
174-
for spec in aff_specs
175-
if len(spec) >= 3
176-
]
159+
aff_spec = get_version_range(affected_version_range)
177160

178161
if fixed_version_range:
179-
fix_specs = normalize_ranges(fixed_version_range)
180-
fix_spec = [
181-
VersionRange.from_scheme_version_spec_string("semver", spec)
182-
for spec in fix_specs
183-
if len(spec) >= 3
184-
]
162+
fix_spec = get_version_range(fixed_version_range)
163+
185164
aff_ver, fix_ver = set(), set()
165+
186166
# Unaffected version is that version which is in the fixed_version_range
187167
# or which is absent in the affected_version_range
188-
for ver in all_versions:
189-
ver_obj = SemverVersion(ver)
168+
for ver in get_all_versions(all_versions):
190169

191-
if not any([ver_obj in spec for spec in aff_spec]) or any(
192-
[ver_obj in spec for spec in fix_spec]
193-
):
170+
if not any([ver in spec for spec in aff_spec]) or any([ver in spec for spec in fix_spec]):
194171
fix_ver.add(ver)
195172
else:
196173
aff_ver.add(ver)
197174

198175
return aff_ver, fix_ver
176+
177+
178+
def get_version_range(version_range) -> List[NpmVersionRange]:
179+
fix_specs = normalize_ranges(version_range)
180+
ver_range_objs = []
181+
for spec in fix_specs:
182+
if len(spec) >= 3:
183+
try:
184+
ver_range_objs.append(NpmVersionRange.from_string(f"vers:npm/{spec}"))
185+
except InvalidVersion:
186+
logger.error(f"InvalidVersionRange {spec}")
187+
188+
return ver_range_objs
189+
190+
191+
def get_all_versions(all_versions) -> List[SemverVersion]:
192+
"""
193+
194+
Args:
195+
all_versions:
196+
197+
Returns:
198+
199+
"""
200+
ver_objs = []
201+
for ver in all_versions:
202+
try:
203+
ver_objs.append(SemverVersion(ver.value))
204+
except InvalidVersion:
205+
logger.error(f"InvalidVersion {ver.value}")
206+
return ver_objs

0 commit comments

Comments
 (0)