Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
36 changes: 24 additions & 12 deletions marimo/_runtime/packages/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -187,22 +187,39 @@ 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_packages(package: str) -> list[str]:
"""
Splits a package string into a list of packages.
"""Split a requirement or compact CLI-style package list.

This can handle editable packages (i.e. local directories)
A valid PEP 508 requirement is always returned unchanged. The fallback
supports legacy whitespace-separated input such as `pandas numpy` and
editable installs. Requirements containing whitespace must be passed one
at a time because whitespace also separates packages in the legacy form.

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"]
"package1 -e /path/to/package1" -> ["package1 -e /path/to/package1"]
"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]
Comment thread
Light2Dark marked this conversation as resolved.

packages: list[str] = []
current_package: list[str] = []
in_environment_marker = False
Expand All @@ -211,12 +228,7 @@ def split_packages(package: str) -> list[str]:
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(";"):
Expand All @@ -239,7 +251,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
Expand Down
12 changes: 10 additions & 2 deletions marimo/_server/api/endpoints/packages.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand Down Expand Up @@ -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,
Expand Down
25 changes: 25 additions & 0 deletions tests/_runtime/packages/test_package_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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'"
) == [
Expand All @@ -125,6 +132,24 @@ 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]


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]"
Expand Down
82 changes: 82 additions & 0 deletions tests/_server/api/endpoints/test_packages.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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:
Expand Down
Loading