diff --git a/homeassistant/components/anglian_water/config_flow.py b/homeassistant/components/anglian_water/config_flow.py index bd16e88aa777..5beade71aa92 100644 --- a/homeassistant/components/anglian_water/config_flow.py +++ b/homeassistant/components/anglian_water/config_flow.py @@ -1,5 +1,6 @@ """Config flow for the Anglian Water integration.""" +from collections.abc import Mapping import logging from typing import TYPE_CHECKING, Any, override @@ -9,10 +10,12 @@ from pyanglianwater import AnglianWater from pyanglianwater.auth import MSOB2CAuth from pyanglianwater.exceptions import ( ConsentRequiredError, + ExpiredAccessTokenError, InvalidAccountIdError, MFARequiredError, SelfAssertedError, SmartMeterUnavailableError, + UnknownEndpointError, ) from homeassistant.config_entries import ConfigFlow, ConfigFlowResult @@ -45,7 +48,7 @@ STEP_MFA_DATA_SCHEMA = probatio.Schema( ) -async def validate_credentials(auth: MSOB2CAuth) -> str | MSOB2CAuth: +async def validate_credentials(auth: MSOB2CAuth) -> str | None: """Validate the provided credentials.""" try: await auth.send_login_request() @@ -58,7 +61,7 @@ async def validate_credentials(auth: MSOB2CAuth) -> str | MSOB2CAuth: except Exception: _LOGGER.exception("Unexpected exception") return "unknown" - return auth + return None def humanize_account_data(account: dict) -> str: @@ -83,14 +86,14 @@ async def get_accounts(auth: MSOB2CAuth) -> list[selector.SelectOptionDict]: ] -async def validate_account(auth: MSOB2CAuth, account_number: str) -> str | MSOB2CAuth: +async def validate_account(auth: MSOB2CAuth, account_number: str) -> str | None: """Validate the provided account number.""" _aw = AnglianWater(authenticator=auth) try: await _aw.validate_smart_meter(account_number) except InvalidAccountIdError, SmartMeterUnavailableError: return "smart_meter_unavailable" - return auth + return None class AnglianWaterConfigFlow(ConfigFlow, domain=DOMAIN): @@ -102,6 +105,44 @@ class AnglianWaterConfigFlow(ConfigFlow, domain=DOMAIN): self.accounts: list[selector.SelectOptionDict] = [] self.user_input: dict[str, Any] | None = None + def _create_authenticator(self, user_input: dict[str, Any]) -> MSOB2CAuth: + """Create an MSOB2CAuth instance with the provided user input.""" + return MSOB2CAuth( + username=user_input[CONF_USERNAME], + password=user_input[CONF_PASSWORD], + session=async_create_clientsession( + self.hass, + cookie_jar=CookieJar(quote_cookie=False), + ), + ) + + async def _async_validate_mfa(self, code: str) -> str | None: + """Validate the provided MFA code.""" + if TYPE_CHECKING: + assert self.authenticator + try: + await self.authenticator.send_mfa_request(code) + except MFARequiredError: + return "invalid_code" + except Exception: + _LOGGER.exception("Unexpected exception") + return "unknown" + return None + + async def _async_get_accounts(self) -> str | None: + """Retrieve the list of accounts associated with the authenticated user.""" + if TYPE_CHECKING: + assert self.authenticator + try: + self.accounts = await get_accounts(self.authenticator) + except ExpiredAccessTokenError, UnknownEndpointError: + _LOGGER.exception("Error retrieving accounts") + return "cannot_connect" + except Exception: + _LOGGER.exception("Unexpected exception") + return "unknown" + return None + @override async def async_step_user( self, user_input: dict[str, Any] | None = None @@ -109,24 +150,19 @@ class AnglianWaterConfigFlow(ConfigFlow, domain=DOMAIN): """Handle the initial step.""" errors: dict[str, str] = {} if user_input is not None: - self.authenticator = MSOB2CAuth( - username=user_input[CONF_USERNAME], - password=user_input[CONF_PASSWORD], - session=async_create_clientsession( - self.hass, - cookie_jar=CookieJar(quote_cookie=False), - ), - ) - validation_response = await validate_credentials(self.authenticator) - if isinstance(validation_response, str): - if validation_response == "mfa_required": - self.user_input = user_input - return await self.async_step_mfa() - errors["base"] = validation_response - else: - self.accounts = await get_accounts(self.authenticator) + self.authenticator = self._create_authenticator(user_input) + validation_error = await validate_credentials(self.authenticator) + if validation_error == "mfa_required": self.user_input = user_input - return await self.async_step_select_account() + return await self.async_step_mfa() + if validation_error: + errors["base"] = validation_error + else: + account_error = await self._async_get_accounts() + if not account_error: + self.user_input = user_input + return await self.async_step_select_account() + errors["base"] = account_error return self.async_show_form( step_id="user", data_schema=STEP_USER_DATA_SCHEMA, errors=errors @@ -140,16 +176,14 @@ class AnglianWaterConfigFlow(ConfigFlow, domain=DOMAIN): if user_input is not None: if TYPE_CHECKING: assert self.authenticator - try: - await self.authenticator.send_mfa_request(user_input[CONF_CODE]) - except MFARequiredError: - errors["base"] = "invalid_code" - except Exception: - _LOGGER.exception("Unexpected exception") - errors["base"] = "unknown" + error = await self._async_validate_mfa(user_input[CONF_CODE]) + if error: + errors["base"] = error else: - self.accounts = await get_accounts(self.authenticator) - return await self.async_step_select_account() + account_error = await self._async_get_accounts() + if not account_error: + return await self.async_step_select_account() + errors["base"] = account_error return self.async_show_form( step_id="mfa", data_schema=STEP_MFA_DATA_SCHEMA, errors=errors ) @@ -173,10 +207,9 @@ class AnglianWaterConfigFlow(ConfigFlow, domain=DOMAIN): self.authenticator, user_input[CONF_ACCOUNT_NUMBER], ) - if isinstance(validation_result, str): - errors["base"] = validation_result - else: + if not validation_result: return await self.async_step_complete(user_input) + errors["base"] = validation_result return self.async_show_form( step_id="select_account", data_schema=probatio.Schema( @@ -209,3 +242,69 @@ class AnglianWaterConfigFlow(ConfigFlow, domain=DOMAIN): title=user_input[CONF_ACCOUNT_NUMBER], data=config_entry_data, ) + + async def async_step_reauth( + self, entry_data: Mapping[str, Any] + ) -> ConfigFlowResult: + """Initial configuration step via reauthentication.""" + return await self.async_step_reauth_confirm() + + async def async_step_reauth_confirm( + self, user_input: dict[str, Any] | None = None + ) -> ConfigFlowResult: + """Handle receiving username/password.""" + errors: dict[str, str] = {} + if user_input is not None: + self.authenticator = self._create_authenticator(user_input) + validation_response = await validate_credentials(self.authenticator) + if not validation_response: + return await self._async_finish_reauth() + if validation_response == "mfa_required": + self.user_input = user_input + return await self.async_step_reauth_mfa() + errors["base"] = validation_response + return self.async_show_form( + step_id="reauth_confirm", data_schema=STEP_USER_DATA_SCHEMA, errors=errors + ) + + async def async_step_reauth_mfa( + self, user_input: dict[str, Any] | None = None + ) -> ConfigFlowResult: + """Handle the MFA step during reauthentication.""" + errors: dict[str, str] = {} + if user_input is not None: + if TYPE_CHECKING: + assert self.authenticator + error = await self._async_validate_mfa(user_input[CONF_CODE]) + if not error: + return await self._async_finish_reauth() + errors["base"] = error + return self.async_show_form( + step_id="reauth_mfa", data_schema=STEP_MFA_DATA_SCHEMA, errors=errors + ) + + async def _async_finish_reauth(self) -> ConfigFlowResult: + """Verify the account and update its access token.""" + if TYPE_CHECKING: + assert self.authenticator + entry = self._get_reauth_entry() + account_error = await self._async_get_accounts() + if account_error: + return self.async_show_form( + step_id="reauth_confirm", + data_schema=STEP_USER_DATA_SCHEMA, + errors={"base": account_error}, + ) + if not any( + account["value"] == entry.data[CONF_ACCOUNT_NUMBER] + for account in self.accounts + ): + return self.async_show_form( + step_id="reauth_confirm", + data_schema=STEP_USER_DATA_SCHEMA, + errors={"base": "account_not_found"}, + ) + return self.async_update_reload_and_abort( + entry, + data_updates={CONF_ACCESS_TOKEN: self.authenticator.refresh_token}, + ) diff --git a/homeassistant/components/anglian_water/quality_scale.yaml b/homeassistant/components/anglian_water/quality_scale.yaml index 4b6eae7e99a7..6092f8b3271a 100644 --- a/homeassistant/components/anglian_water/quality_scale.yaml +++ b/homeassistant/components/anglian_water/quality_scale.yaml @@ -43,7 +43,7 @@ rules: integration-owner: done log-when-unavailable: done parallel-updates: done - reauthentication-flow: todo + reauthentication-flow: done test-coverage: done # Gold diff --git a/homeassistant/components/anglian_water/strings.json b/homeassistant/components/anglian_water/strings.json index d117b6629325..e838928f4357 100644 --- a/homeassistant/components/anglian_water/strings.json +++ b/homeassistant/components/anglian_water/strings.json @@ -4,6 +4,7 @@ "already_configured": "[%key:common::config_flow::abort::already_configured_device%]" }, "error": { + "account_not_found": "These credentials do not have access to the configured billing account.", "cannot_connect": "[%key:common::config_flow::error::cannot_connect%]", "consent_required": "You need to accept the terms and conditions for your Anglian Water account before using this integration. Log in to their website for further information.", "invalid_auth": "[%key:common::config_flow::error::invalid_auth%]", @@ -21,6 +22,26 @@ }, "description": "Two-factor authentication is enabled on your account. Check your email for the security code. It will come from `no-reply@anglianwater.co.uk`." }, + "reauth_confirm": { + "data": { + "password": "[%key:common::config_flow::data::password%]", + "username": "[%key:common::config_flow::data::username%]" + }, + "data_description": { + "password": "[%key:component::anglian_water::config::step::user::data_description::password%]", + "username": "[%key:component::anglian_water::config::step::user::data_description::username%]" + }, + "description": "[%key:component::anglian_water::config::step::user::description%]" + }, + "reauth_mfa": { + "data": { + "code": "[%key:component::anglian_water::config::step::mfa::data::code%]" + }, + "data_description": { + "code": "[%key:component::anglian_water::config::step::mfa::data_description::code%]" + }, + "description": "[%key:component::anglian_water::config::step::mfa::description%]" + }, "select_account": { "data": { "account_number": "Billing account number" diff --git a/tests/components/anglian_water/test_config_flow.py b/tests/components/anglian_water/test_config_flow.py index dd984b0d51fa..7b52a0364d26 100644 --- a/tests/components/anglian_water/test_config_flow.py +++ b/tests/components/anglian_water/test_config_flow.py @@ -4,10 +4,12 @@ from unittest.mock import AsyncMock from pyanglianwater.exceptions import ( ConsentRequiredError, + ExpiredAccessTokenError, InvalidAccountIdError, MFARequiredError, SelfAssertedError, SmartMeterUnavailableError, + UnknownEndpointError, ) import pytest @@ -232,6 +234,98 @@ async def test_single_account_flow_with_mfa_exception( assert result["result"].unique_id == ACCOUNT_NUMBER +@pytest.mark.parametrize( + ("exception_type", "expected_error"), + [ + (ExpiredAccessTokenError, "cannot_connect"), + ( + UnknownEndpointError(status=500, response="Service Unavailable"), + "cannot_connect", + ), + (ValueError, "unknown"), + ], +) +async def test_account_fetch_exception( + hass: HomeAssistant, + mock_anglian_water_authenticator: AsyncMock, + mock_anglian_water_client: AsyncMock, + exception_type: Exception, + expected_error: str, +) -> None: + """Test that the flow handles account-fetch exceptions.""" + result = await hass.config_entries.flow.async_init( + DOMAIN, context={"source": SOURCE_USER} + ) + assert result is not None + assert result["type"] is FlowResultType.FORM + assert result["step_id"] == "user" + + mock_anglian_water_client.api.get_associated_accounts.side_effect = exception_type + + result = await hass.config_entries.flow.async_configure( + result["flow_id"], + user_input={ + CONF_USERNAME: USERNAME, + CONF_PASSWORD: PASSWORD, + }, + ) + + assert result["type"] is FlowResultType.FORM + assert result["step_id"] == "user" + assert result["errors"] == {"base": expected_error} + + +@pytest.mark.parametrize( + ("exception_type", "expected_error"), + [ + (ExpiredAccessTokenError, "cannot_connect"), + ( + UnknownEndpointError(status=500, response="Service Unavailable"), + "cannot_connect", + ), + (ValueError, "unknown"), + ], +) +async def test_mfa_account_fetch_exception( + hass: HomeAssistant, + mock_anglian_water_authenticator: AsyncMock, + mock_anglian_water_client: AsyncMock, + exception_type: Exception, + expected_error: str, +) -> None: + """Test that the MFA flow handles account-fetch exceptions.""" + result = await hass.config_entries.flow.async_init( + DOMAIN, context={"source": SOURCE_USER} + ) + assert result is not None + assert result["type"] is FlowResultType.FORM + assert result["step_id"] == "user" + + mock_anglian_water_authenticator.send_login_request.side_effect = MFARequiredError + + result = await hass.config_entries.flow.async_configure( + result["flow_id"], + user_input={ + CONF_USERNAME: USERNAME, + CONF_PASSWORD: PASSWORD, + }, + ) + + assert result["type"] is FlowResultType.FORM + assert result["step_id"] == "mfa" + + mock_anglian_water_client.api.get_associated_accounts.side_effect = exception_type + + result = await hass.config_entries.flow.async_configure( + result["flow_id"], + user_input={CONF_CODE: "123456"}, + ) + + assert result["type"] is FlowResultType.FORM + assert result["step_id"] == "mfa" + assert result["errors"] == {"base": expected_error} + + @pytest.mark.usefixtures("mock_setup_entry") async def test_already_configured( hass: HomeAssistant, @@ -403,3 +497,315 @@ async def test_account_recover_exception( assert result["data"][CONF_ACCESS_TOKEN] == ACCESS_TOKEN assert result["data"][CONF_ACCOUNT_NUMBER] == ACCOUNT_NUMBER assert result["result"].unique_id == ACCOUNT_NUMBER + + +async def test_reauth_flow( + hass: HomeAssistant, + mock_config_entry: MockConfigEntry, + mock_anglian_water_authenticator: AsyncMock, + mock_anglian_water_client: AsyncMock, +) -> None: + """Test the reauth flow.""" + mock_config_entry.add_to_hass(hass) + mock_anglian_water_authenticator.refresh_token = "new_access_token" + original_data = dict(mock_config_entry.data) + + result = await hass.config_entries.flow.async_init( + DOMAIN, + context={ + "source": config_entries.SOURCE_REAUTH, + "entry_id": mock_config_entry.entry_id, + }, + data=mock_config_entry.data, + ) + assert result is not None + assert result["type"] is FlowResultType.FORM + assert result["step_id"] == "reauth_confirm" + + result = await hass.config_entries.flow.async_configure( + result["flow_id"], + user_input={ + CONF_USERNAME: USERNAME, + CONF_PASSWORD: PASSWORD, + }, + ) + + assert result["type"] is FlowResultType.ABORT + assert result["reason"] == "reauth_successful" + assert mock_config_entry.data == { + **original_data, + CONF_ACCESS_TOKEN: "new_access_token", + } + + +@pytest.mark.parametrize( + ("exception_type", "expected_error"), + [ + (SelfAssertedError, "invalid_auth"), + (ConsentRequiredError, "consent_required"), + (ValueError, "unknown"), + ], +) +async def test_reauth_flow_auth_exception( + hass: HomeAssistant, + mock_config_entry: MockConfigEntry, + mock_anglian_water_authenticator: AsyncMock, + mock_anglian_water_client: AsyncMock, + exception_type: type[Exception], + expected_error: str, +) -> None: + """Test that the reauth flow can recover from an auth exception.""" + mock_config_entry.add_to_hass(hass) + + result = await hass.config_entries.flow.async_init( + DOMAIN, + context={ + "source": config_entries.SOURCE_REAUTH, + "entry_id": mock_config_entry.entry_id, + }, + data=mock_config_entry.data, + ) + assert result is not None + assert result["type"] is FlowResultType.FORM + assert result["step_id"] == "reauth_confirm" + + mock_anglian_water_authenticator.send_login_request.side_effect = exception_type + + result = await hass.config_entries.flow.async_configure( + result["flow_id"], + user_input={ + CONF_USERNAME: USERNAME, + CONF_PASSWORD: PASSWORD, + }, + ) + + assert result["type"] is FlowResultType.FORM + assert result["step_id"] == "reauth_confirm" + assert result["errors"] == {"base": expected_error} + + mock_anglian_water_authenticator.send_login_request.side_effect = None + + result = await hass.config_entries.flow.async_configure( + result["flow_id"], + user_input={ + CONF_USERNAME: USERNAME, + CONF_PASSWORD: PASSWORD, + }, + ) + assert result["type"] is FlowResultType.ABORT + assert result["reason"] == "reauth_successful" + + +async def test_reauth_flow_account_not_found( + hass: HomeAssistant, + mock_config_entry: MockConfigEntry, + mock_anglian_water_authenticator: AsyncMock, + mock_anglian_water_client: AsyncMock, +) -> None: + """Test that reauth does not update an entry for another account.""" + mock_config_entry.add_to_hass(hass) + original_data = dict(mock_config_entry.data) + mock_anglian_water_client.api.get_associated_accounts.return_value = { + "result": {"active": []} + } + + result = await hass.config_entries.flow.async_init( + DOMAIN, + context={ + "source": config_entries.SOURCE_REAUTH, + "entry_id": mock_config_entry.entry_id, + }, + data=mock_config_entry.data, + ) + assert result is not None + + result = await hass.config_entries.flow.async_configure( + result["flow_id"], + user_input={ + CONF_USERNAME: USERNAME, + CONF_PASSWORD: PASSWORD, + }, + ) + + assert result["type"] is FlowResultType.FORM + assert result["step_id"] == "reauth_confirm" + assert result["errors"] == {"base": "account_not_found"} + assert mock_config_entry.data == original_data + + mock_anglian_water_client.api.get_associated_accounts.return_value = ( + await async_load_json_object_fixture( + hass, "multi_associated_accounts.json", DOMAIN + ) + ) + result = await hass.config_entries.flow.async_configure( + result["flow_id"], + user_input={ + CONF_USERNAME: USERNAME, + CONF_PASSWORD: PASSWORD, + }, + ) + + assert result["type"] is FlowResultType.ABORT + assert result["reason"] == "reauth_successful" + + +async def test_reauth_flow_mfa_required( + hass: HomeAssistant, + mock_config_entry: MockConfigEntry, + mock_anglian_water_authenticator: AsyncMock, + mock_anglian_water_client: AsyncMock, +) -> None: + """Test the reauth flow with MFA required.""" + mock_config_entry.add_to_hass(hass) + + result = await hass.config_entries.flow.async_init( + DOMAIN, + context={ + "source": config_entries.SOURCE_REAUTH, + "entry_id": mock_config_entry.entry_id, + }, + data=mock_config_entry.data, + ) + assert result is not None + assert result["type"] is FlowResultType.FORM + assert result["step_id"] == "reauth_confirm" + + mock_anglian_water_authenticator.send_login_request.side_effect = MFARequiredError + + result = await hass.config_entries.flow.async_configure( + result["flow_id"], + user_input={ + CONF_USERNAME: USERNAME, + CONF_PASSWORD: PASSWORD, + }, + ) + + assert result["type"] is FlowResultType.FORM + assert result["step_id"] == "reauth_mfa" + + result = await hass.config_entries.flow.async_configure( + result["flow_id"], + user_input={ + CONF_CODE: "123456", + }, + ) + + assert result["type"] is FlowResultType.ABORT + assert result["reason"] == "reauth_successful" + + +@pytest.mark.parametrize( + ("exception_type", "expected_error"), + [ + (MFARequiredError, "invalid_code"), + (ValueError, "unknown"), + ], +) +async def test_reauth_flow_mfa_exception( + hass: HomeAssistant, + mock_config_entry: MockConfigEntry, + mock_anglian_water_authenticator: AsyncMock, + mock_anglian_water_client: AsyncMock, + exception_type: type[Exception], + expected_error: str, +) -> None: + """Test that the reauth flow can recover from an MFA exception.""" + mock_config_entry.add_to_hass(hass) + + result = await hass.config_entries.flow.async_init( + DOMAIN, + context={ + "source": config_entries.SOURCE_REAUTH, + "entry_id": mock_config_entry.entry_id, + }, + data=mock_config_entry.data, + ) + assert result is not None + assert result["type"] is FlowResultType.FORM + assert result["step_id"] == "reauth_confirm" + + mock_anglian_water_authenticator.send_login_request.side_effect = MFARequiredError + + result = await hass.config_entries.flow.async_configure( + result["flow_id"], + user_input={ + CONF_USERNAME: USERNAME, + CONF_PASSWORD: PASSWORD, + }, + ) + + assert result["type"] is FlowResultType.FORM + assert result["step_id"] == "reauth_mfa" + + mock_anglian_water_authenticator.send_mfa_request.side_effect = exception_type + + result = await hass.config_entries.flow.async_configure( + result["flow_id"], + user_input={ + CONF_CODE: "123456", + }, + ) + + assert result["type"] is FlowResultType.FORM + assert result["step_id"] == "reauth_mfa" + assert result["errors"] == {"base": expected_error} + + mock_anglian_water_authenticator.send_mfa_request.side_effect = None + + result = await hass.config_entries.flow.async_configure( + result["flow_id"], + user_input={ + CONF_CODE: "123456", + }, + ) + + assert result["type"] is FlowResultType.ABORT + assert result["reason"] == "reauth_successful" + + +@pytest.mark.parametrize( + ("exception", "expected_error"), + [ + (ExpiredAccessTokenError, "cannot_connect"), + ( + UnknownEndpointError(status=500, response="Service Unavailable"), + "cannot_connect", + ), + (ValueError, "unknown"), + ], +) +async def test_reauth_flow_account_fetch_exception( + hass: HomeAssistant, + mock_config_entry: MockConfigEntry, + mock_anglian_water_authenticator: AsyncMock, + mock_anglian_water_client: AsyncMock, + exception: Exception, + expected_error: str, +) -> None: + """Test that reauth handles account-fetch exceptions.""" + mock_config_entry.add_to_hass(hass) + original_data = dict(mock_config_entry.data) + + result = await hass.config_entries.flow.async_init( + DOMAIN, + context={ + "source": config_entries.SOURCE_REAUTH, + "entry_id": mock_config_entry.entry_id, + }, + data=mock_config_entry.data, + ) + + mock_anglian_water_client.api.get_associated_accounts.side_effect = exception + + result = await hass.config_entries.flow.async_configure( + result["flow_id"], + user_input={ + CONF_USERNAME: USERNAME, + CONF_PASSWORD: PASSWORD, + }, + ) + + assert result["type"] is FlowResultType.FORM + assert result["step_id"] == "reauth_confirm" + assert result["errors"] == {"base": expected_error} + assert mock_config_entry.data == original_data