""" Helper functions for Alexa Media Player. SPDX-License-Identifier: Apache-2.0 For more details about this platform, please refer to the documentation at https://community.home-assistant.io/t/echo-devices-alexa-as-media-player-testers-needed/58639 """ import asyncio import functools import hashlib import logging from typing import Any, Callable, Optional, TypeVar, overload from alexapy import AlexapyLoginCloseRequested, AlexapyLoginError, hide_email from alexapy.alexalogin import AlexaLogin from dictor import dictor from homeassistant.const import CONF_EMAIL, CONF_URL from homeassistant.core import HomeAssistant from homeassistant.exceptions import ConditionErrorMessage from homeassistant.helpers.entity import Entity from homeassistant.helpers.instance_id import async_get as async_get_instance_id import wrapt from .const import DATA_ALEXAMEDIA, EXCEPTION_TEMPLATE _LOGGER = logging.getLogger(__name__) ArgType = TypeVar("ArgType") def _norm_filter_token(value: Any) -> str | None: """Normalize a single filter token for reliable matching.""" if value is None: return None s = str(value).strip() if not s: return None return s.casefold() def _coerce_filter(value: Any) -> set[str]: """Coerce include/exclude filter input into a normalized set[str]. Accepts: - None / empty -> empty set - comma-separated str -> split on commas - list/set/tuple -> per-item normalization - anything else -> single token (best effort) """ if not value: return set() # Legacy/back-compat: allow comma-separated string if isinstance(value, str): out = set() for part in value.split(","): token = _norm_filter_token(part) if token: out.add(token) return out if isinstance(value, (list, set, tuple)): out = set() for v in value: token = _norm_filter_token(v) if token: out.add(token) return out token = _norm_filter_token(value) return {token} if token else set() async def add_devices( account: str, devices: list[Entity], add_devices_callback: Callable[[list[Entity], bool], None], include_filter: str | list[str] | set[str] | tuple[str, ...] | None = None, exclude_filter: str | list[str] | set[str] | tuple[str, ...] | None = None, ) -> bool: """Add devices using add_devices_callback.""" include_filter_set = _coerce_filter(include_filter) exclude_filter_set = _coerce_filter(exclude_filter) if include_filter_set: _LOGGER.debug( "%s: include_filter_set: %s", account, include_filter_set, ) if exclude_filter_set: _LOGGER.debug( "%s: exclude_filter_set: %s", account, exclude_filter_set, ) def _device_name(dev: Entity) -> str | None: """Best-effort name before entity_id is assigned. For AMP switches, reconstruct the legacy " switch" name only if those attributes were explicitly set. """ # First prefer explicitly set name attributes (works for tests + most entities) name = ( getattr(dev, "name", None) or getattr(dev, "_attr_name", None) or getattr(dev, "_name", None) or getattr(dev, "_device_name", None) or getattr(dev, "_friendly_name", None) ) if name: return name # Only attempt switch reconstruction if attributes were explicitly defined # (avoids MagicMock auto-attribute trap in tests) dev_dict = getattr(dev, "__dict__", {}) client = dev_dict.get("_client") suffix = dev_dict.get("_unique_id_suffix") if client and suffix: client_dict = getattr(client, "__dict__", {}) base = ( client_dict.get("name") or client_dict.get("_attr_name") or client_dict.get("_name") or client_dict.get("_device_name") ) if base: return f"{base} {suffix} switch" return None def _device_label(dev: Entity) -> str: """Return a compact, stable identifier for logging.""" name = _device_name(dev) entity_id = getattr(dev, "entity_id", None) # often not set yet dev_type = type(dev).__name__ if name and entity_id: return f"{name} ({dev_type}, {entity_id})" if name: return f"{name} ({dev_type})" return f" ({dev_type})" def _devices_preview(devs: list[Entity]) -> str: max_items = 8 labels = [_device_label(d) for d in devs[:max_items]] suffix = f" …(+{len(devs) - max_items} more)" if len(devs) > max_items else "" return ", ".join(labels) + suffix def _filter_devices( devs: list[Entity], include_set: set[str], exclude_set: set[str], ) -> list[Entity]: selected: list[Entity] = [] include_mode = bool(include_set) if include_mode and exclude_set: _LOGGER.debug( "%s: include_devices set; ignoring exclude_devices per documented precedence", account, ) for dev in devs: dev_name = _norm_filter_token(_device_name(dev)) # INCLUDE MODE: only include explicitly listed names if include_mode: if dev_name and dev_name in include_set: selected.append(dev) else: if not dev_name: _LOGGER.debug( "%s: Not including device (no name yet): %s", account, _device_label(dev), ) else: _LOGGER.debug( "%s: Not including device: %s (match key=%r)", account, _device_label(dev), dev_name, ) continue # EXCLUDE MODE: exclude listed names if exclude_set and dev_name and dev_name in exclude_set: _LOGGER.debug( "%s: Excluding device: %s (match key=%r)", account, _device_label(dev), dev_name, ) continue selected.append(dev) return selected devices = _filter_devices(devices, include_filter_set, exclude_filter_set) if not devices: return True _LOGGER.debug( "%s: Adding %d device(s): %s", account, len(devices), _devices_preview(devices), ) try: add_devices_callback(devices, False) except ConditionErrorMessage as exception_: message: str = exception_.message if message.startswith("Entity id already exists"): _LOGGER.debug("%s: Device already added: %s", account, message) else: _LOGGER.debug( "%s: Unable to add %d device(s): %s", account, len(devices), message, ) except Exception as ex: # pylint: disable=broad-except _LOGGER.debug( "%s: Unable to add %d device(s): %s", account, len(devices), EXCEPTION_TEMPLATE.format(type(ex).__name__, ex.args), ) else: return True return False def retry_async( limit: int = 5, delay: float = 1, catch_exceptions: bool = True ) -> Callable: """Wrap function with retry logic. The function will retry until true or the limit is reached. It will delay for the period of time specified exponentially increasing the delay. Parameters ---------- limit : int The max number of retries. delay : float The delay in seconds between retries. catch_exceptions : bool Whether exceptions should be caught and treated as failures or thrown. Returns ------- def Wrapped function. """ def wrap(func) -> Callable: @functools.wraps(func) async def wrapper(*args, **kwargs) -> Any: _LOGGER.debug( "%s.%s: Trying with limit %s delay %s catch_exceptions %s", func.__module__[func.__module__.find(".") + 1 :], func.__name__, limit, delay, catch_exceptions, ) retries: int = 0 result: bool = False next_try: int = 0 while not result and retries < limit: if retries != 0: next_try = delay * 2**retries await asyncio.sleep(next_try) retries += 1 try: result = await func(*args, **kwargs) except Exception as ex: # pylint: disable=broad-except if not catch_exceptions: raise _LOGGER.debug( "%s.%s: failure caught due to exception: %s", func.__module__[func.__module__.find(".") + 1 :], func.__name__, EXCEPTION_TEMPLATE.format(type(ex).__name__, ex.args), ) _LOGGER.debug( "%s.%s: Try: %s/%s after waiting %s seconds result: %s", func.__module__[func.__module__.find(".") + 1 :], func.__name__, retries, limit, next_try, result, ) return result return wrapper return wrap @wrapt.decorator async def _catch_login_errors(func, instance, args, kwargs) -> Any: """Detect AlexapyLoginError and attempt relogin.""" result = None if instance is None and args: instance = args[0] if hasattr(instance, "check_login_changes"): # _LOGGER.debug( # "%s checking for login changes", instance, # ) instance.check_login_changes() try: result = await func(*args, **kwargs) except AlexapyLoginCloseRequested: _LOGGER.debug( "%s.%s: Ignoring attempt to access Alexa after HA shutdown", func.__module__[func.__module__.find(".") + 1 :], func.__name__, ) return None except AlexapyLoginError as ex: login = None email = None all_args = list(args) + list(kwargs.values()) # _LOGGER.debug("Func %s instance %s %s %s", func, instance, args, kwargs) if instance: if hasattr(instance, "_login"): login = instance._login # pylint: disable=protected-access hass = instance.hass else: for arg in all_args: _LOGGER.debug("Checking %s", arg) if isinstance(arg, AlexaLogin): login = arg break if hasattr(arg, "_login"): login = instance._login hass = instance.hass break if login: # Try to re-login email = login.email if await login.test_loggedin(): _LOGGER.info( "%s.%s: Successful re-login after a login error for %s", func.__module__[func.__module__.find(".") + 1 :], func.__name__, hide_email(email), ) return None _LOGGER.debug( "%s.%s: detected bad login for %s: %s", func.__module__[func.__module__.find(".") + 1 :], func.__name__, hide_email(email), EXCEPTION_TEMPLATE.format(type(ex).__name__, ex.args), ) try: hass except NameError: hass = None report_relogin_required(hass, login, email) return None return result def report_relogin_required(hass, login, email) -> bool: """Send message for relogin required.""" if hass and login and email: if login.status: _LOGGER.debug( "Reporting need to relogin to %s with %s stats: %s", login.url, hide_email(email), login.stats, ) hass.bus.async_fire( "alexa_media_relogin_required", event_data={ "email": hide_email(email), "url": login.url, "stats": login.stats, }, ) return True return False def _existing_serials(hass, login_obj) -> list: """Retrieve existing serial numbers for a given login object.""" email: str = login_obj.email if ( DATA_ALEXAMEDIA in hass.data and "accounts" in hass.data[DATA_ALEXAMEDIA] and email in hass.data[DATA_ALEXAMEDIA]["accounts"] ): existing_serials = list( hass.data[DATA_ALEXAMEDIA]["accounts"][email]["entities"][ "media_player" ].keys() ) device_data = ( hass.data[DATA_ALEXAMEDIA]["accounts"][email] .get("devices", {}) .get("media_player", {}) ) for serial in existing_serials[:]: device = device_data.get(serial, {}) if "appDeviceList" in device and device["appDeviceList"]: apps = [ x["serialNumber"] for x in device["appDeviceList"] if "serialNumber" in x ] existing_serials.extend(apps) else: _LOGGER.warning( "No accounts data found for %s. Skipping serials retrieval.", email ) existing_serials = [] return existing_serials async def calculate_uuid(hass, email: str, url: str) -> dict: """Return uuid and index of email/url. Args hass (bool): Hass entity url (str): url for account email (str): email for account Returns dict: dictionary with uuid and index """ result = {} return_index = 0 if hass.config_entries.async_entries(DATA_ALEXAMEDIA): for index, entry in enumerate( hass.config_entries.async_entries(DATA_ALEXAMEDIA) ): if entry.data.get(CONF_EMAIL) == email and entry.data.get(CONF_URL) == url: return_index = index break uuid = await async_get_instance_id(hass) result["uuid"] = hex( int(uuid, 16) # increment uuid for second accounts + return_index # hash email/url in case HA uuid duplicated + int( hashlib.sha256((email.lower() + url.lower()).encode()).hexdigest(), 16, # nosec ) )[-32:] result["index"] = return_index _LOGGER.debug("%s: Returning uuid %s", hide_email(email), result) return result def alarm_just_dismissed( alarm: dict[str, Any], previous_status: Optional[str], previous_version: Optional[str], ) -> bool: """Given the previous state of an alarm, determine if it has just been dismissed.""" if ( previous_status not in ("SNOOZED", "ON") # The alarm had to be in a status that supported being dismissed or previous_version is None # The alarm was probably just created or not alarm # The alarm that was probably just deleted. or alarm.get("status") not in ("OFF", "ON") # A dismissed alarm is guaranteed to be turned off(one-off alarm) or left on(recurring alarm) or previous_version == alarm.get("version") # A dismissal always has a changed version. or int(alarm.get("version", "0")) > 1 + int(previous_version) ): # This is an absurd thing to check, but it solves many, many edge cases. # Experimentally, when an alarm is dismissed, the version always increases by 1 # When an alarm is edited either via app or voice, its version always increases by 2+ return False # It seems obvious that a check involving time should be necessary. It is not. # We know there was a change and that it wasn't an edit. # We also know the alarm's status rules out a snooze. # The only remaining possibility is that this alarm was just dismissed. return True def is_http2_enabled(hass: HomeAssistant | None, login_email: str) -> bool: """Whether HTTP2 push is enabled for the current account session""" if hass: return bool( safe_get( hass.data, [DATA_ALEXAMEDIA, "accounts", login_email, "http2"], ) ) return False @overload def safe_get( data: Any, path_list: list[str | int] | None = None, checknone: bool = False, ignorecase: bool = False, pathsep: str = ".", search: Any = None, pretty: bool = False, rtype: str | None = None, ) -> Any | None: ... @overload def safe_get( data: Any, path_list: list[str | int] | None, default: ArgType, *args, **kwargs ) -> ArgType: ... def safe_get( data: Any, path_list: list[str | int] | None = None, *args, **kwargs ) -> None | Any: """Safely get nested value using path segments with optional type checking. Args: data: Source data structure path_list: List of path segments (dots in segment names are auto-escaped) *args: Positional arguments passed to dictor (e.g., default value) **kwargs: Keyword arguments passed to dictor (checknone, ignorecase) Returns: The value at the specified path, or None if: - The path doesn't exist and no default is provided or default if: - A default is provided and the path doesn't exist - A default is provided and the retrieved value's type doesn't match the default's type Note: - Do not pass 'pathsep' in kwargs as the path is pre-built. - Type checking: When a default value is provided and a non-None value is retrieved, the result is validated against the default's type. If types don't match, default is returned. This prevents silent type errors from malformed data structures. Examples: >>> safe_get({"a": {"b": "value"}}, ["a", "b"]) 'value' >>> safe_get({"a": {"b": 123}}, ["a", "b"], "default") 'default' # Type mismatch: int vs str >>> safe_get({"a": {"b": "value"}}, ["a", "b"], "default") 'value' # Type matches """ if not path_list: raise ValueError("path_list cannot be empty") if "pathsep" in kwargs: kwargs.pop("pathsep") # Ignore pathsep since we build the path escaped_segments = (str(seg).replace(".", "\\.") for seg in path_list) path = ".".join(escaped_segments) default = args[0] if args else (kwargs.get("default") if kwargs else None) result = dictor(data, path, *args, **kwargs) if default is not None and result is not None: if not isinstance(result, type(default)): result = default return result