diff --git a/homeassistant/components/portainer/button.py b/homeassistant/components/portainer/button.py index e5ee60f25495..e8d6784b9d77 100644 --- a/homeassistant/components/portainer/button.py +++ b/homeassistant/components/portainer/button.py @@ -31,7 +31,6 @@ from .entity import ( PortainerEndpointEntity, PortainerStackEntity, ) -from .util import async_call_portainer PARALLEL_UPDATES = 1 @@ -269,7 +268,7 @@ class PortainerBaseButton(ButtonEntity): @override async def async_press(self) -> None: """Trigger the Portainer button press service.""" - await async_call_portainer(self._async_press_call()) + await self.coordinator.async_call_portainer(self._async_press_call()) await self.coordinator.async_request_refresh() diff --git a/homeassistant/components/portainer/coordinator.py b/homeassistant/components/portainer/coordinator.py index 17d839bc3441..c97cdd3bea9f 100644 --- a/homeassistant/components/portainer/coordinator.py +++ b/homeassistant/components/portainer/coordinator.py @@ -2,13 +2,13 @@ from abc import abstractmethod import asyncio -from collections.abc import Callable +from collections.abc import Awaitable, Callable import dataclasses from dataclasses import dataclass from datetime import datetime, timedelta import logging import time -from typing import override +from typing import Any, override from pyportainer import ( DockerContainerState, @@ -40,7 +40,7 @@ from yarl import URL from homeassistant.config_entries import ConfigEntry from homeassistant.const import CONF_URL from homeassistant.core import HomeAssistant, callback -from homeassistant.exceptions import ConfigEntryAuthFailed +from homeassistant.exceptions import ConfigEntryAuthFailed, HomeAssistantError import homeassistant.helpers.device_registry as dr from homeassistant.helpers.device_registry import DeviceEntryType from homeassistant.helpers.update_coordinator import DataUpdateCoordinator, UpdateFailed @@ -196,6 +196,27 @@ class PortainerBaseCoordinator[_DataT](DataUpdateCoordinator[_DataT]): translation_key="timeout_connect", ) from err + async def async_call_portainer(self, coroutine: Awaitable[Any]) -> None: + """Await a Portainer call, mapping library errors to HomeAssistantError.""" + try: + await coroutine + except PortainerAuthenticationError as err: + self.config_entry.async_start_reauth(self.hass) + raise HomeAssistantError( + translation_domain=DOMAIN, + translation_key="invalid_auth", + ) from err + except PortainerConnectionError as err: + raise HomeAssistantError( + translation_domain=DOMAIN, + translation_key="cannot_connect", + ) from err + except PortainerTimeoutError as err: + raise HomeAssistantError( + translation_domain=DOMAIN, + translation_key="timeout_connect", + ) from err + class PortainerCoordinator( PortainerBaseCoordinator[dict[int, PortainerCoordinatorData]] diff --git a/homeassistant/components/portainer/services.py b/homeassistant/components/portainer/services.py index 2294e62ccf9b..2debad41c069 100644 --- a/homeassistant/components/portainer/services.py +++ b/homeassistant/components/portainer/services.py @@ -16,7 +16,6 @@ from homeassistant.helpers import ( from .const import DOMAIN from .coordinator import PortainerConfigEntry -from .util import async_call_portainer class PortainerService(StrEnum): @@ -130,12 +129,12 @@ async def prune_images(call: ServiceCall) -> None: coordinator = config_entry.runtime_data endpoint_id = _async_get_endpoint_id(device, config_entry) - await async_call_portainer( + await coordinator.async_call_portainer( coordinator.portainer.images_prune( endpoint_id=endpoint_id, until=call.data.get(PortainerServiceArgument.UNTIL), dangling=call.data.get(PortainerServiceArgument.DANGLING, False), - ) + ), ) @@ -145,12 +144,12 @@ async def prune_build_cache(call: ServiceCall) -> None: coordinator = config_entry.runtime_data endpoint_id = _async_get_endpoint_id(device, config_entry) - await async_call_portainer( + await coordinator.async_call_portainer( coordinator.portainer.prune_build_cache( endpoint_id, all_cache=call.data[PortainerServiceArgument.ALL], until=call.data.get(PortainerServiceArgument.UNTIL), - ) + ), ) @@ -165,13 +164,13 @@ async def recreate_container(call: ServiceCall) -> None: ) timeout: timedelta | None = call.data.get(PortainerServiceArgument.TIMEOUT) - await async_call_portainer( + await coordinator.async_call_portainer( coordinator.portainer.container_recreate( endpoint_id=endpoint_id, container_id=container_id, **({"timeout": timeout} if timeout is not None else {}), pull_image=call.data.get(PortainerServiceArgument.PULL_IMAGE, False), - ) + ), ) await coordinator.async_request_refresh() diff --git a/homeassistant/components/portainer/switch.py b/homeassistant/components/portainer/switch.py index a9c93666387d..5193428eddaf 100644 --- a/homeassistant/components/portainer/switch.py +++ b/homeassistant/components/portainer/switch.py @@ -25,7 +25,6 @@ from .entity import ( PortainerCoordinatorData, PortainerStackEntity, ) -from .util import async_call_portainer @dataclass(frozen=True, kw_only=True) @@ -54,7 +53,7 @@ async def _perform_action( coroutine: Coroutine[Any, Any, Any], ) -> None: """Perform a Portainer action with error handling and coordinator refresh.""" - await async_call_portainer(coroutine) + await coordinator.async_call_portainer(coroutine) await coordinator.async_request_refresh() diff --git a/homeassistant/components/portainer/update.py b/homeassistant/components/portainer/update.py index b521034b928c..c39712d8e2bf 100644 --- a/homeassistant/components/portainer/update.py +++ b/homeassistant/components/portainer/update.py @@ -6,10 +6,6 @@ from datetime import timedelta from typing import Any, override from pyportainer import Portainer -from pyportainer.exceptions import ( - PortainerAuthenticationError, - PortainerConnectionError, -) from pyportainer.models.docker import ( DockerContainer, LocalImageInformation, @@ -23,10 +19,8 @@ from homeassistant.components.update import ( ) from homeassistant.const import EntityCategory from homeassistant.core import HomeAssistant -from homeassistant.exceptions import HomeAssistantError from homeassistant.helpers.entity_platform import AddConfigEntryEntitiesCallback -from .const import DOMAIN from .coordinator import ( PortainerConfigEntry, PortainerContainerData, @@ -164,22 +158,11 @@ class PortainerContainerImageUpdateEntity(PortainerContainerEntity, UpdateEntity self, version: str | None, backup: bool, **kwargs: Any ) -> None: """Install update.""" - try: - await self.entity_description.update_func( + await self.coordinator.async_call_portainer( + self.entity_description.update_func( self.coordinator.portainer, self.endpoint_id, self.container_data.container.id, - ) - except PortainerAuthenticationError as ex: - self.coordinator.config_entry.async_start_reauth(self.hass) - raise HomeAssistantError( - translation_domain=DOMAIN, - translation_key="invalid_auth", - ) from ex - except PortainerConnectionError as ex: - raise HomeAssistantError( - translation_domain=DOMAIN, - translation_key="cannot_connect", - ) from ex - else: - await self.coordinator.async_request_refresh() + ), + ) + await self.coordinator.async_request_refresh() diff --git a/homeassistant/components/portainer/util.py b/homeassistant/components/portainer/util.py index e4f7d463e304..8a4e503e976c 100644 --- a/homeassistant/components/portainer/util.py +++ b/homeassistant/components/portainer/util.py @@ -1,40 +1,6 @@ """Utility functions for the Portainer integration.""" -from collections.abc import Coroutine -from typing import Any - -from pyportainer import ( - PortainerAuthenticationError, - PortainerConnectionError, - PortainerTimeoutError, -) - -from homeassistant.exceptions import HomeAssistantError - -from .const import DOMAIN - def sanitize_container_name(container_name: str) -> str: """Sanitize to get a proper container name.""" return container_name.replace("/", " ").strip() - - -async def async_call_portainer(coroutine: Coroutine[Any, Any, Any]) -> None: - """Await a Portainer call, mapping library errors to HomeAssistantError.""" - try: - await coroutine - except PortainerAuthenticationError as err: - raise HomeAssistantError( - translation_domain=DOMAIN, - translation_key="invalid_auth", - ) from err - except PortainerConnectionError as err: - raise HomeAssistantError( - translation_domain=DOMAIN, - translation_key="cannot_connect", - ) from err - except PortainerTimeoutError as err: - raise HomeAssistantError( - translation_domain=DOMAIN, - translation_key="timeout_connect", - ) from err diff --git a/tests/components/portainer/test_button.py b/tests/components/portainer/test_button.py index 83db541c97ac..ebb1d5ea745e 100644 --- a/tests/components/portainer/test_button.py +++ b/tests/components/portainer/test_button.py @@ -14,6 +14,7 @@ from syrupy.assertion import SnapshotAssertion from homeassistant.components.button import SERVICE_PRESS from homeassistant.components.portainer.const import DOMAIN +from homeassistant.config_entries import SOURCE_REAUTH from homeassistant.const import ATTR_ENTITY_ID, Platform from homeassistant.core import HomeAssistant from homeassistant.exceptions import HomeAssistantError @@ -112,6 +113,31 @@ async def test_buttons_containers_exceptions( ) +async def test_buttons_invalid_auth_starts_reauth( + hass: HomeAssistant, + mock_portainer_client: AsyncMock, + mock_config_entry: MockConfigEntry, +) -> None: + """Test an invalid API key on press starts a reauth flow.""" + await setup_integration(hass, mock_config_entry) + mock_portainer_client.restart_container.side_effect = PortainerAuthenticationError( + "auth" + ) + + with pytest.raises(HomeAssistantError): + await hass.services.async_call( + BUTTON_DOMAIN, + SERVICE_PRESS, + {ATTR_ENTITY_ID: "button.practical_morse_restart_container"}, + blocking=True, + ) + await hass.async_block_till_done() + + flows = hass.config_entries.flow.async_progress() + assert len(flows) == 1 + assert flows[0]["context"]["source"] == SOURCE_REAUTH + + @pytest.mark.parametrize( ("action", "client_method"), [ diff --git a/tests/components/portainer/test_services.py b/tests/components/portainer/test_services.py index 70ce63e7a3d2..332cd9cb73e4 100644 --- a/tests/components/portainer/test_services.py +++ b/tests/components/portainer/test_services.py @@ -17,6 +17,7 @@ from homeassistant.components.portainer.services import ( PortainerServiceArgument, _async_get_device_and_entry, ) +from homeassistant.config_entries import SOURCE_REAUTH from homeassistant.const import ATTR_DEVICE_ID from homeassistant.core import Context, HomeAssistant from homeassistant.exceptions import ( @@ -508,3 +509,33 @@ async def test_service_portainer_exceptions( blocking=True, ) mock_portainer_client.images_prune.assert_called_once() + + +async def test_service_invalid_auth_starts_reauth( + hass: HomeAssistant, + device_registry: DeviceRegistry, + mock_portainer_client: AsyncMock, + mock_config_entry: MockConfigEntry, +) -> None: + """Test an invalid API key in an action starts a reauth flow.""" + await setup_integration(hass, mock_config_entry) + device = device_registry.async_get_device_by_identifier( + (DOMAIN, TEST_DEVICE_IDENTIFIER), mock_config_entry.entry_id + ) + assert device is not None + mock_portainer_client.images_prune.side_effect = PortainerAuthenticationError( + "auth" + ) + + with pytest.raises(HomeAssistantError): + await hass.services.async_call( + DOMAIN, + PortainerService.PRUNE_IMAGES, + {ATTR_DEVICE_ID: device.id}, + blocking=True, + ) + await hass.async_block_till_done() + + flows = hass.config_entries.flow.async_progress() + assert len(flows) == 1 + assert flows[0]["context"]["source"] == SOURCE_REAUTH diff --git a/tests/components/portainer/test_switch.py b/tests/components/portainer/test_switch.py index 026f852475a2..4d82063ba7fd 100644 --- a/tests/components/portainer/test_switch.py +++ b/tests/components/portainer/test_switch.py @@ -15,6 +15,7 @@ from homeassistant.components.switch import ( SERVICE_TURN_OFF, SERVICE_TURN_ON, ) +from homeassistant.config_entries import SOURCE_REAUTH from homeassistant.const import ATTR_ENTITY_ID, Platform from homeassistant.core import HomeAssistant from homeassistant.exceptions import HomeAssistantError @@ -140,3 +141,28 @@ async def test_turn_off_on_exceptions( {ATTR_ENTITY_ID: entity_id}, blocking=True, ) + + +async def test_switch_invalid_auth_starts_reauth( + hass: HomeAssistant, + mock_portainer_client: AsyncMock, + mock_config_entry: MockConfigEntry, +) -> None: + """Test an invalid API key on a switch action starts a reauth flow.""" + await setup_integration(hass, mock_config_entry) + mock_portainer_client.stop_container.side_effect = PortainerAuthenticationError( + "auth" + ) + + with pytest.raises(HomeAssistantError): + await hass.services.async_call( + SWITCH_DOMAIN, + SERVICE_TURN_OFF, + {ATTR_ENTITY_ID: "switch.practical_morse_container"}, + blocking=True, + ) + await hass.async_block_till_done() + + flows = hass.config_entries.flow.async_progress() + assert len(flows) == 1 + assert flows[0]["context"]["source"] == SOURCE_REAUTH diff --git a/tests/components/portainer/test_update.py b/tests/components/portainer/test_update.py index 832a7be5e0a9..292cbee32154 100644 --- a/tests/components/portainer/test_update.py +++ b/tests/components/portainer/test_update.py @@ -7,6 +7,7 @@ from freezegun.api import FrozenDateTimeFactory from pyportainer.exceptions import ( PortainerAuthenticationError, PortainerConnectionError, + PortainerTimeoutError, ) from pyportainer.models.docker import DockerContainer, PortainerImageUpdateStatus from pyportainer.watcher import PortainerImageWatcherResult @@ -16,6 +17,7 @@ from syrupy.assertion import SnapshotAssertion from homeassistant.components.portainer.const import DOMAIN from homeassistant.components.portainer.coordinator import DEFAULT_SCAN_INTERVAL from homeassistant.components.update import ATTR_INSTALLED_VERSION, ATTR_LATEST_VERSION +from homeassistant.config_entries import SOURCE_REAUTH from homeassistant.const import STATE_OFF, STATE_ON, STATE_UNKNOWN, Platform from homeassistant.core import HomeAssistant from homeassistant.exceptions import HomeAssistantError @@ -95,6 +97,7 @@ async def test_update_install( [ (PortainerAuthenticationError("auth"), "invalid_auth_no_details"), (PortainerConnectionError("conn"), "cannot_connect_no_details"), + (PortainerTimeoutError("timeout"), "timeout_connect_no_details"), ], ) async def test_update_install_errors( @@ -123,6 +126,37 @@ async def test_update_install_errors( ) +async def test_update_install_invalid_auth_starts_reauth( + hass: HomeAssistant, + mock_portainer_client: AsyncMock, + mock_portainer_watcher: MagicMock, + mock_config_entry: MockConfigEntry, +) -> None: + """Test an invalid API key on install starts a reauth flow.""" + mock_portainer_client.container_recreate.side_effect = PortainerAuthenticationError( + "auth" + ) + + with patch( + "homeassistant.components.portainer._PLATFORMS", + [Platform.UPDATE], + ): + await setup_integration(hass, mock_config_entry) + + with pytest.raises(HomeAssistantError): + await hass.services.async_call( + "update", + "install", + {"entity_id": ENTITY_ID}, + blocking=True, + ) + await hass.async_block_till_done() + + flows = hass.config_entries.flow.async_progress() + assert len(flows) == 1 + assert flows[0]["context"]["source"] == SOURCE_REAUTH + + @pytest.mark.parametrize("repo_digests", [None, []], ids=["missing", "empty"]) async def test_update_installed_version_without_repo_digest( hass: HomeAssistant,