Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
334 changes: 211 additions & 123 deletions homeassistant/helpers/event.py
Original file line number Diff line number Diff line change
Expand Up @@ -61,13 +61,27 @@
TRACK_ENTITY_REGISTRY_UPDATED_CALLBACKS = "track_entity_registry_updated_callbacks"
TRACK_ENTITY_REGISTRY_UPDATED_LISTENER = "track_entity_registry_updated_listener"

_TEMPLATE_ALL_LISTENER = "all"
_TEMPLATE_DOMAINS_LISTENER = "domains"
_TEMPLATE_ENTITIES_LISTENER = "entities"
_ALL_LISTENER = "all"
_DOMAINS_LISTENER = "domains"
_ENTITIES_LISTENER = "entities"

_LOGGER = logging.getLogger(__name__)


@dataclass
class TrackStates:
"""Class for keeping track of states being tracked.

all_states: All states on the system are being tracked
entities: Entities to track
domains: Domains to track
"""

all_states: bool
entities: Set
domains: Set


@dataclass
class TrackTemplate:
"""Class for keeping track of a template with variables.
Expand Down Expand Up @@ -452,6 +466,158 @@ def _async_string_to_lower_list(instr: Union[str, Iterable[str]]) -> List[str]:
return [mstr.lower() for mstr in instr]


class _TrackStateChangeFiltered:
"""Handle removal / refresh of tracker."""

def __init__(
self,
hass: HomeAssistant,
track_states: TrackStates,
action: Callable[[Event], Any],
):
"""Handle removal / refresh of tracker init."""
self.hass = hass
self._action = action
self._listeners: Dict[str, Callable] = {}
self._last_track_states: TrackStates = track_states

@callback
def async_setup(self) -> None:
"""Create listeners to track states."""
track_states = self._last_track_states

if (
not track_states.all_states
and not track_states.domains
and not track_states.entities
):
return

if track_states.all_states:
self._setup_all_listener()
return

self._setup_domains_listener(track_states.domains)
self._setup_entities_listener(track_states.domains, track_states.entities)

@property
def listeners(self) -> Dict:
"""State changes that will cause a re-render."""
track_states = self._last_track_states
return {
_ALL_LISTENER: track_states.all_states,
_ENTITIES_LISTENER: track_states.entities,
_DOMAINS_LISTENER: track_states.domains,
}

@callback
def async_update_listeners(self, new_track_states: TrackStates) -> None:
"""Update the listeners based on the new TrackStates."""
last_track_states = self._last_track_states
self._last_track_states = new_track_states

had_all_listener = last_track_states.all_states

if new_track_states.all_states:
if had_all_listener:
return
self._cancel_listener(_DOMAINS_LISTENER)
self._cancel_listener(_ENTITIES_LISTENER)
self._setup_all_listener()
return

if had_all_listener:
self._cancel_listener(_ALL_LISTENER)

domains_changed = new_track_states.domains != last_track_states.domains

if had_all_listener or domains_changed:
domains_changed = True
self._cancel_listener(_DOMAINS_LISTENER)
self._setup_domains_listener(new_track_states.domains)

if (
had_all_listener
or domains_changed
or new_track_states.entities != last_track_states.entities
):
self._cancel_listener(_ENTITIES_LISTENER)
self._setup_entities_listener(
new_track_states.domains, new_track_states.entities
)

@callback
def async_remove(self) -> None:
"""Cancel the listeners."""
for key in list(self._listeners):
self._listeners.pop(key)()

@callback
def _cancel_listener(self, listener_name: str) -> None:
if listener_name not in self._listeners:
return

self._listeners.pop(listener_name)()

@callback
def _setup_entities_listener(self, domains: Set, entities: Set) -> None:
if domains:
entities = entities.copy()
entities.update(self.hass.states.async_entity_ids(domains))

# Entities has changed to none
if not entities:
return

self._listeners[_ENTITIES_LISTENER] = async_track_state_change_event(
self.hass, entities, self._action
)

@callback
def _setup_domains_listener(self, domains: Set) -> None:
if not domains:
return

self._listeners[_DOMAINS_LISTENER] = async_track_state_added_domain(
self.hass, domains, self._action
)

@callback
def _setup_all_listener(self) -> None:
self._listeners[_ALL_LISTENER] = self.hass.bus.async_listen(
EVENT_STATE_CHANGED, self._action
)


@callback
@bind_hass
def async_track_state_change_filtered(
hass: HomeAssistant,
track_states: TrackStates,
action: Callable[[Event], Any],
) -> _TrackStateChangeFiltered:
"""Track state changes with a TrackStates filter that can be updated.

Parameters
----------
hass
Home assistant object.
track_states
A TrackStates data class.
action
Callable to call with results.

Returns
-------
Object used to update the listeners (async_update_listeners) with a new TrackStates or
cancel the tracking (async_remove).

