/
/
/
1"""Discovery core controller."""
2
3from __future__ import annotations
4
5import asyncio
6import contextlib
7import inspect
8import logging
9import os
10import re
11from collections import defaultdict
12from ipaddress import IPv4Address
13from typing import TYPE_CHECKING, Any
14
15from aiohttp import ClientTimeout
16from music_assistant_models.config_entries import ConfigEntry
17from music_assistant_models.enums import ConfigEntryType, EventType
18from zeroconf import (
19 NonUniqueNameException,
20 ServiceStateChange,
21 Zeroconf,
22)
23from zeroconf.asyncio import AsyncServiceBrowser, AsyncServiceInfo, AsyncZeroconf
24
25from music_assistant.constants import (
26 CONF_ENTRY_ZEROCONF_INTERFACES,
27 CONF_ZEROCONF_INTERFACES,
28 INGRESS_SERVER_PORT,
29 VERBOSE_LOG_LEVEL,
30)
31from music_assistant.helpers.util import get_ip_pton, get_zeroconf_args
32from music_assistant.models.core_controller import CoreController
33
34if TYPE_CHECKING:
35 from async_upnp_client.utils import CaseInsensitiveDict
36 from music_assistant_models.config_entries import CoreConfig
37 from music_assistant_models.event import MassEvent
38
39 from music_assistant.models import ProviderInstanceType
40
41# RAOP cache keys prefix the device name with the device MAC, e.g. "aabbccddeeff@Kelder".
42# Cache keys are lowercased, so the hex is matched in lowercase.
43RAOP_MAC_PREFIX = re.compile(r"^[0-9a-f]{12}@")
44
45CONF_UPNP_NETWORK_SCAN = "upnp_network_scan"
46UPNP_DISCOVERY_INTERVAL = 300
47UPNP_DISCOVERY_BROADCAST_TARGET = (str(IPv4Address("255.255.255.255")), 1900)
48UPNP_DISCOVERY_TASK_ID = "discovery_upnp_cycle"
49UPNP_DISCOVERY_TIMER_ID = "discovery_upnp_timer"
50
51# Re-announce daily so the HA integration token rotates before expiry, also on long uptimes
52HA_ANNOUNCE_INTERVAL = 86400
53HA_ANNOUNCE_TIMER_ID = "discovery_ha_announce_timer"
54
55
56async def async_upnp_search(*args: Any, **kwargs: Any) -> None:
57 """Run async_upnp_client SSDP search with lazy import."""
58 from async_upnp_client.search import async_search # noqa: PLC0415
59
60 await async_search(*args, **kwargs)
61
62
63class DiscoveryController(CoreController):
64 """Core controller that manages mDNS/Zeroconf and SSDP/UPnP discovery."""
65
66 domain = "discovery"
67
68 def __init__(self, *args: Any, **kwargs: Any) -> None:
69 """Initialize discovery controller."""
70 super().__init__(*args, **kwargs)
71 self.manifest.name = "Discovery"
72 self.manifest.description = (
73 "Handles mDNS/Zeroconf, SSDP/UPnP discovery and Music Assistant network broadcast."
74 )
75 self.manifest.icon = "radar"
76 self._aiozc: AsyncZeroconf | None = None
77 self._mdns_browser: AsyncServiceBrowser | None = None
78 self._mass_service_info: AsyncServiceInfo | None = None
79 self._mdns_locks: dict[str, asyncio.Lock] = {}
80 self._upnp_locks: dict[str, asyncio.Lock] = {}
81 self._upnp_run_lock = asyncio.Lock()
82 self._mdns_waiters: list[asyncio.Event] = []
83
84 @property
85 def aiozc(self) -> AsyncZeroconf:
86 """Return the shared AsyncZeroconf instance for discovery consumers."""
87 assert self._aiozc is not None, "DiscoveryController is not initialized"
88 return self._aiozc
89
90 async def setup(self, config: CoreConfig) -> None:
91 """Initialize discovery controller."""
92 self.config = config
93 if self._aiozc is None:
94 self._aiozc = self._create_aiozc(config)
95 self._configure_library_loggers()
96 await self._setup_mdns_browser()
97 await self._register_mass_service()
98 # the mdns record embeds the server info, so refresh it when that info
99 # changes (e.g. the server was renamed or its state/urls changed)
100 self.mass.subscribe(self._on_core_state_updated, EventType.CORE_STATE_UPDATED)
101 if self.mass.running_as_hass_addon:
102 # (re)announce to HA supervisor to make sure that HA picks it up
103 await self._announce_to_homeassistant()
104 self._schedule_periodic_ha_announce()
105 self._schedule_periodic_upnp_discovery()
106
107 async def get_config_entries(self) -> tuple[ConfigEntry, ...]:
108 """Return config entries for the discovery controller."""
109 return (
110 CONF_ENTRY_ZEROCONF_INTERFACES,
111 ConfigEntry(
112 key=CONF_UPNP_NETWORK_SCAN,
113 type=ConfigEntryType.BOOLEAN,
114 default_value=False,
115 requires_reload=False,
116 ),
117 )
118
119 async def close(self) -> None:
120 """Handle logic on server stop."""
121 self.mass.cancel_timer(UPNP_DISCOVERY_TIMER_ID)
122 self.mass.cancel_task(UPNP_DISCOVERY_TASK_ID)
123 self.mass.cancel_timer(HA_ANNOUNCE_TIMER_ID)
124
125 await self._cancel_mdns_browser()
126
127 aiozc = self._aiozc
128 if self._mass_service_info:
129 with contextlib.suppress(Exception):
130 assert aiozc is not None
131 await aiozc.async_unregister_service(self._mass_service_info)
132 self._mass_service_info = None
133
134 self._mdns_locks.clear()
135 self._upnp_locks.clear()
136 if aiozc is not None:
137 with contextlib.suppress(Exception):
138 await aiozc.async_close()
139 self._aiozc = None
140
141 async def run_provider_discovery(self, provider: ProviderInstanceType) -> None:
142 """Run discovery for a specific provider."""
143 if provider.manifest.mdns_discovery:
144 await self._replay_mdns_discovery(provider)
145 if provider.manifest.upnp_discovery:
146 await self._run_upnp_discovery_cycle(set(provider.manifest.upnp_discovery))
147 self._schedule_periodic_upnp_discovery()
148
149 def on_provider_unload(self, instance_id: str) -> None:
150 """Clean up provider-specific discovery state."""
151 self._mdns_locks.pop(instance_id, None)
152 self._upnp_locks.pop(instance_id, None)
153 self._schedule_periodic_upnp_discovery()
154
155 async def async_find_mdns_service(
156 self, service_type: str, name_filter: str, timeout: float = 3.0
157 ) -> AsyncServiceInfo | None:
158 """
159 Find an mDNS service by exact device name match, checking cache first then waiting.
160
161 :param service_type: The mDNS service type (e.g., "_raop._tcp.local.").
162 :param name_filter: Device name that must exactly match the service name portion.
163 :param timeout: Maximum time to wait in seconds.
164 """
165 deadline = asyncio.get_event_loop().time() + timeout
166 # Cache keys are lowercased DNS names, so we must match case-insensitively
167 name_filter_lower = name_filter.lower()
168 service_type_lower = service_type.lower()
169 event = asyncio.Event()
170 self._mdns_waiters.append(event)
171 try:
172 while True:
173 # Clear before scanning so events arriving during the scan are not lost
174 event.clear()
175 # Check cache for a matching entry
176 for mdns_name in set(self.aiozc.zeroconf.cache.cache):
177 if service_type_lower not in mdns_name or mdns_name == service_type_lower:
178 continue
179 # Use exact matching on the device name portion to prevent a device named
180 # "Foo" from cross-matching another device named "ATV Foo".
181 # mDNS names are either "[email protected]." or "DeviceName.service.local."
182 # Strip the MAC prefix only when present, so device names that legitimately
183 # contain "@" are not truncated.
184 device_part = mdns_name.split(".")[0]
185 device_name = RAOP_MAC_PREFIX.sub("", device_part, count=1)
186 if device_name != name_filter_lower:
187 continue
188 info = AsyncServiceInfo(service_type, mdns_name)
189 if await info.async_request(self.aiozc.zeroconf, 3000):
190 return info
191 remaining = deadline - asyncio.get_event_loop().time()
192 if remaining <= 0:
193 return None
194 # Wait for the next mDNS state change event, then re-check the cache
195 try:
196 await asyncio.wait_for(event.wait(), timeout=remaining)
197 except TimeoutError:
198 return None
199 finally:
200 self._mdns_waiters.remove(event)
201
202 def _configure_library_loggers(self) -> None:
203 """Align third-party discovery logging with the discovery controller log level."""
204 library_log_level = (
205 logging.DEBUG if self.logger.isEnabledFor(VERBOSE_LOG_LEVEL) else self.logger.level + 10
206 )
207 for logger_name in ("async_upnp_client", "zeroconf"):
208 logging.getLogger(logger_name).setLevel(library_log_level)
209
210 def _create_aiozc(self, config: CoreConfig) -> AsyncZeroconf:
211 """Create the shared AsyncZeroconf instance for the discovery controller."""
212 zeroconf_interfaces = str(config.get_value(CONF_ZEROCONF_INTERFACES, "default"))
213 use_all_interfaces = zeroconf_interfaces == "all"
214 zc_args = get_zeroconf_args(use_all_interfaces)
215 self.logger.debug("Zeroconf configuration: %s", zc_args)
216 return AsyncZeroconf(
217 ip_version=zc_args["ip_version"],
218 interfaces=zc_args["interfaces"],
219 )
220
221 async def _setup_mdns_browser(self) -> None:
222 """Create the global mDNS browser for all subscribed provider types."""
223 await self._cancel_mdns_browser()
224
225 all_types: set[str] = set()
226 for manifest in self.mass.get_provider_manifests():
227 if manifest.mdns_discovery:
228 all_types.update(manifest.mdns_discovery)
229
230 if not all_types:
231 return
232
233 self._mdns_browser = AsyncServiceBrowser(
234 self.aiozc.zeroconf,
235 list(all_types),
236 handlers=[self._on_mdns_service_state_change],
237 )
238
239 async def _cancel_mdns_browser(self) -> None:
240 """Cancel the active mDNS browser if one is running."""
241 if not self._mdns_browser:
242 return
243
244 async_cancel = getattr(self._mdns_browser, "async_cancel", None)
245 cancel = getattr(self._mdns_browser, "cancel", None)
246 if callable(async_cancel):
247 result = async_cancel()
248 if inspect.isawaitable(result):
249 await result
250 elif callable(cancel):
251 cancel()
252 elif callable(cancel):
253 cancel()
254 self._mdns_browser = None
255
256 async def _register_mass_service(self) -> None:
257 """Register the Music Assistant server on the network via Zeroconf."""
258 zeroconf_type = "_mass._tcp.local."
259 server_id = self.mass.server_id
260 self.logger.debug("Starting Zeroconf broadcast...")
261 info = AsyncServiceInfo(
262 zeroconf_type,
263 name=f"{server_id}.{zeroconf_type}",
264 addresses=[
265 await get_ip_pton(address) for address in self.mass.webserver.publish_addresses
266 ],
267 port=self.mass.webserver.publish_port,
268 properties=self.mass.get_server_info().to_dict(),
269 server="mass.local.",
270 )
271 try:
272 if self._mass_service_info:
273 await self.aiozc.async_update_service(info)
274 else:
275 await self.aiozc.async_register_service(info)
276 self._mass_service_info = info
277 except NonUniqueNameException:
278 self.logger.error(
279 "Music Assistant instance with identical name present in the local network!"
280 )
281
282 async def _on_core_state_updated(self, event: MassEvent) -> None:
283 """Refresh the advertised mdns record after the server info changed."""
284 if self._aiozc is None or self.mass.closing:
285 return
286 await self._register_mass_service()
287
288 def _on_mdns_service_state_change(
289 self,
290 zeroconf: Zeroconf,
291 service_type: str,
292 name: str,
293 state_change: ServiceStateChange,
294 ) -> None:
295 """Handle mDNS service state callbacks."""
296
297 async def process_mdns_state_change(provider: ProviderInstanceType) -> None:
298 lock = self._mdns_locks.setdefault(provider.instance_id, asyncio.Lock())
299 if state_change == ServiceStateChange.Removed:
300 info = None
301 else:
302 info = AsyncServiceInfo(service_type, name)
303 await info.async_request(zeroconf, 3000)
304 async with lock:
305 await provider.on_mdns_service_state_change(name, state_change, info)
306
307 self.logger.log(
308 VERBOSE_LOG_LEVEL,
309 "Service %s of type %s state changed: %s",
310 name,
311 service_type,
312 state_change,
313 )
314 # Notify any waiters that a new mDNS event arrived
315 for waiter in self._mdns_waiters:
316 waiter.set()
317 for provider in list(self.mass.providers):
318 if not provider.available or not provider.manifest.mdns_discovery:
319 continue
320 if service_type in provider.manifest.mdns_discovery:
321 self.mass.create_task(process_mdns_state_change(provider))
322
323 async def _replay_mdns_discovery(self, provider: ProviderInstanceType) -> None:
324 """Replay cached mDNS results for a provider after it loads."""
325 lock = self._mdns_locks.setdefault(provider.instance_id, asyncio.Lock())
326 async with lock:
327 for mdns_type in provider.manifest.mdns_discovery or []:
328 for mdns_name in set(self.aiozc.zeroconf.cache.cache):
329 if mdns_type not in mdns_name or mdns_type == mdns_name:
330 continue
331 info = AsyncServiceInfo(mdns_type, mdns_name)
332 if await info.async_request(self.aiozc.zeroconf, 3000):
333 await provider.on_mdns_service_state_change(
334 mdns_name, ServiceStateChange.Added, info
335 )
336
337 def _schedule_periodic_upnp_discovery(self) -> None:
338 """Ensure the periodic SSDP discovery cycle matches active subscriptions."""
339 subscriptions, _ = self._get_upnp_subscriptions()
340 if not subscriptions:
341 self.mass.cancel_timer(UPNP_DISCOVERY_TIMER_ID)
342 self.mass.cancel_task(UPNP_DISCOVERY_TASK_ID)
343 return
344
345 def run_discovery() -> None:
346 self.mass.create_task(
347 self._run_upnp_discovery_cycle(),
348 task_id=UPNP_DISCOVERY_TASK_ID,
349 )
350
351 self.mass.call_later(
352 UPNP_DISCOVERY_INTERVAL,
353 run_discovery,
354 task_id=UPNP_DISCOVERY_TIMER_ID,
355 )
356
357 def _get_upnp_subscriptions(
358 self,
359 ) -> tuple[dict[str, list[ProviderInstanceType]], set[str]]:
360 """Return active SSDP subscriptions keyed by search target."""
361 subscriptions: dict[str, list[ProviderInstanceType]] = defaultdict(list)
362 broadcast_targets: set[str] = set()
363 allow_network_scan = bool(self.config.get_value(CONF_UPNP_NETWORK_SCAN))
364 for provider in list(self.mass.providers):
365 if not provider.available or not provider.manifest.upnp_discovery:
366 continue
367 for search_target in provider.manifest.upnp_discovery:
368 subscriptions[search_target].append(provider)
369 if allow_network_scan:
370 broadcast_targets.add(search_target)
371 return subscriptions, broadcast_targets
372
373 async def _run_upnp_discovery_cycle(self, search_targets: set[str] | None = None) -> None:
374 """Run one SSDP discovery cycle for all active subscriptions."""
375 try:
376 async with self._upnp_run_lock:
377 subscriptions, broadcast_targets = self._get_upnp_subscriptions()
378 if search_targets is not None:
379 subscriptions = {
380 search_target: providers
381 for search_target, providers in subscriptions.items()
382 if search_target in search_targets
383 }
384 broadcast_targets.intersection_update(search_targets)
385 if not subscriptions:
386 return
387
388 for search_target, providers in subscriptions.items():
389 seen_results: set[tuple[str | None, str | None, str | None]] = set()
390 await self._run_upnp_search(search_target, providers, seen_results)
391 if search_target in broadcast_targets:
392 await self._run_upnp_search(
393 search_target,
394 providers,
395 seen_results,
396 target=UPNP_DISCOVERY_BROADCAST_TARGET,
397 )
398 finally:
399 self._schedule_periodic_upnp_discovery()
400
401 async def _run_upnp_search(
402 self,
403 search_target: str,
404 providers: list[ProviderInstanceType],
405 seen_results: set[tuple[str | None, str | None, str | None]],
406 target: tuple[str, int] | None = None,
407 ) -> None:
408 """Run one SSDP search and dispatch matching responses to subscribed providers."""
409
410 async def on_response(discovery_info: CaseInsensitiveDict) -> None:
411 dedupe_key = (
412 discovery_info.get("usn"),
413 discovery_info.get("location"),
414 discovery_info.get("_host"),
415 )
416 if dedupe_key in seen_results:
417 return
418 seen_results.add(dedupe_key)
419 for provider in providers:
420 if not provider.available:
421 continue
422 await self._dispatch_upnp_discovery(provider, search_target, discovery_info)
423
424 try:
425 if target is None:
426 await async_upnp_search(on_response, search_target=search_target)
427 else:
428 await async_upnp_search(on_response, search_target=search_target, target=target)
429 except OSError as err:
430 target_label = f" via {target[0]}" if target else ""
431 self.logger.warning(
432 "UPnP discovery for %s failed%s: %s",
433 search_target,
434 target_label,
435 err,
436 )
437
438 async def _dispatch_upnp_discovery(
439 self,
440 provider: ProviderInstanceType,
441 search_target: str,
442 discovery_info: CaseInsensitiveDict,
443 ) -> None:
444 """Dispatch one SSDP discovery result to a provider callback."""
445 lock = self._upnp_locks.setdefault(provider.instance_id, asyncio.Lock())
446 async with lock:
447 try:
448 await provider.on_upnp_service_discovered(search_target, discovery_info)
449 except Exception as err:
450 self.logger.warning(
451 "Error handling UPnP discovery for %s: %s",
452 provider.name,
453 err,
454 exc_info=err if self.logger.isEnabledFor(logging.DEBUG) else None,
455 )
456
457 async def _announce_to_homeassistant(self) -> None:
458 """Announce Music Assistant Ingress server to Home Assistant via Supervisor API."""
459 supervisor_token = os.environ["SUPERVISOR_TOKEN"]
460 addon_hostname = os.environ["HOSTNAME"]
461 ha_integration_token = await self.mass.webserver.auth.get_homeassistant_system_user_token()
462 discovery_payload = {
463 "service": "music_assistant",
464 "config": {
465 "host": addon_hostname,
466 "port": INGRESS_SERVER_PORT,
467 "auth_token": ha_integration_token,
468 },
469 }
470 try:
471 async with self.mass.http_session_no_ssl.post(
472 "http://supervisor/discovery",
473 headers={"Authorization": f"Bearer {supervisor_token}"},
474 json=discovery_payload,
475 timeout=ClientTimeout(total=10),
476 ) as response:
477 response.raise_for_status()
478 result = await response.json()
479 self.logger.debug(
480 "Successfully announced to Home Assistant. Discovery UUID: %s",
481 result.get("uuid"),
482 )
483 except Exception as err:
484 self.logger.warning("Failed to announce to Home Assistant: %s", err)
485
486 def _schedule_periodic_ha_announce(self) -> None:
487 """Schedule the periodic (re)announce to Home Assistant."""
488
489 def run_announce() -> None:
490 self.mass.create_task(self._announce_to_homeassistant())
491 self._schedule_periodic_ha_announce()
492
493 self.mass.call_later(
494 HA_ANNOUNCE_INTERVAL,
495 run_announce,
496 task_id=HA_ANNOUNCE_TIMER_ID,
497 )
498