From 68c7d524927aa6da4bc68567c34e849041691771 Mon Sep 17 00:00:00 2001 From: Lalatendu Mohanty Date: Thu, 10 Sep 2026 16:19:22 -0400 Subject: [PATCH] fix(resolver): delegate hook source providers Route hook-backed source resolver profiles through `get_resolver_provider`, forwarding resolver context and server URL while preserving legacy hook signatures. Validate returned providers and cover cooldown and error behavior. Refs: #1324 Co-Authored-By: Codex Signed-off-by: Lalatendu Mohanty --- docs/concepts/resolver-architecture.rst | 8 + docs/reference/hooks.rst | 29 +++- src/fromager/packagesettings/_resolver.py | 108 +++++++++++- src/fromager/sources.py | 7 +- tests/test_packagesettings_resolver.py | 191 +++++++++++++++++++++- tests/test_sources.py | 80 ++++++++- 6 files changed, 406 insertions(+), 17 deletions(-) diff --git a/docs/concepts/resolver-architecture.rst b/docs/concepts/resolver-architecture.rst index 1865151ae..8359eb9bf 100644 --- a/docs/concepts/resolver-architecture.rst +++ b/docs/concepts/resolver-architecture.rst @@ -94,9 +94,17 @@ CLI commands interact with providers through a common Packages with ``source:`` configuration now select their provider through the configured source resolver. +.. versionchanged:: 0.96.0 + The ``hook-sdist`` and ``hook-prebuilt`` profiles create providers through + the required ``get_resolver_provider`` override hook. They do not fall back + to PyPI, and hook-backed artifact downloading remains a separate feature. + Per-package settings in YAML can select which provider to use and configure its parameters (index URL, tag pattern, etc.). When a package has a ``source:`` resolver configured, that resolver creates the provider. +The ``hook-sdist`` and ``hook-prebuilt`` profiles delegate provider creation +to the package's required ``get_resolver_provider`` hook and validate that it +returns a :class:`~fromager.resolver.BaseProvider`. Otherwise, override plugins can replace the provider for a specific package via the ``get_resolver_provider`` hook, after which the legacy resolver settings are used. diff --git a/docs/reference/hooks.rst b/docs/reference/hooks.rst index fbfe217ad..5e10f0a8b 100644 --- a/docs/reference/hooks.rst +++ b/docs/reference/hooks.rst @@ -138,10 +138,28 @@ Resolver hooks The arguments are the ``WorkContext``, the ``Requirement`` being evaluated, a boolean indicating whether source distributions should be included, a boolean indicating whether built wheels should be - included, and the URL for the sdist server. + included, and the URL for the sdist server. The hook also receives + ``req_type`` and ``ignore_platform`` when those parameters are supported + by its signature. - The return value must be an instance of a class that implements the - ``resolvelib.providers.AbstractProvider`` API. + The ``hook-sdist`` and ``hook-prebuilt`` source profiles require this hook. + They pass the following keyword arguments: + + * ``ctx`` + * ``req`` + * ``include_sdists`` + * ``include_wheels`` + * ``sdist_server_url`` + * ``req_type`` + * ``ignore_platform`` + + Older hooks may omit newer arguments; unsupported keyword arguments are + filtered for compatibility. A missing hook or a return value that is not a + :class:`~fromager.resolver.BaseProvider` is an error. Hook profiles never + fall back to the default PyPI provider. + + The return value must be an instance of + :class:`~fromager.resolver.BaseProvider`. The expectation is that it acts as an engine for any sort of package resolution whether it is for wheels or sources. The provider can @@ -169,7 +187,10 @@ Resolver hooks return VERSIONS.items() - def get_resolver_provider(ctx, req, include_sdists, include_wheels, sdist_server_url): + def get_resolver_provider( + ctx, req, include_sdists, include_wheels, sdist_server_url, + req_type=None, ignore_platform=False, + ): return resolver.GenericProvider(version_source=_version_source, constraints=ctx.constraints) ``GenericProvider``, ``GitHubTagProvider``, and ``GitLabTagProvider`` take diff --git a/src/fromager/packagesettings/_resolver.py b/src/fromager/packagesettings/_resolver.py index 0059b1677..27b980264 100644 --- a/src/fromager/packagesettings/_resolver.py +++ b/src/fromager/packagesettings/_resolver.py @@ -10,7 +10,7 @@ import pydantic -from .. import downloads, resolver +from .. import downloads, overrides, resolver from ..candidate import Cooldown from ._typedefs import MODEL_CONFIG @@ -68,6 +68,8 @@ def resolver_provider( ctx: context.WorkContext, req: Requirement, req_type: requirements_file.RequirementType | None, + *, + sdist_server_url: str = resolver.PYPI_SERVER_URL, ) -> resolver.BaseProvider: """Return a resolver provider for the given requirement.""" raise NotImplementedError @@ -199,6 +201,8 @@ def resolver_provider( ctx: context.WorkContext, req: Requirement, req_type: requirements_file.RequirementType | None, + *, + sdist_server_url: str = resolver.PYPI_SERVER_URL, ) -> resolver.PyPIProvider: return resolver.PyPIProvider( include_sdists=True, @@ -250,6 +254,8 @@ def resolver_provider( ctx: context.WorkContext, req: Requirement, req_type: requirements_file.RequirementType | None, + *, + sdist_server_url: str = resolver.PYPI_SERVER_URL, ) -> resolver.PyPIProvider: return resolver.PyPIProvider( include_sdists=False, @@ -323,6 +329,8 @@ def resolver_provider( ctx: context.WorkContext, req: Requirement, req_type: requirements_file.RequirementType | None, + *, + sdist_server_url: str = resolver.PYPI_SERVER_URL, ) -> resolver.PyPIProvider: return resolver.PyPIProvider( include_sdists=True, @@ -405,6 +413,8 @@ def resolver_provider( ctx: context.WorkContext, req: Requirement, req_type: requirements_file.RequirementType | None, + *, + sdist_server_url: str = resolver.PYPI_SERVER_URL, ) -> resolver.PyPIProvider: download_url = f"git+{self.clone_url}@refs/tags/{self.tag}" return resolver.PyPIProvider( @@ -572,6 +582,8 @@ def resolver_provider( ctx: context.WorkContext, req: Requirement, req_type: requirements_file.RequirementType | None, + *, + sdist_server_url: str = resolver.PYPI_SERVER_URL, ) -> resolver.GitHubTagProvider: return self._github_provider( ctx=ctx, @@ -613,6 +625,8 @@ def resolver_provider( ctx: context.WorkContext, req: Requirement, req_type: requirements_file.RequirementType | None, + *, + sdist_server_url: str = resolver.PYPI_SERVER_URL, ) -> resolver.GitHubTagProvider: return self._github_provider( ctx=ctx, @@ -652,6 +666,8 @@ def resolver_provider( ctx: context.WorkContext, req: Requirement, req_type: requirements_file.RequirementType | None, + *, + sdist_server_url: str = resolver.PYPI_SERVER_URL, ) -> resolver.GitLabTagProvider: return self._gitlab_provider( ctx=ctx, @@ -693,6 +709,8 @@ def resolver_provider( ctx: context.WorkContext, req: Requirement, req_type: requirements_file.RequirementType | None, + *, + sdist_server_url: str = resolver.PYPI_SERVER_URL, ) -> resolver.GitLabTagProvider: return self._gitlab_provider( ctx=ctx, @@ -720,6 +738,8 @@ def resolver_provider( ctx: context.WorkContext, req: Requirement, req_type: requirements_file.RequirementType | None, + *, + sdist_server_url: str = resolver.PYPI_SERVER_URL, ) -> resolver.BaseProvider: raise ValueError(f"package {req.name} is not available") @@ -739,14 +759,58 @@ class AbstractHookResolver(AbstractResolver, CooldownMixin): supports_override_hooks: typing.ClassVar[bool] = True """Hook resolvers support override hooks.""" + def _resolver_provider_from_hook( + self, + *, + ctx: context.WorkContext, + req: Requirement, + req_type: requirements_file.RequirementType | None, + sdist_server_url: str, + include_sdists: bool, + include_wheels: bool, + ignore_platform: bool, + ) -> resolver.BaseProvider: + """Invoke the required package resolver hook and validate its result.""" + hook = overrides.find_override_method(req.name, "get_resolver_provider") + if hook is None: + raise ValueError( + f"{req.name}: source resolver {self.provider!r} requires a " + "get_resolver_provider override hook" + ) + + try: + provider = overrides.invoke( + hook, + ctx=ctx, + req=req, + include_sdists=include_sdists, + include_wheels=include_wheels, + sdist_server_url=sdist_server_url, + req_type=req_type, + ignore_platform=ignore_platform, + ) + except Exception as err: + raise RuntimeError( + f"{req.name}: {self.provider!r} get_resolver_provider hook failed" + ) from err + + if not isinstance(provider, resolver.BaseProvider): + raise TypeError( + f"{req.name}: {self.provider!r} get_resolver_provider hook " + f"returned {type(provider).__name__}, expected " + "fromager.resolver.BaseProvider" + ) + return provider + def resolver_provider( self, ctx: context.WorkContext, req: Requirement, req_type: requirements_file.RequirementType | None, + *, + sdist_server_url: str = resolver.PYPI_SERVER_URL, ) -> resolver.BaseProvider: - # TODO - raise NotImplementedError("Hook resolver needs a hook") + raise NotImplementedError def download( self, @@ -776,6 +840,25 @@ class HookSDistResolver(AbstractHookResolver): {DownloadKind.sdist, DownloadKind.tarball, DownloadKind.git_checkout} ) + def resolver_provider( + self, + ctx: context.WorkContext, + req: Requirement, + req_type: requirements_file.RequirementType | None, + *, + sdist_server_url: str = resolver.PYPI_SERVER_URL, + ) -> resolver.BaseProvider: + """Return a source provider from the package resolver hook.""" + return self._resolver_provider_from_hook( + ctx=ctx, + req=req, + req_type=req_type, + sdist_server_url=sdist_server_url, + include_sdists=True, + include_wheels=False, + ignore_platform=False, + ) + class HookPrebuiltResolver(AbstractHookResolver): """Call resolver_provider and download_source hook, use pre-built wheel @@ -795,6 +878,25 @@ class HookPrebuiltResolver(AbstractHookResolver): ) resolves_prebuilt_wheel: typing.ClassVar[bool] = True + def resolver_provider( + self, + ctx: context.WorkContext, + req: Requirement, + req_type: requirements_file.RequirementType | None, + *, + sdist_server_url: str = resolver.PYPI_SERVER_URL, + ) -> resolver.BaseProvider: + """Return a pre-built wheel provider from the package resolver hook.""" + return self._resolver_provider_from_hook( + ctx=ctx, + req=req, + req_type=req_type, + sdist_server_url=sdist_server_url, + include_sdists=False, + include_wheels=True, + ignore_platform=False, + ) + SourceResolver = typing.Annotated[ PyPISDistResolver diff --git a/src/fromager/sources.py b/src/fromager/sources.py index 6cd8633a9..742741fd8 100644 --- a/src/fromager/sources.py +++ b/src/fromager/sources.py @@ -125,7 +125,12 @@ def get_source_provider( source_resolver = pbi.source_resolver if source_resolver is not None: - provider = source_resolver.resolver_provider(ctx, req, req_type) + provider = source_resolver.resolver_provider( + ctx, + req, + req_type, + sdist_server_url=sdist_server_url, + ) if req_type == RequirementType.TOP_LEVEL and resolver._has_equality_pin(req): provider.cooldown = None else: diff --git a/tests/test_packagesettings_resolver.py b/tests/test_packagesettings_resolver.py index 8132b4a49..2dffa5705 100644 --- a/tests/test_packagesettings_resolver.py +++ b/tests/test_packagesettings_resolver.py @@ -679,10 +679,55 @@ def test_parse(self) -> None: {DownloadKind.sdist, DownloadKind.tarball, DownloadKind.git_checkout} ) - def test_resolver_provider(self, tmp_context: WorkContext) -> None: + @mock.patch("fromager.packagesettings._resolver.overrides.find_override_method") + def test_resolver_provider( + self, + find_override_method: mock.Mock, + tmp_context: WorkContext, + ) -> None: r = _parse(self.YAML) - with pytest.raises(NotImplementedError): - r.resolver_provider(tmp_context, _REQ, _REQ_TYPE) + provider = resolver.GenericProvider(version_source=lambda identifier: []) + captured: dict[str, object] = {} + + def hook( + ctx: WorkContext, + req: Requirement, + include_sdists: bool, + include_wheels: bool, + sdist_server_url: str, + req_type: RequirementType | None = None, + ignore_platform: bool = False, + ) -> resolver.BaseProvider: + captured.update( + ctx=ctx, + req=req, + include_sdists=include_sdists, + include_wheels=include_wheels, + sdist_server_url=sdist_server_url, + req_type=req_type, + ignore_platform=ignore_platform, + ) + return provider + + find_override_method.return_value = hook + + result = r.resolver_provider( + tmp_context, + _REQ, + _REQ_TYPE, + sdist_server_url="https://index.test/simple", + ) + + assert result is provider + assert captured == { + "ctx": tmp_context, + "req": _REQ, + "include_sdists": True, + "include_wheels": False, + "sdist_server_url": "https://index.test/simple", + "req_type": _REQ_TYPE, + "ignore_platform": False, + } def test_download(self, tmp_context: WorkContext) -> None: r = _parse(self.YAML) @@ -704,10 +749,40 @@ def test_parse(self) -> None: assert r.resolves_prebuilt_wheel is True assert r.download_kinds == frozenset({DownloadKind.prebuilt_wheel}) - def test_resolver_provider(self, tmp_context: WorkContext) -> None: + @mock.patch("fromager.packagesettings._resolver.overrides.find_override_method") + def test_resolver_provider( + self, + find_override_method: mock.Mock, + tmp_context: WorkContext, + ) -> None: r = _parse(self.YAML) - with pytest.raises(NotImplementedError): - r.resolver_provider(tmp_context, _REQ, _REQ_TYPE) + provider = resolver.GenericProvider(version_source=lambda identifier: []) + + def hook( + ctx: WorkContext, + req: Requirement, + include_sdists: bool, + include_wheels: bool, + sdist_server_url: str, + req_type: RequirementType | None = None, + ignore_platform: bool = False, + ) -> resolver.BaseProvider: + return provider + + find_override_method.return_value = hook + + result = r.resolver_provider( + tmp_context, + _REQ, + _REQ_TYPE, + sdist_server_url="https://index.test/simple", + ) + + assert result is provider + find_override_method.assert_called_once_with( + _REQ.name, + "get_resolver_provider", + ) def test_download(self, tmp_context: WorkContext) -> None: r = _parse(self.YAML) @@ -715,6 +790,110 @@ def test_download(self, tmp_context: WorkContext) -> None: r.download(tmp_context, _REQ, _CANDIDATE_WHEEL) +@mock.patch("fromager.packagesettings._resolver.overrides.find_override_method") +def test_hook_resolver_supports_legacy_hook_signature( + find_override_method: mock.Mock, + tmp_context: WorkContext, +) -> None: + """Legacy five-argument resolver hooks remain compatible.""" + provider = resolver.GenericProvider(version_source=lambda identifier: []) + + def legacy_hook( + ctx: WorkContext, + req: Requirement, + include_sdists: bool, + include_wheels: bool, + sdist_server_url: str, + ) -> resolver.BaseProvider: + return provider + + find_override_method.return_value = legacy_hook + resolver_model = HookSDistResolver(provider="hook-sdist") + + result = resolver_model.resolver_provider(tmp_context, _REQ, _REQ_TYPE) + + assert result is provider + + +@mock.patch("fromager.packagesettings._resolver.resolver.default_resolver_provider") +@mock.patch( + "fromager.packagesettings._resolver.overrides.find_override_method", + return_value=None, +) +def test_hook_resolver_does_not_fall_back_when_hook_is_missing( + find_override_method: mock.Mock, + default_resolver_provider: mock.Mock, + tmp_context: WorkContext, +) -> None: + """A missing hook fails instead of selecting the default PyPI provider.""" + resolver_model = HookSDistResolver(provider="hook-sdist") + + with pytest.raises(ValueError, match="requires a get_resolver_provider"): + resolver_model.resolver_provider(tmp_context, _REQ, _REQ_TYPE) + + find_override_method.assert_called_once_with( + _REQ.name, + "get_resolver_provider", + ) + default_resolver_provider.assert_not_called() + + +@mock.patch("fromager.packagesettings._resolver.overrides.find_override_method") +def test_hook_resolver_rejects_invalid_provider( + find_override_method: mock.Mock, + tmp_context: WorkContext, +) -> None: + """Reject values that are not Fromager resolver providers.""" + + def hook( + ctx: WorkContext, + req: Requirement, + include_sdists: bool, + include_wheels: bool, + sdist_server_url: str, + req_type: RequirementType | None = None, + ignore_platform: bool = False, + ) -> object: + return object() + + find_override_method.return_value = hook + resolver_model = HookPrebuiltResolver(provider="hook-prebuilt") + + with pytest.raises( + TypeError, + match=r"expected fromager\.resolver\.BaseProvider", + ): + resolver_model.resolver_provider(tmp_context, _REQ, _REQ_TYPE) + + +@mock.patch("fromager.packagesettings._resolver.overrides.find_override_method") +def test_hook_resolver_chains_hook_exception( + find_override_method: mock.Mock, + tmp_context: WorkContext, +) -> None: + """Add resolver context while preserving the hook exception as the cause.""" + + def hook( + ctx: WorkContext, + req: Requirement, + include_sdists: bool, + include_wheels: bool, + sdist_server_url: str, + req_type: RequirementType | None = None, + ignore_platform: bool = False, + ) -> resolver.BaseProvider: + raise ValueError("hook failure") + + find_override_method.return_value = hook + resolver_model = HookSDistResolver(provider="hook-sdist") + + with pytest.raises(RuntimeError, match="hook-sdist") as exc_info: + resolver_model.resolver_provider(tmp_context, _REQ, _REQ_TYPE) + + assert isinstance(exc_info.value.__cause__, ValueError) + assert str(exc_info.value.__cause__) == "hook failure" + + # -- Discriminated union validation ------------------------------------------- diff --git a/tests/test_sources.py b/tests/test_sources.py index 74107abc3..e69135738 100644 --- a/tests/test_sources.py +++ b/tests/test_sources.py @@ -37,11 +37,16 @@ def test_get_source_provider_uses_configured_source_resolver( result = sources.get_source_provider( ctx=tmp_context, req=req, - sdist_server_url=resolver.PYPI_SERVER_URL, + sdist_server_url="https://caller.test/simple", ) assert result is provider - source_resolver.resolver_provider.assert_called_once_with(tmp_context, req, None) + source_resolver.resolver_provider.assert_called_once_with( + tmp_context, + req, + None, + sdist_server_url="https://caller.test/simple", + ) find_and_invoke.assert_not_called() @@ -64,7 +69,7 @@ def test_get_source_provider_uses_pypi_sdist_source_resolver( provider = sources.get_source_provider( ctx=tmp_context, req=req, - sdist_server_url=resolver.PYPI_SERVER_URL, + sdist_server_url="https://caller.test/simple", ) assert isinstance(provider, resolver.PyPIProvider) @@ -73,6 +78,75 @@ def test_get_source_provider_uses_pypi_sdist_source_resolver( assert provider.sdist_server_url == "https://pypi.test/simple" +@patch("fromager.packagesettings._resolver.overrides.find_override_method") +def test_get_source_provider_hook_forwards_url_and_applies_cooldown( + find_override_method: Mock, + tmp_context: context.WorkContext, +) -> None: + """Hook providers receive the runtime URL and centralized cooldown.""" + req = Requirement("test-pkg") + initial_cooldown = Cooldown(min_age=datetime.timedelta(days=1)) + provider = resolver.GenericProvider( + version_source=lambda identifier: [], + cooldown=initial_cooldown, + ) + captured: dict[str, object] = {} + + def hook( + ctx: context.WorkContext, + req: Requirement, + include_sdists: bool, + include_wheels: bool, + sdist_server_url: str, + req_type: RequirementType | None = None, + ignore_platform: bool = False, + ) -> resolver.BaseProvider: + captured.update( + ctx=ctx, + req=req, + include_sdists=include_sdists, + include_wheels=include_wheels, + sdist_server_url=sdist_server_url, + req_type=req_type, + ignore_platform=ignore_platform, + ) + return provider + + find_override_method.return_value = hook + tmp_context.cooldown = Cooldown(min_age=datetime.timedelta(days=7)) + source_resolver = packagesettings.HookSDistResolver( + provider="hook-sdist", + min_release_age=14, + ) + + with patch.object( + packagesettings.PackageBuildInfo, + "source_resolver", + new_callable=PropertyMock, + return_value=source_resolver, + ): + result = sources.get_source_provider( + ctx=tmp_context, + req=req, + sdist_server_url="https://caller.test/simple", + req_type=RequirementType.INSTALL, + ) + + assert result is provider + assert captured == { + "ctx": tmp_context, + "req": req, + "include_sdists": True, + "include_wheels": False, + "sdist_server_url": "https://caller.test/simple", + "req_type": RequirementType.INSTALL, + "ignore_platform": False, + } + assert result.cooldown is not initial_cooldown + assert result.cooldown is not None + assert result.cooldown.min_age == datetime.timedelta(days=14) + + def test_get_source_provider_forwards_req_type_to_source_resolver( tmp_context: context.WorkContext, ) -> None: