@@ -119,71 +119,40 @@ def requests_with_5xx_retry(max_retries=5, backoff_factor=0.5):
119119
120120
121121def nearest_patched_package (vulnerable_packages , resolved_packages ):
122- if not vulnerable_packages :
123- return []
124-
125- def create_package_by_version_obj_mapping (packages , overwrite = False ):
126- # overwrite=True, returns version->PackageURL mapping.
127- # overwrite=False, returns version->[PackageURL] mapping.
128- if not packages :
129- return {}
130-
131- package_by_version_obj_mapping = {}
132- version_class = version_class_by_package_type [packages [0 ].type ]
133- for package in packages :
134- version_object = version_class (package .version )
135- if not overwrite :
136- if version_object in package_by_version_obj_mapping :
137- package_by_version_obj_mapping [version_object ].append (package )
138- else :
139- package_by_version_obj_mapping [version_object ] = [package ]
140- else :
141- package_by_version_obj_mapping [version_object ] = package
142-
143- return package_by_version_obj_mapping
144-
145- vulnerable_packages_by_version_obj = create_package_by_version_obj_mapping (vulnerable_packages )
146- resolved_package_by_version_obj = create_package_by_version_obj_mapping (
147- resolved_packages , overwrite = True
122+ # This class is used to get around bisect module's lack of supplying custom
123+ # compartor. Get rid of this once we use python 3.10 which supports this.
124+ # See https://github.com/python/cpython/pull/20556
125+ class PackageURLWithVersionComparator :
126+ def __init__ (self , package ):
127+ self .package = package
128+ self .version_object = version_class_by_package_type [package .type ](package .version )
129+
130+ def __eq__ (self , other ):
131+ return self .version_object == other .version_object
132+
133+ def __lt__ (self , other ):
134+ return self .version_object < other .version_object
135+
136+ vulnerable_packages = sorted (
137+ [PackageURLWithVersionComparator (package ) for package in vulnerable_packages ]
148138 )
149-
150- vulnerable_versions = list (vulnerable_packages_by_version_obj .keys ())
151- resolved_versions = list (resolved_package_by_version_obj .keys ())
152-
153- patched_version_by_vulnerable_versions = nearest_patched_versions (
154- vulnerable_versions , resolved_versions
139+ resolved_packages = sorted (
140+ [PackageURLWithVersionComparator (package ) for package in resolved_packages ]
155141 )
156142
157- affected_packages_with_patched_package = []
158- for vulnerable_version , patched_version in patched_version_by_vulnerable_versions .items ():
159- for vulnerable_package in vulnerable_packages_by_version_obj [vulnerable_version ]:
160- affected_packages_with_patched_package .append (
161- AffectedPackageWithPatchedPackage (
162- vulnerable_package = vulnerable_package ,
163- patched_package = resolved_package_by_version_obj .get (patched_version ),
164- )
165- )
166-
167- return affected_packages_with_patched_package
143+ resolved_package_count = len (resolved_packages )
144+ affected_package_with_patched_package_objects = []
168145
146+ for vulnerable_package in vulnerable_packages :
147+ patched_package_index = bisect .bisect_right (resolved_packages , vulnerable_package )
148+ patched_package = None
149+ if patched_package_index < resolved_package_count :
150+ patched_package = resolved_packages [patched_package_index ].package
169151
170- def nearest_patched_versions (vulnerable_versions , resolved_versions ):
171- """
172- Returns a mapping of vulnerable_version -> nearest_safe_version
173- """
152+ affected_package_with_patched_package_objects .append (
153+ AffectedPackageWithPatchedPackage (
154+ vulnerable_package = vulnerable_package .package , patched_package = patched_package
155+ )
156+ )
174157
175- vulnerable_versions = sorted (vulnerable_versions )
176- resolved_versions = sorted (resolved_versions )
177- resolved_version_count = len (resolved_versions )
178- nearest_patch_for_version = {}
179- for vulnerable_version in vulnerable_versions :
180- nearest_patch_for_version [vulnerable_version ] = None
181- if not resolved_versions :
182- continue
183-
184- patched_version_index = bisect .bisect_right (resolved_versions , vulnerable_version )
185- if patched_version_index >= resolved_version_count :
186- continue
187- nearest_patch_for_version [vulnerable_version ] = resolved_versions [patched_version_index ]
188-
189- return nearest_patch_for_version
158+ return affected_package_with_patched_package_objects
0 commit comments