Skip to content
Open
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
24 changes: 20 additions & 4 deletions src/substrait/extension_registry/registry.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

import re
from collections import defaultdict
from importlib.resources import as_file
from importlib.resources import files as importlib_files
from pathlib import Path
from typing import Optional, Union
Expand Down Expand Up @@ -34,10 +35,25 @@ def __init__(self, load_default_extensions=True) -> None:
# extension relation's output schema can be derived during inference.
self._extension_relations: dict = {}
if load_default_extensions:
for fpath in importlib_files("substrait_extensions.extensions").glob( # type: ignore
"functions*.yaml"
):
self.register_extension_yaml(fpath)
# NB: iterate + filter instead of ``.glob("functions*.yaml")``.
# ``importlib.resources.files`` returns a ``Traversable``, which is
# not guaranteed to be a filesystem path: when
# ``substrait_extensions.extensions`` resolves as a namespace
# package the reader returns a ``MultiplexedPath``, and that type
# implements ``iterdir``/``open``/``joinpath`` but not ``glob``
# (calling ``.glob`` raises ``AttributeError``). ``iterdir`` is part
# of the ``Traversable`` protocol and works across ``Path``,
# ``MultiplexedPath``, and zip-based readers alike.
#
# ``as_file`` turns each entry into a real filesystem path for the
# duration of the ``with`` block (a no-op for on-disk resources, an
# extraction to a temp file for zip-backed ones). Without it,
# ``register_extension_yaml`` -> ``Path(fname)`` would raise
# ``TypeError`` on a ``zipfile.Path``.
Comment on lines +38 to +52

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Please cut this down to the lasting reason; the history fits better in the PR description.

Suggested change
# NB: iterate + filter instead of ``.glob("functions*.yaml")``.
# ``importlib.resources.files`` returns a ``Traversable``, which is
# not guaranteed to be a filesystem path: when
# ``substrait_extensions.extensions`` resolves as a namespace
# package the reader returns a ``MultiplexedPath``, and that type
# implements ``iterdir``/``open``/``joinpath`` but not ``glob``
# (calling ``.glob`` raises ``AttributeError``). ``iterdir`` is part
# of the ``Traversable`` protocol and works across ``Path``,
# ``MultiplexedPath``, and zip-based readers alike.
#
# ``as_file`` turns each entry into a real filesystem path for the
# duration of the ``with`` block (a no-op for on-disk resources, an
# extraction to a temp file for zip-backed ones). Without it,
# ``register_extension_yaml`` -> ``Path(fname)`` would raise
# ``TypeError`` on a ``zipfile.Path``.
# Traversable has no portable glob() (MultiplexedPath lacks it), so
# filter iterdir() by name instead.

for fpath in importlib_files("substrait_extensions.extensions").iterdir():

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Please collect the matches first and raise if there are none, so a broken install fails here rather than building an empty registry that only shows up later as "no matching overload" errors. A loader without resource support gives a Traversable whose iterdir() is empty: base raised AttributeError there, this PR silently loads nothing. With the round-trip through as_file dropped, something like:

            ext_dir = importlib_files("substrait_extensions.extensions")
            yamls = [
                f
                for f in ext_dir.iterdir()
                if f.name.startswith("functions") and f.name.endswith(".yaml")
            ]
            if not yamls:
                raise RuntimeError(
                    f"no functions*.yaml found in {ext_dir!r}; "
                    "is the substrait_extensions package data installed?"
                )
            for fpath in yamls:
                self.register_extension_yaml(fpath)

if fpath.name.startswith("functions") and fpath.name.endswith(".yaml"):
with as_file(fpath) as real_path:
self.register_extension_yaml(real_path)
Comment on lines +55 to +56

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Please drop as_file here and have register_extension_yaml read the Traversable directly, then pass fpath straight through (self.register_extension_yaml(fpath)) and remove the as_file import. as_file copies each YAML out to a temp file for zip installs, so those still fail when no writable temp dir is available, and opening in binary mode also fixes an existing crash on non-UTF-8 locales such as Japanese Windows (functions_arithmetic.yaml contains non-ASCII characters). This is the other option CodeRabbit offered; register_extension_yaml is outside the diff, so I can't attach a suggestion, but this is the change for registry.py:64-67:

            fname: Path to the YAML file, or any ``Traversable`` such as an entry
                from ``importlib.resources.files``
        """
        if not hasattr(fname, "open"):
            fname = Path(fname)
        # Binary mode lets PyYAML detect the encoding instead of using the locale's.
        with fname.open("rb") as f:


def register_extension_yaml(
self,
Expand Down
73 changes: 73 additions & 0 deletions tests/extension_registry/test_multiplexed_path.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,73 @@
"""Regression tests for loading default extensions off a non-``Path`` resource.

