|
1 | 1 | import sys |
2 | 2 | import os |
| 3 | +import py_compile |
3 | 4 | import marshal |
4 | 5 | import glob |
5 | 6 | import importlib |
@@ -300,6 +301,52 @@ def testBadMTime(self): |
300 | 301 | TESTMOD + pyc_ext: badtime_pyc} |
301 | 302 | self.doTest(".py", files, TESTMOD) |
302 | 303 |
|
| 304 | + def testCorruptDeflateDataRaisesZipImportError(self): |
| 305 | + # Regression test: a corrupt deflate stream inside a .pyc must |
| 306 | + # raise ZipImportError, not leak zlib.error. |
| 307 | + # See gh-158494. |
| 308 | + if self.compression != ZIP_DEFLATED: |
| 309 | + self.skipTest("requires ZIP_DEFLATED") |
| 310 | + |
| 311 | + modname = TESTMOD |
| 312 | + |
| 313 | + src = os.path.join(TEMP_DIR, modname + ".py") |
| 314 | + self.addCleanup(os_helper.unlink, TEMP_ZIP) |
| 315 | + self.addCleanup(os_helper.rmtree, TEMP_DIR) |
| 316 | + |
| 317 | + os.makedirs(TEMP_DIR, exist_ok=True) |
| 318 | + with open(src, "w") as f: |
| 319 | + f.write(test_src) |
| 320 | + |
| 321 | + # Build a real .pyc so its magic number matches this build. |
| 322 | + pyc_src = os.path.join(TEMP_DIR, modname + ".pyc") |
| 323 | + py_compile.compile(src, cfile=pyc_src, doraise=True) |
| 324 | + |
| 325 | + with ZipFile(TEMP_ZIP, "w", compression=ZIP_DEFLATED) as z: |
| 326 | + z.write(src, modname + ".py") |
| 327 | + z.write(pyc_src, modname + ".pyc") |
| 328 | + |
| 329 | + # Corrupt the compressed bytes of the .pyc entry in place. |
| 330 | + with open(TEMP_ZIP, "r+b") as f: |
| 331 | + data = f.read() |
| 332 | + idx = data.find((modname + ".pyc").encode()) |
| 333 | + self.assertNotEqual(idx, -1) |
| 334 | + local_header = data.rfind(b"PK\x03\x04", 0, idx) |
| 335 | + self.assertNotEqual(local_header, -1) |
| 336 | + name_len, extra_len = struct.unpack_from( |
| 337 | + "<HH", data, local_header + 26) |
| 338 | + data_start = local_header + 30 + name_len + extra_len |
| 339 | + f.seek(data_start) |
| 340 | + f.write(b"\xff" * 8) |
| 341 | + |
| 342 | + zipimport._zip_directory_cache.clear() |
| 343 | + sys.path.insert(0, TEMP_ZIP) |
| 344 | + self.addCleanup(sys.modules.pop, modname, None) |
| 345 | + importlib.invalidate_caches() |
| 346 | + |
| 347 | + with self.assertRaises(zipimport.ZipImportError): |
| 348 | + importlib.import_module(modname) |
| 349 | + |
303 | 350 | def test2038MTime(self): |
304 | 351 | # Make sure we can handle mtimes larger than what a 32-bit signed number |
305 | 352 | # can hold. |
|
0 commit comments