Repository navigation
fix(extension_registry): load default extensions via iterdir() #285
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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 | ||
|
|
@@ -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``. | ||
| for fpath in importlib_files("substrait_extensions.extensions").iterdir(): | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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 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
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Please drop 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, | ||
|
|
||
| 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
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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 |
||
| """ | ||
|
|
||
| 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
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Please make the tests check that the patched 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) | ||
There was a problem hiding this comment.
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.