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
14 changes: 7 additions & 7 deletions lm15/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -293,12 +293,12 @@
# nothing on the request path needs them, and lm15's import time is a
# promise, so these names resolve lazily (PEP 562).
_LAZY_LOGIN = {
"Auth": ("lm15.login", "Auth"),
"AsyncAuth": ("lm15.login", "AsyncAuth"),
"BoundClient": ("lm15.login", "BoundClient"),
"TerminalUI": ("lm15.login", "TerminalUI"),
"providers": ("lm15.login", "providers"),
"connect": ("lm15.interactive", "connect"),
"Auth": (".login", "Auth"),
"AsyncAuth": (".login", "AsyncAuth"),
"BoundClient": (".login", "BoundClient"),
"TerminalUI": (".login", "TerminalUI"),
"providers": (".login", "providers"),
"connect": (".interactive", "connect"),
}


Expand All @@ -308,6 +308,6 @@ def __getattr__(name: str):
raise AttributeError(f"module 'lm15' has no attribute {name!r}")
import importlib

value = getattr(importlib.import_module(target[0]), target[1])
value = getattr(importlib.import_module(target[0], __package__), target[1])
globals()[name] = value
return value
2 changes: 1 addition & 1 deletion lm15/login/flows/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -52,7 +52,7 @@ def _account_flow(provider: str) -> ProviderFlow | None:
return None
import importlib

module = importlib.import_module(f"lm15.login.flows.{module_name}")
module = importlib.import_module(f".{module_name}", __package__)
for name in dir(module):
candidate = getattr(module, name)
if isinstance(candidate, type) and issubclass(candidate, ProviderFlow) and candidate is not ProviderFlow:
Expand Down
39 changes: 39 additions & 0 deletions tests/test_api_papercuts.py
Original file line number Diff line number Diff line change
Expand Up @@ -98,3 +98,42 @@ def test_factories_stay_in_types(self):

for name in self.FACTORY_NAMES:
assert callable(getattr(types, name))

@pytest.mark.parametrize("surface", ["exports", "flows"])
def test_lazy_imports_use_the_loaded_package(self, surface):
import subprocess
import sys
from pathlib import Path

code = """
import importlib.util
import sys
from pathlib import Path

root = Path(sys.argv[1])
sys.modules['lm15'] = None # No unrelated top-level installation may be used.
spec = importlib.util.spec_from_file_location(
'vendored_lm15', root / '__init__.py', submodule_search_locations=[str(root)],
)
package = importlib.util.module_from_spec(spec)
sys.modules[spec.name] = package
spec.loader.exec_module(package)
assert 'vendored_lm15.login' not in sys.modules
if sys.argv[2] == 'exports':
for name in package.__all__:
getattr(package, name)
from vendored_lm15.login import Auth
assert package.Auth is Auth
assert package.connect.__module__ == 'vendored_lm15.interactive'
else:
from vendored_lm15.login.flows import descriptor, flow, ProviderFlow
assert descriptor('xai').id == 'xai'
assert isinstance(flow('xai', 'device'), ProviderFlow)
assert type(flow('xai', 'device')).__module__ == 'vendored_lm15.login.flows.xai'
"""
result = subprocess.run(
[sys.executable, "-I", "-c", code,
str(Path(__file__).resolve().parents[1] / "lm15"), surface],
capture_output=True, text=True, encoding="utf-8", timeout=30,
)
assert result.returncode == 0, result.stderr
Loading