Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
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
86 changes: 73 additions & 13 deletions marimo/_runtime/packages/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -187,36 +187,96 @@ 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"]
"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

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(";"):
Expand All @@ -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
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
82 changes: 82 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,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]"
Expand Down
18 changes: 18 additions & 0 deletions tests/_runtime/packages/test_pypi_package_manager.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"},
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