diff --git a/sandbox/hass_client/hass_client/approved_domains.py b/sandbox/hass_client/hass_client/approved_domains.py index 6d480dc42131..ddd01c1fcd7c 100644 --- a/sandbox/hass_client/hass_client/approved_domains.py +++ b/sandbox/hass_client/hass_client/approved_domains.py @@ -26,11 +26,14 @@ Domain comparison is case-insensitive; everything is normalised to lowercase at insertion time so the lookups stay cheap. """ -from collections.abc import Iterable +from collections.abc import Callable, Iterable +import contextlib import logging _LOGGER = logging.getLogger(__name__) +ApproveListener = Callable[[str], None] + class ApprovedDomains: """Mutable set of domains the sandbox runtime is allowed to own.""" @@ -38,14 +41,38 @@ class ApprovedDomains: def __init__(self, initial: Iterable[str] | None = None) -> None: """Initialise the gate, optionally seeded with a starter set.""" self._counts: dict[str, int] = {} + self._listeners: list[ApproveListener] = [] if initial is not None: for domain in initial: self.add(domain) + def add_listener(self, listener: ApproveListener) -> None: + """Subscribe ``listener`` to absent→present domain transitions. + + The listener fires (with the lowercased domain) only the first time + a domain becomes approved, not on every refcount bump. Used by the + :class:`~hass_client.service_mirror.ServiceMirror` to replay services + a domain registered before its approval landed. + """ + self._listeners.append(listener) + + def remove_listener(self, listener: ApproveListener) -> None: + """Drop a previously-added approve listener (no-op if absent).""" + with contextlib.suppress(ValueError): + self._listeners.remove(listener) + def add(self, domain: str) -> None: - """Approve ``domain``; multiple ``add`` calls bump a refcount.""" + """Approve ``domain``; multiple ``add`` calls bump a refcount. + + The first ``add`` for a domain (absent→present) notifies any + registered approve listeners. + """ key = domain.lower() + is_new = key not in self._counts self._counts[key] = self._counts.get(key, 0) + 1 + if is_new: + for listener in self._listeners: + listener(key) def remove(self, domain: str) -> None: """Drop one ``add`` for ``domain``; harmless when over-removed.""" diff --git a/sandbox/hass_client/hass_client/service_mirror.py b/sandbox/hass_client/hass_client/service_mirror.py index 52fe169cb981..f191f63862aa 100644 --- a/sandbox/hass_client/hass_client/service_mirror.py +++ b/sandbox/hass_client/hass_client/service_mirror.py @@ -64,9 +64,16 @@ class ServiceMirror: self._unsub_removed = self.hass.bus.async_listen( EVENT_SERVICE_REMOVED, self._on_service_removed ) + # Replay services a domain registered *before* it became approved: + # EVENT_SERVICE_REGISTERED fires synchronously during + # ``async_setup_entry``, but the entry runner only approves the + # domain once setup returns, so those early registrations were + # dropped. Re-mirror them the moment the domain is approved. + self.approved.add_listener(self.async_sync_domain) async def async_stop(self) -> None: """Detach the bus listeners.""" + self.approved.remove_listener(self.async_sync_domain) if self._unsub_registered is not None: self._unsub_registered() self._unsub_registered = None @@ -74,6 +81,21 @@ class ServiceMirror: self._unsub_removed() self._unsub_removed = None + @callback + def async_sync_domain(self, domain: str) -> None: + """Mirror every already-registered service of a freshly-approved domain. + + Invoked as an :class:`ApprovedDomains` approve listener. Services + already mirrored (``_mirrored``) are skipped, so this is a no-op for + a domain that was approved before its services registered. + """ + if self._channel is None or self._channel.closed: + return + if not self.approved.approves(domain): + return + for service in self.hass.services.async_services_for_domain(domain): + self._mirror_service(domain, service) + @callback def _on_service_registered(self, event: Event) -> None: if self._channel is None or self._channel.closed: @@ -89,6 +111,15 @@ class ServiceMirror: sorted(self.approved.domains), ) return + self._mirror_service(domain, service) + + @callback + def _mirror_service(self, domain: str, service: str) -> None: + """Push one ``register_service`` for ``domain.service`` (once). + + Shared by the live ``EVENT_SERVICE_REGISTERED`` path and the + approval-replay path. Caller guarantees the domain is approved. + """ key = (domain.lower(), service.lower()) if key in self._mirrored: return diff --git a/sandbox/hass_client/tests/test_service_mirror.py b/sandbox/hass_client/tests/test_service_mirror.py index 1b56050ba8da..633ee7845695 100644 --- a/sandbox/hass_client/tests/test_service_mirror.py +++ b/sandbox/hass_client/tests/test_service_mirror.py @@ -153,6 +153,60 @@ async def test_unapproved_domain_is_rejected( await mirror.async_stop() +async def test_service_registered_before_approval_is_replayed( + channels: tuple[Channel, Channel], hass_runtime: Any +) -> None: + """A service registered before its domain is approved is replayed on approval. + + Mirrors the real race: ``EVENT_SERVICE_REGISTERED`` fires synchronously + while a service is registered inside ``async_setup_entry``, before the + entry runner approves the domain. The early registration is dropped at + fire time, then replayed the moment ``ApprovedDomains.add`` approves it. + """ + main, sandbox = channels + register_calls: list[pb.RegisterService] = [] + + async def _on_register(msg: pb.RegisterService) -> pb.RegisterServiceResult: + register_calls.append(msg) + return pb.RegisterServiceResult(ok=True, installed=True) + + main.register("sandbox/register_service", _on_register) + main.start() + sandbox.start() + + # Start with an empty gate — the domain is NOT yet approved. + approved = ApprovedDomains() + mirror = ServiceMirror(hass_runtime, approved) + mirror.register(sandbox) + + async def _svc(_call: Any) -> None: + return None + + # Service registers while the domain is still unapproved (the setup race). + hass_runtime.services.async_register( + "phase1_demo", + "do_thing", + _svc, + supports_response=SupportsResponse.NONE, + ) + # Give the (dropped) registration a few ticks; nothing should reach main. + for _ in range(20): + await asyncio.sleep(0) + assert register_calls == [] + + # Domain becomes approved — exactly what the entry runner does after + # ``async_setup`` returns — and the dropped service is now replayed. + approved.add("phase1_demo") + + await _wait_until(lambda: bool(register_calls)) + + assert len(register_calls) == 1 + assert register_calls[0].domain == "phase1_demo" + assert register_calls[0].service == "do_thing" + + await mirror.async_stop() + + async def test_unregister_service_propagates( channels: tuple[Channel, Channel], hass_runtime: Any ) -> None: