diff --git a/pipreqs/pipreqs.py b/pipreqs/pipreqs.py index 0877173..43b737b 100644 --- a/pipreqs/pipreqs.py +++ b/pipreqs/pipreqs.py @@ -227,6 +227,29 @@ def output_requirements(imports, symbol): generate_requirements_file("-", imports, symbol) +def _normalize_package_name(name): + return re.sub(r"[-_.]+", "-", name).lower() + + +def _deduplicate_imports(imports): + result = [] + seen = {} + + for item in imports: + name = item["name"] + key = _normalize_package_name(name) + if key in seen: + existing = result[seen[key]] + if not existing["version"] and item["version"]: + result[seen[key]] = item + continue + + seen[key] = len(result) + result.append(item) + + return result + + def get_imports_info(imports, pypi_server="https://pypi.python.org/pypi/", proxy=None): result = [] @@ -584,6 +607,8 @@ def init(args): clean(args["--clean"], imports) return + imports = _deduplicate_imports(imports) + if args["--mode"]: scheme = args.get("--mode") if scheme in ["compat", "gt", "no-pin"]: diff --git a/tests/test_pipreqs.py b/tests/test_pipreqs.py index 5e046e0..8bda318 100644 --- a/tests/test_pipreqs.py +++ b/tests/test_pipreqs.py @@ -142,6 +142,24 @@ def test_get_pkg_names(self): expected_output = ["camel", "Caroline", "Japan", "jury"] self.assertEqual(actual_output, expected_output) + def test_deduplicate_imports_uses_normalized_package_names(self): + imports = [ + {"name": "PyYAML", "version": "6.0.2"}, + {"name": "pyyaml", "version": "6.0.2"}, + {"name": "opencv_python", "version": "4.12.0"}, + {"name": "opencv-python", "version": "4.12.0"}, + {"name": "example_pkg", "version": None}, + {"name": "example-pkg", "version": "1.2.3"}, + ] + + actual_output = pipreqs._deduplicate_imports(imports) + expected_output = [ + {"name": "PyYAML", "version": "6.0.2"}, + {"name": "opencv_python", "version": "4.12.0"}, + {"name": "example-pkg", "version": "1.2.3"}, + ] + self.assertEqual(actual_output, expected_output) + def test_get_use_local_only(self): """ Test without checking PyPI, check to see if names of local @@ -209,6 +227,78 @@ def test_init_local_only(self): item = item.strip().split("==") self.assertTrue(item[0].lower() in self.local) + @patch("pipreqs.pipreqs.get_all_imports") + @patch("pipreqs.pipreqs.get_import_local") + def test_init_deduplicates_normalized_requirements(self, get_import_local_mock, get_all_imports_mock): + """ + Test that equivalent package names are only written once. + """ + get_all_imports_mock.return_value = ["yaml", "win32gui"] + get_import_local_mock.return_value = [ + {"name": "PyYAML", "version": "6.0.2", "exports": ["yaml"]}, + {"name": "pyyaml", "version": "6.0.2", "exports": ["yaml"]}, + {"name": "pywin32", "version": "310", "exports": ["win32gui"]}, + {"name": "pywin32", "version": "311", "exports": ["win32gui"]}, + ] + + pipreqs.init( + { + "": self.project, + "--savepath": None, + "--print": False, + "--use-local": True, + "--force": True, + "--proxy": None, + "--pypi-server": None, + "--diff": None, + "--clean": None, + "--mode": None, + } + ) + + assert os.path.exists(self.requirements_path) == 1 + with open(self.requirements_path, "r") as f: + data = f.read().splitlines() + + self.assertEqual(["pywin32==310", "PyYAML==6.0.2"], data) + + @patch("pipreqs.pipreqs.get_all_imports") + @patch("pipreqs.pipreqs.get_import_local") + def test_init_clean_keeps_normalized_alias_requirements( + self, get_import_local_mock, get_all_imports_mock + ): + """ + Test that --clean keeps raw requirement names before output deduplication. + """ + get_all_imports_mock.return_value = ["cv2"] + get_import_local_mock.return_value = [ + {"name": "opencv-python", "version": "4.12.0", "exports": ["cv2"]}, + {"name": "opencv_python", "version": "4.12.0", "exports": ["cv2"]}, + ] + with open(self.requirements_path, "w") as f: + f.write("opencv_python==4.12.0\nunused-package==1.0.0\n") + + pipreqs.init( + { + "": self.project, + "--savepath": None, + "--print": False, + "--use-local": True, + "--force": True, + "--proxy": None, + "--pypi-server": None, + "--diff": None, + "--clean": self.requirements_path, + "--mode": None, + } + ) + + with open(self.requirements_path, "r") as f: + data = f.read() + + self.assertIn("opencv_python==4.12.0", data) + self.assertNotIn("unused-package", data) + def test_init_savepath(self): """ Test that we can save requirements.txt correctly