mirror of
https://github.com/home-assistant/core.git
synced 2026-10-06 06:15:47 -04:00
Start Portainer reauth on an invalid API key (#184002)
This commit is contained in:
@@ -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()
|
||||
|
||||
|
||||
|
||||
@@ -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]]
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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()
|
||||
|
||||
|
||||
|
||||
@@ -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()
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"),
|
||||
[
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user