"""
tracker = _TrackStateChangeFiltered(hass, track_states, action)
tracker.async_setup()
return tracker


@callback
@bind_hass
def async_track_template(
Expand Down Expand Up @@ -557,12 +723,9 @@ def __init__(
track_template_.template.hass = hass
self._track_templates = track_templates

self._listeners: Dict[str, Callable] = {}

self._last_result: Dict[Template, Union[str, TemplateError]] = {}
self._info: Dict[Template, RenderInfo] = {}
self._last_domains: Set = set()
self._last_entities: Set = set()
self._track_state_changes: Optional[_TrackStateChangeFiltered] = None

def async_setup(self, raise_on_template_error: bool) -> None:
"""Activation of template tracking."""
Expand All @@ -580,7 +743,9 @@ def async_setup(self, raise_on_template_error: bool) -> None:
exc_info=self._info[template].exception,
)

self._create_listeners()
self._track_state_changes = async_track_state_change_filtered(
self.hass, _render_infos_to_track_states(self._info.values()), self._refresh
)
_LOGGER.debug(
"Template group %s listens for %s",
self._track_templates,
Expand All @@ -590,123 +755,14 @@ def async_setup(self, raise_on_template_error: bool) -> None:
@property
def listeners(self) -> Dict:
"""State changes that will cause a re-render."""
return {
"all": _TEMPLATE_ALL_LISTENER in self._listeners,
"entities": self._last_entities,
"domains": self._last_domains,
}

@property
def _needs_all_listener(self) -> bool:
for info in self._info.values():
# Tracking all states
if info.all_states or info.all_states_lifecycle:
return True

# Previous call had an exception
# so we do not know which states
# to track
if info.exception:
return True

return False

@property
def _all_templates_are_static(self) -> bool:
for info in self._info.values():
if not info.is_static:
return False

return True

@callback
def _create_listeners(self) -> None:
if self._all_templates_are_static:
return

if self._needs_all_listener:
self._setup_all_listener()
return

self._last_entities, self._last_domains = _entities_domains_from_info(
self._info.values()
)
self._setup_domains_listener(self._last_domains)
self._setup_entities_listener(self._last_domains, self._last_entities)

@callback
def _cancel_listener(self, listener_name: str) -> None:
if listener_name not in self._listeners:
return

self._listeners.pop(listener_name)()

@callback
def _update_listeners(self) -> None:
had_all_listener = _TEMPLATE_ALL_LISTENER in self._listeners

if self._needs_all_listener:
if had_all_listener:
return
self._last_domains = set()
self._last_entities = set()
self._cancel_listener(_TEMPLATE_DOMAINS_LISTENER)
self._cancel_listener(_TEMPLATE_ENTITIES_LISTENER)
self._setup_all_listener()
return

if had_all_listener:
self._cancel_listener(_TEMPLATE_ALL_LISTENER)

entities, domains = _entities_domains_from_info(self._info.values())
domains_changed = domains != self._last_domains

if had_all_listener or domains_changed:
domains_changed = True
self._cancel_listener(_TEMPLATE_DOMAINS_LISTENER)
self._setup_domains_listener(domains)

if had_all_listener or domains_changed or entities != self._last_entities:
self._cancel_listener(_TEMPLATE_ENTITIES_LISTENER)
self._setup_entities_listener(domains, entities)

self._last_domains = domains
self._last_entities = entities

@callback
def _setup_entities_listener(self, domains: Set, entities: Set) -> None:
if domains:
entities = entities.copy()
entities.update(self.hass.states.async_entity_ids(domains))

# Entities has changed to none
if not entities:
return

self._listeners[_TEMPLATE_ENTITIES_LISTENER] = async_track_state_change_event(
self.hass, entities, self._refresh
)

@callback
def _setup_domains_listener(self, domains: Set) -> None:
if not domains:
return

self._listeners[_TEMPLATE_DOMAINS_LISTENER] = async_track_state_added_domain(
self.hass, domains, self._refresh
)

@callback
def _setup_all_listener(self) -> None:
self._listeners[_TEMPLATE_ALL_LISTENER] = self.hass.bus.async_listen(
EVENT_STATE_CHANGED, self._refresh
)
assert self._track_state_changes
return self._track_state_changes.listeners

@callback
def async_remove(self) -> None:
"""Cancel the listener."""
for key in list(self._listeners):
self._listeners.pop(key)()
assert self._track_state_changes
self._track_state_changes.async_remove()

@callback
def async_refresh(self) -> None:
Expand Down Expand Up @@ -765,7 +821,10 @@ def _refresh(self, event: Optional[Event]) -> None:
updates.append(TrackTemplateResult(template, last_result, result))

if info_changed:
self._update_listeners()
assert self._track_state_changes
self._track_state_changes.async_update_listeners(
_render_infos_to_track_states(self._info.values()),
)
_LOGGER.debug(
"Template group %s listens for %s",
self._track_templates,
Expand Down Expand Up @@ -1229,7 +1288,10 @@ def process_state_match(
return lambda state: state in parameter_set


def _entities_domains_from_info(render_infos: Iterable[RenderInfo]) -> Tuple[Set, Set]:
@callback
def _entities_domains_from_render_infos(
render_infos: Iterable[RenderInfo],
) -> Tuple[Set, Set]:
"""Combine from multiple RenderInfo."""
entities = set()
domains = set()
Expand All @@ -1242,3 +1304,29 @@ def _entities_domains_from_info(render_infos: Iterable[RenderInfo]) -> Tuple[Set
if render_info.domains_lifecycle:
domains.update(render_info.domains_lifecycle)
return entities, domains


@callback
def _render_infos_needs_all_listener(render_infos: Iterable[RenderInfo]) -> bool:
"""Determine if an all listener is needed from RenderInfo."""
for render_info in render_infos:
# Tracking all states
if render_info.all_states or render_info.all_states_lifecycle:
return True

# Previous call had an exception
# so we do not know which states
# to track
if render_info.exception:
return True

return False


@callback
def _render_infos_to_track_states(render_infos: Iterable[RenderInfo]) -> TrackStates:
"""Create a TrackStates dataclass from the latest RenderInfo."""
if _render_infos_needs_all_listener(render_infos):
return TrackStates(True, set(), set())

return TrackStates(False, *_entities_domains_from_render_infos(render_infos))
Loading