``importlib.resources.files`` returns a ``Traversable``, not necessarily a
filesystem ``Path``. The registry must load the bundled extension YAMLs
regardless of which concrete ``Traversable`` the resource reader hands back:

* ``MultiplexedPath`` -- returned for a *namespace* package (no ``__init__.py``).
It implements ``iterdir``/``open``/``joinpath`` but *not* ``glob``, so the old
``files(...).glob("functions*.yaml")`` raised
``AttributeError: 'MultiplexedPath' object has no attribute 'glob'`` (observed
under the pure-Python WASI guest interpreter).
* ``zipfile.Path`` -- returned for a zip-imported package. It is not an
``os.PathLike``, so ``register_extension_yaml``'s ``Path(fname)`` raised
``TypeError``; ``as_file`` now materialises a real path first.
Comment on lines +12 to +14

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Please say the old code "failed" here (and in the comment at line 71) rather than "raised TypeError": zipfile.Path has no glob before Python 3.12, so on 3.10/3.11 it raised AttributeError first. Since the file now covers zip too, a name like test_default_extensions.py would fit better.

"""

import zipfile
from importlib.resources import files

import substrait.extension_registry.registry as registry_module
from substrait.builders.type import i8
from substrait.extension_registry import ExtensionRegistry

try: # Python 3.11+
from importlib.resources.readers import MultiplexedPath
except ModuleNotFoundError: # Python 3.10
from importlib.readers import MultiplexedPath


def _assert_defaults_loaded(reg: ExtensionRegistry) -> None:
"""The default extension set was parsed and registered."""
assert reg._function_mapping
assert reg.lookup_function(
urn="extension:io.substrait:functions_arithmetic",
function_name="add",
signature=[i8(nullable=False), i8(nullable=False)],
)
Comment on lines +30 to +37

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Please make the tests check that the patched files() was actually used and that every bundled functions*.yaml loaded: as written, putting back importlib.resources.files(...).glob(...) still passes both, because on CPython that returns a regular Path. Replacing this helper with the one below (plus import yaml, dropping the i8 import and the two monkeypatch.setattr lines, and calling _load_defaults_from(monkeypatch, multiplexed) / _load_defaults_from(monkeypatch, zip_dir)) makes both tests fail under that regression:

def _load_defaults_from(monkeypatch, traversable) -> None:
    """Build a default registry with ``files()`` resolving to ``traversable``."""
    calls = []
    monkeypatch.setattr(
        registry_module,
        "importlib_files",
        lambda pkg: calls.append(pkg) or traversable,
    )
    reg = ExtensionRegistry(load_default_extensions=True)
    assert calls == ["substrait_extensions.extensions"]
    expected = {
        yaml.safe_load(f.read_bytes())["urn"]
        for f in files("substrait_extensions.extensions").iterdir()
        if f.name.startswith("functions") and f.name.endswith(".yaml")
    }
    assert set(reg.urns()) == expected



def test_load_default_extensions_via_multiplexed_path(monkeypatch):
# Wrap the real extensions directory in a MultiplexedPath so ``files()``
# yields exactly what it would for a namespace package. Guard the premise:
# MultiplexedPath must not expose ``glob`` (otherwise this test proves
# nothing).
ext_dir = files("substrait_extensions.extensions")
multiplexed = MultiplexedPath(ext_dir)
assert not hasattr(multiplexed, "glob")

monkeypatch.setattr(registry_module, "importlib_files", lambda _pkg: multiplexed)

# Previously raised AttributeError here; must now load cleanly.
reg = ExtensionRegistry(load_default_extensions=True)
_assert_defaults_loaded(reg)


def test_load_default_extensions_via_zipfile_path(monkeypatch, tmp_path):
# Mirror a zip-imported install: copy the real YAMLs into a zip and resolve
# ``files()`` to a zipfile.Path over the archive. The entries are not
# os.PathLike, so register_extension_yaml's ``Path(fname)`` would raise
# TypeError without the ``as_file`` materialisation.
ext_dir = files("substrait_extensions.extensions")
archive = tmp_path / "substrait_extensions.zip"
with zipfile.ZipFile(archive, "w") as zf:
for child in ext_dir.iterdir():
if child.name.endswith(".yaml"):
zf.writestr(f"extensions/{child.name}", child.read_bytes())

zip_dir = zipfile.Path(archive, "extensions/")
monkeypatch.setattr(registry_module, "importlib_files", lambda _pkg: zip_dir)

# Previously raised TypeError in register_extension_yaml; must now load.
reg = ExtensionRegistry(load_default_extensions=True)
_assert_defaults_loaded(reg)
Loading