diff --git a/homeassistant/components/smtp/notify.py b/homeassistant/components/smtp/notify.py index c17753de588f..be9fb367b5f1 100644 --- a/homeassistant/components/smtp/notify.py +++ b/homeassistant/components/smtp/notify.py @@ -40,11 +40,7 @@ from homeassistant.const import ( ) from homeassistant.core import HomeAssistant from homeassistant.data_entry_flow import FlowResultType -from homeassistant.exceptions import ( - ConfigEntryAuthFailed, - HomeAssistantError, - ServiceValidationError, -) +from homeassistant.exceptions import HomeAssistantError, ServiceValidationError from homeassistant.helpers import config_validation as cv, entity_registry as er from homeassistant.helpers.device_registry import DeviceEntryType, DeviceInfo from homeassistant.helpers.entity_platform import AddConfigEntryEntitiesCallback @@ -325,7 +321,8 @@ class MailNotifyEntity(NotifyEntity): await client.send_message(msg) break except aiosmtplib.SMTPAuthenticationError as e: - raise ConfigEntryAuthFailed( + self._entry.async_start_reauth(self.hass) + raise HomeAssistantError( translation_domain=DOMAIN, translation_key="authentication_error", ) from e diff --git a/tests/components/smtp/test_notify.py b/tests/components/smtp/test_notify.py index 529ea2a58c43..8132cd396e1c 100644 --- a/tests/components/smtp/test_notify.py +++ b/tests/components/smtp/test_notify.py @@ -6,7 +6,10 @@ import re from smtplib import SMTPException, SMTPServerDisconnected from unittest.mock import AsyncMock, MagicMock, patch -import aiosmtplib +from aiosmtplib import ( + SMTPAuthenticationError, + SMTPException as aiosmtplib_SMTPException, +) import pytest from syrupy.assertion import SnapshotAssertion @@ -29,7 +32,7 @@ from homeassistant.components.smtp.const import ( DOMAIN, ) from homeassistant.components.smtp.notify import MailNotificationService -from homeassistant.config_entries import ConfigEntryState +from homeassistant.config_entries import SOURCE_REAUTH, ConfigEntryState from homeassistant.const import ATTR_ENTITY_ID, CONF_NAME, CONF_RECIPIENT, STATE_UNKNOWN from homeassistant.core import HomeAssistant from homeassistant.exceptions import HomeAssistantError, ServiceValidationError @@ -266,8 +269,8 @@ async def test_notify_send_message( @pytest.mark.parametrize( ("exception", "translation_key", "call_count"), [ - (aiosmtplib.SMTPAuthenticationError(0, ""), "authentication_error", 1), - (aiosmtplib.SMTPException(""), "send_mail_connection_error", 2), + (SMTPAuthenticationError(0, ""), "authentication_error", 1), + (aiosmtplib_SMTPException(""), "send_mail_connection_error", 2), ], ) @pytest.mark.usefixtures("make_msgid", "smtp") @@ -305,6 +308,48 @@ async def test_notify_send_message_exceptions( assert aiosmtplib.__aenter__.return_value.send_message.call_count == call_count +@pytest.mark.usefixtures("make_msgid", "smtp") +async def test_notify_send_message_reauth_flow( + hass: HomeAssistant, + config_entry: MockConfigEntry, + aiosmtplib: AsyncMock, +) -> None: + """Test authentication error starts reauth flow.""" + + config_entry.add_to_hass(hass) + await hass.config_entries.async_setup(config_entry.entry_id) + await hass.async_block_till_done() + + assert config_entry.state is ConfigEntryState.LOADED + + aiosmtplib.__aenter__.return_value.send_message.side_effect = ( + SMTPAuthenticationError(0, "") + ) + + with pytest.raises(HomeAssistantError) as e: + await hass.services.async_call( + NOTIFY_DOMAIN, + SERVICE_SEND_MESSAGE, + { + ATTR_ENTITY_ID: "notify.home_assistant_recipient", + ATTR_MESSAGE: "Hello World", + }, + blocking=True, + ) + + assert e.value.translation_key == "authentication_error" + assert aiosmtplib.__aenter__.return_value.send_message.call_count == 1 + + flows = hass.config_entries.flow.async_progress() + assert len(flows) == 1 + + flow = flows[0] + assert flow["step_id"] == "reauth_confirm" + assert flow["handler"] == DOMAIN + assert flow["context"]["source"] == SOURCE_REAUTH + assert flow["context"]["entry_id"] == config_entry.entry_id + + @pytest.mark.parametrize("exception", [SMTPServerDisconnected, SMTPException]) @pytest.mark.usefixtures("aiosmtplib") async def test_legacy_notify_exception(