diff --git a/marimo/_runtime/packages/utils.py b/marimo/_runtime/packages/utils.py index 310e65575d0..986214217e8 100644 --- a/marimo/_runtime/packages/utils.py +++ b/marimo/_runtime/packages/utils.py @@ -187,13 +187,73 @@ def append_version(pkg_name: str, version: str | None) -> str: return f"{pkg_name}=={version}" +def _is_pep508_requirement(package: str) -> bool: + from packaging.requirements import InvalidRequirement, Requirement + + try: + Requirement(package) + except InvalidRequirement: + return False + return True + + +def _split_around_extras(package: str) -> list[str]: + """Split on whitespace except within package extras.""" + parts: list[str] = [] + current: list[str] = [] + current_is_package_name = False + bracket_depth = 0 + quote: str | None = None + + for index, char in enumerate(package): + if char.isspace() and bracket_depth == 0 and quote is None: + if current and current[-1].isspace(): + current.append(char) + continue + if not current: + continue + next_index = index + 1 + while next_index < len(package) and package[next_index].isspace(): + next_index += 1 + if ( + next_index < len(package) + and package[next_index] == "[" + and current_is_package_name + ): + current.append(char) + continue + if current: + parts.append("".join(current)) + current = [] + current_is_package_name = False + continue + + if char in ("'", '"') and bracket_depth == 0: + quote = None if quote == char else quote or char + + if char == "[" and quote is None and current_is_package_name: + bracket_depth += 1 + elif char == "]" and bracket_depth > 0: + bracket_depth -= 1 + + if not current: + current_is_package_name = char.isascii() and char.isalnum() + elif not (char.isascii() and (char.isalnum() or char in "._-")): + current_is_package_name = False + current.append(char) + + if current: + parts.append("".join(current)) + return parts + + def split_packages(package: str) -> list[str]: - """ - Splits a package string into a list of packages. + """Split one or more package specifications. - This can handle editable packages (i.e. local directories) + PEP 508 requirements are parsed before falling back to handling editable + installs, paths, and URLs. - e.g. + Examples: "package1[extra1,extra2]==1.0.0" -> ["package1[extra1,extra2]==1.0.0"] "package1 package2" -> ["package1", "package2"] "package1==1.0.0 package2==2.0.0" -> ["package1==1.0.0", "package2==2.0.0"] @@ -201,22 +261,22 @@ def split_packages(package: str) -> list[str]: "package1 --editable /path/to/package1" -> ["package1 --editable /path/to/package1"] "package1 -e /path/to/package1 package2" -> ["package1 -e /path/to/package1", "package2"] "package1 @ /path/to/package1" -> ["package1 @ /path/to/package1"] - "foo==1.0; python_version>'3.6' bar==2.0; sys_platform=='win32'" -> ["foo==1.0; python_version>'3.6'", "bar==2.0; sys_platform=='win32'"] """ + package = package.strip() + if not package: + return [] + if _is_pep508_requirement(package): + return [package] + packages: list[str] = [] current_package: list[str] = [] in_environment_marker = False - for part in package.split(): + for part in _split_around_extras(package): if ( part in ["-e", "--editable", "@"] or current_package - and current_package[-1] - in [ - "-e", - "--editable", - "@", - ] + and current_package[-1] in ["-e", "--editable", "@"] ): current_package.append(part) elif part.endswith(";"): @@ -239,7 +299,7 @@ def split_packages(package: str) -> list[str]: if current_package: packages.append(" ".join(current_package)) - return [pkg.strip() for pkg in packages] + return packages @dataclasses.dataclass diff --git a/marimo/_server/api/endpoints/packages.py b/marimo/_server/api/endpoints/packages.py index fa0c04ac59d..b0caeac8482 100644 --- a/marimo/_server/api/endpoints/packages.py +++ b/marimo/_server/api/endpoints/packages.py @@ -63,7 +63,11 @@ async def add_package(request: Request) -> PackageOperationResponse: # Update the script metadata filename = _get_filename(request) - if filename is not None and GLOBAL_SETTINGS.MANAGE_SCRIPT_METADATA: + if ( + success + and filename is not None + and GLOBAL_SETTINGS.MANAGE_SCRIPT_METADATA + ): await asyncio.to_thread( package_manager.update_notebook_script_metadata, filepath=filename, @@ -112,7 +116,11 @@ async def remove_package(request: Request) -> PackageOperationResponse: # Update the script metadata filename = _get_filename(request) - if filename is not None and GLOBAL_SETTINGS.MANAGE_SCRIPT_METADATA: + if ( + success + and filename is not None + and GLOBAL_SETTINGS.MANAGE_SCRIPT_METADATA + ): await asyncio.to_thread( package_manager.update_notebook_script_metadata, filepath=filename, diff --git a/tests/_runtime/packages/test_package_utils.py b/tests/_runtime/packages/test_package_utils.py index 4352ae19d59..ff1ea63a28e 100644 --- a/tests/_runtime/packages/test_package_utils.py +++ b/tests/_runtime/packages/test_package_utils.py @@ -111,6 +111,13 @@ def test_split_packages() -> None: "foo @ https://example.com/foo.tar.gz", "bar @ https://example.com/bar.whl", ] + assert split_packages( + "git+https://example.com/foo.git https://example.com/bar.whl /tmp/baz.whl" + ) == [ + "git+https://example.com/foo.git", + "https://example.com/bar.whl", + "/tmp/baz.whl", + ] assert split_packages( "foo==1.0; python_version>'3.6' bar==2.0; sys_platform=='win32'" ) == [ @@ -125,6 +132,81 @@ def test_split_packages() -> None: ] +@pytest.mark.parametrize( + "requirement", + [ + "pydantic-ai[duckduckgo, web-fetch]", + "foo [extra1, extra2] == 1.2.3", + "foo >= 1.0, < 2.0", + "bar ( >= 3.0 )", + "foo ; python_version >= '3.10' and sys_platform == 'darwin'", + "foo; python_version < '3.10' or implementation_name == 'pypy'", + "foo; python_version>'3' and(sys_platform=='win32')", + "foo; python_version in'3.8 3.9'", + "foo; python_version>='3'and sys_platform=='darwin'", + ], +) +def test_split_packages_preserves_valid_requirement(requirement: str) -> None: + assert split_packages(requirement) == [requirement] + + +@pytest.mark.parametrize( + ("package_input", "expected"), + [ + ( + "pydantic-ai[duckduckgo, web-fetch] matplotlib", + ["pydantic-ai[duckduckgo, web-fetch]", "matplotlib"], + ), + ( + "matplotlib pydantic-ai[duckduckgo, web-fetch]", + ["matplotlib", "pydantic-ai[duckduckgo, web-fetch]"], + ), + ( + "foo [bar, baz] matplotlib", + ["foo [bar, baz]", "matplotlib"], + ), + ( + "matplotlib foo [bar, baz]", + ["matplotlib", "foo [bar, baz]"], + ), + ( + "numpy pydantic-ai[duckduckgo, web-fetch] matplotlib", + [ + "numpy", + "pydantic-ai[duckduckgo, web-fetch]", + "matplotlib", + ], + ), + ( + "https://example.com/foo.whl pydantic-ai[duckduckgo, web-fetch]", + [ + "https://example.com/foo.whl", + "pydantic-ai[duckduckgo, web-fetch]", + ], + ), + ( + "pydantic-ai[duckduckgo, web-fetch] -e /tmp/foo", + ["pydantic-ai[duckduckgo, web-fetch] -e /tmp/foo"], + ), + ], +) +def test_split_packages_with_spaced_extras( + package_input: str, expected: list[str] +) -> None: + assert split_packages(package_input) == expected + + +def test_split_packages_ignores_brackets_outside_extras() -> None: + assert split_packages("foo; os_name == '[abc def' bar") == [ + "foo; os_name == '[abc def'", + "bar", + ] + assert split_packages("foo @ https://example.com/a[bc matplotlib") == [ + "foo @ https://example.com/a[bc", + "matplotlib", + ] + + def test_strip_requirement_name() -> None: assert strip_requirement_name("package==1.0.0") == "package" assert strip_requirement_name("package[extra]>=1.0") == "package[extra]" diff --git a/tests/_runtime/packages/test_pypi_package_manager.py b/tests/_runtime/packages/test_pypi_package_manager.py index 16ed3ffe2bb..050d9b1411b 100644 --- a/tests/_runtime/packages/test_pypi_package_manager.py +++ b/tests/_runtime/packages/test_pypi_package_manager.py @@ -371,6 +371,24 @@ def test_uv_can_target_python_without_mutating_project() -> None: ] +@patch.object(UvPackageManager, "_uv_bin", "uv") +def test_uv_install_command_preserves_spaced_extras() -> None: + mgr = UvPackageManager.for_pip_install("/server/python") + + assert mgr.install_command( + "matplotlib pydantic-ai[duckduckgo, web-fetch]", upgrade=True + ) == [ + "uv", + "pip", + "install", + "--upgrade", + "matplotlib", + "pydantic-ai[duckduckgo, web-fetch]", + "-p", + "/server/python", + ] + + @patch.dict( "os.environ", {"VIRTUAL_ENV": "/path/to/venv", "UV": "/path/to/venv"}, diff --git a/tests/_server/api/endpoints/test_packages.py b/tests/_server/api/endpoints/test_packages.py index c940f2a09af..6f1aec110d7 100644 --- a/tests/_server/api/endpoints/test_packages.py +++ b/tests/_server/api/endpoints/test_packages.py @@ -507,6 +507,63 @@ def test_add_package_with_metadata_update( mock_package_manager_with_metadata.update_notebook_script_metadata.assert_called_once() +def test_add_package_with_spaced_extras_updates_metadata( + client: TestClient, mock_package_manager_with_metadata: Mock +) -> None: + package = "pydantic-ai[duckduckgo, web-fetch]" + + with ( + patch( + "marimo._config.settings.GLOBAL_SETTINGS.MANAGE_SCRIPT_METADATA", + True, + ), + patch( + "marimo._server.api.endpoints.packages._get_filename", + return_value="test.py", + ), + ): + response = client.post( + "/api/packages/add", + headers=HEADERS, + json={"package": package, "upgrade": True}, + ) + + assert response.json() == {"success": True, "error": None} + mock_package_manager_with_metadata.install.assert_awaited_once_with( + package, version=None, upgrade=True, group=None + ) + mock_package_manager_with_metadata.update_notebook_script_metadata.assert_called_once_with( + filepath="test.py", + packages_to_add=[package], + upgrade=True, + ) + + +def test_add_package_failure_does_not_update_metadata( + client: TestClient, mock_package_manager_with_metadata: Mock +) -> None: + mock_package_manager_with_metadata.install.return_value = False + + with ( + patch( + "marimo._config.settings.GLOBAL_SETTINGS.MANAGE_SCRIPT_METADATA", + True, + ), + patch( + "marimo._server.api.endpoints.packages._get_filename", + return_value="test.py", + ), + ): + response = client.post( + "/api/packages/add", + headers=HEADERS, + json={"package": "test-package"}, + ) + + assert response.json()["success"] is False + mock_package_manager_with_metadata.update_notebook_script_metadata.assert_not_called() + + def test_remove_package_with_metadata_update( client: TestClient, mock_package_manager_with_metadata: Mock ) -> None: @@ -529,6 +586,31 @@ def test_remove_package_with_metadata_update( mock_package_manager_with_metadata.update_notebook_script_metadata.assert_called_once() +def test_remove_package_failure_does_not_update_metadata( + client: TestClient, mock_package_manager_with_metadata: Mock +) -> None: + mock_package_manager_with_metadata.uninstall.return_value = False + + with ( + patch( + "marimo._config.settings.GLOBAL_SETTINGS.MANAGE_SCRIPT_METADATA", + True, + ), + patch( + "marimo._server.api.endpoints.packages._get_filename", + return_value="test.py", + ), + ): + response = client.post( + "/api/packages/remove", + headers=HEADERS, + json={"package": "test-package"}, + ) + + assert response.json()["success"] is False + mock_package_manager_with_metadata.update_notebook_script_metadata.assert_not_called() + + def test_add_package_no_metadata_update_when_disabled( client: TestClient, mock_package_manager_with_metadata: Mock ) -> None: