/
/
1"""Just-in-time clip rendering for AI Radio."""
2# mypy: disable-error-code="attr-defined"
3
4from __future__ import annotations
5
6import asyncio
7import json
8import logging
9from dataclasses import dataclass, replace
10from typing import TYPE_CHECKING, Any, cast
11
12from music_assistant_models.enums import ContentType, StreamType, VolumeNormalizationMode
13from music_assistant_models.errors import (
14 InvalidDataError,
15 MediaNotFoundError,
16 MusicAssistantError,
17)
18from music_assistant_models.media_items import AudioFormat
19from music_assistant_models.streamdetails import StreamDetails
20
21from music_assistant.constants import (
22 CONF_VALUE_DISABLED,
23 CONF_VALUE_ENABLED,
24 CONF_VOLUME_NORMALIZATION,
25 CONF_VOLUME_NORMALIZATION_TARGET,
26 CONF_VOLUME_NORMALIZATION_TRACKS,
27)
28from music_assistant.helpers.audio import parse_loudnorm
29from music_assistant.helpers.ffmpeg import get_ffmpeg_stream
30from music_assistant.helpers.process import check_output
31from music_assistant.helpers.tags import async_parse_tags
32from music_assistant.helpers.tts import (
33 query_tts_engine_with_language_fallback,
34 resolve_tts_language,
35 resolve_tts_stream_path,
36)
37
38from .constants import (
39 ATTR_HOST_ID,
40 ATTR_MAX_CHARS,
41 ATTR_PROMPT,
42 ATTR_RENDERED_TEXT,
43 ATTR_SESSION_ID,
44 ATTR_WEB_SEARCH_MODE,
45 CLIP_STREAMDETAILS_EXPIRATION,
46 CONF_TTS_LOUDNESS_BOOST,
47 DEFAULT_TTS_LOUDNESS_BOOST,
48 DEFERRED_PLACEHOLDERS,
49 LOUDNESS_MEASURE_TIMEOUT,
50 MIN_CLIP_MEDIA_LIFETIME,
51 MIN_LOUDNESS_REFERENCE_SECONDS,
52 TTS_CLIP_PCM_FORMAT,
53 TTS_PEAK_CEILING_DB,
54 TTS_SERVER_ERROR_MARKERS,
55 TTS_SPEECHNORM_FILTER,
56)
57from .helpers import coerce_int, soft_limit_text
58
59if TYPE_CHECKING:
60 from collections.abc import AsyncGenerator
61
62 from music_assistant_models.config_entries import ProviderConfig
63 from music_assistant_models.enums import MediaType
64 from music_assistant_models.queue_item import QueueItem
65
66 from music_assistant.mass import MusicAssistant
67
68 from .models import SessionState
69
70
71@dataclass(slots=True)
72class _CachedClipMedia:
73 """Media previously minted for a clip, kept until it expires."""
74
75 path: str
76 stream_type: StreamType
77 audio_format: AudioFormat
78 duration: int | None
79 minted_at: float
80 loudness: float | None
81
82
83@dataclass(slots=True)
84class _ClipAudio:
85 """What get_audio_stream needs to serve a levelled clip, carried on StreamDetails.data."""
86
87 path: str
88 input_format: AudioFormat
89 gain_db: float
90
91
92class AIRadioRenderMixin:
93 """Renders an AI Radio clip at the moment MA needs its audio."""
94
95 if TYPE_CHECKING:
96 mass: MusicAssistant
97 config: ProviderConfig
98 logger: logging.Logger
99 _hosts: dict[str, dict[str, Any]]
100 _sessions: dict[str, SessionState]
101
102 _render_locks: dict[str, asyncio.Lock]
103 _media_cache: dict[str, _CachedClipMedia]
104 _engine_loudness: dict[tuple[str, str, str], float]
105
106 async def get_stream_details(self, item_id: str, media_type: MediaType) -> StreamDetails:
107 """
108 Render the AI Radio clip with the given id and return its StreamDetails.
109
110 :param item_id: The clip id of the queue item MA wants to play.
111 :param media_type: The media type of the requested item.
112 """
113 queue_item = self._find_clip_item(item_id)
114 if queue_item is None:
115 raise MediaNotFoundError(f"AI Radio clip {item_id} is not in any queue")
116 prompt = str(queue_item.extra_attributes.get(ATTR_PROMPT) or "")
117 if not prompt:
118 self._record_skip(queue_item, "clip has no prompt to render")
119 raise MediaNotFoundError(f"AI Radio clip {item_id} has no prompt to render")
120
121 async with self._lock_for(item_id):
122 text = str(queue_item.extra_attributes.get(ATTR_RENDERED_TEXT) or "")
123 if not text:
124 text = await self._generate_script(queue_item, prompt, item_id)
125 queue_item.extra_attributes[ATTR_RENDERED_TEXT] = text
126 # the signal is what marks the items cache dirty and schedules the persist
127 self.mass.player_queues.signal_update(queue_item.queue_id, items_changed=True)
128 media = await self._cached_clip_media(queue_item, text, item_id)
129
130 streamdetails = StreamDetails(
131 provider=self.instance_id,
132 item_id=item_id,
133 audio_format=media.audio_format,
134 media_type=media_type,
135 stream_type=media.stream_type,
136 path=media.path,
137 duration=media.duration,
138 # a talk clip has nothing worth seeking to, and a seek is the one path that
139 # would re-fetch a possibly-expired HA url mid-playback
140 can_seek=False,
141 allow_seek=False,
142 # a cache hit serves a url that was minted earlier, so it may only claim the life
143 # that url has left or the stream outlives the token behind it
144 expiration=self._remaining_media_lifetime(media),
145 )
146 gain_db = self._loudness_gain(queue_item.queue_id, media.loudness)
147 if gain_db is not None:
148 # core never normalizes a sound effect, so the clip is levelled here or it
149 # airs noticeably quieter than the music around it
150 streamdetails.stream_type = StreamType.CUSTOM
151 # core mirrors what ffmpeg reports onto this object, so it gets a copy of the
152 # constant rather than a handle on the one every clip shares
153 streamdetails.decoded_audio_format = replace(TTS_CLIP_PCM_FORMAT)
154 streamdetails.data = _ClipAudio(media.path, media.audio_format, gain_db)
155 return streamdetails
156
157 async def get_audio_stream(
158 self, streamdetails: StreamDetails, seek_position: int = 0
159 ) -> AsyncGenerator[bytes]:
160 """
161 Return the levelled audio of a spoken clip as PCM.
162
163 :param streamdetails: The StreamDetails previously returned by get_stream_details.
164 :param seek_position: Ignored, a spoken clip cannot be seeked.
165 """
166 clip = cast("_ClipAudio", streamdetails.data)
167 async for chunk in get_ffmpeg_stream(
168 audio_input=clip.path,
169 input_format=clip.input_format,
170 output_format=TTS_CLIP_PCM_FORMAT,
171 filter_params=[
172 TTS_SPEECHNORM_FILTER,
173 f"volume={round(clip.gain_db, 2)}dB",
174 f"alimiter=limit={TTS_PEAK_CEILING_DB}dB:level=false:latency=true",
175 ],
176 ):
177 yield chunk
178
179 def _lock_for(self, clip_id: str) -> asyncio.Lock:
180 """Return the per-clip render lock, creating it on first use."""
181 if not hasattr(self, "_render_locks"):
182 self._render_locks = {}
183 if clip_id not in self._render_locks:
184 self._render_locks[clip_id] = asyncio.Lock()
185 return self._render_locks[clip_id]
186
187 async def _cached_clip_media(
188 self, queue_item: QueueItem, text: str, clip_id: str
189 ) -> _CachedClipMedia:
190 """Return the clip's minted media, re-minting only once the cache entry has expired."""
191 if not hasattr(self, "_media_cache"):
192 self._media_cache = {}
193 now = asyncio.get_running_loop().time()
194 cached = self._media_cache.get(clip_id)
195 if cached is not None and self._remaining_media_lifetime(cached) > MIN_CLIP_MEDIA_LIFETIME:
196 return cached
197 # the caller holds the per-clip render lock, so of the several uncoordinated paths
198 # that resolve the same clip only the first one mints; the rest hit the cache above
199 path, stream_type, audio_format, duration, loudness = await self._mint_clip_media(
200 queue_item, text, clip_id
201 )
202 media = _CachedClipMedia(path, stream_type, audio_format, duration, now, loudness)
203 # clips are minted per queue item, so without pruning the cache grows for as long as
204 # the server runs. an entry past its window can never be served again anyway
205 for expired_id in [
206 key
207 for key, entry in self._media_cache.items()
208 if now - entry.minted_at >= CLIP_STREAMDETAILS_EXPIRATION
209 ]:
210 del self._media_cache[expired_id]
211 self._media_cache[clip_id] = media
212 return media
213
214 def _remaining_media_lifetime(self, media: _CachedClipMedia) -> int:
215 """Return the seconds the given minted media is still usable for."""
216 elapsed = asyncio.get_running_loop().time() - media.minted_at
217 return max(MIN_CLIP_MEDIA_LIFETIME, round(CLIP_STREAMDETAILS_EXPIRATION - elapsed))
218
219 def _wanted_loudness(self, queue_id: str) -> float | None:
220 """Return the level in LUFS a clip should air at, or None when it should air as is."""
221 normalization = self.mass.config.get_effective_player_queue_config_value(
222 queue_id, CONF_VOLUME_NORMALIZATION, CONF_VALUE_ENABLED
223 )
224 if normalization == CONF_VALUE_DISABLED:
225 return None
226 # the queue switch only says normalization may run; the tracks around the clip are
227 # the ones it has to match, and their own preference can still turn it off
228 tracks_mode = self.mass.streams.get_config_value(CONF_VOLUME_NORMALIZATION_TRACKS)
229 if tracks_mode == VolumeNormalizationMode.DISABLED.value:
230 return None
231 target = self.mass.streams.get_config_value(
232 CONF_VOLUME_NORMALIZATION_TARGET, return_type=int
233 )
234 boost = coerce_int(
235 self.config.get_value(CONF_TTS_LOUDNESS_BOOST), DEFAULT_TTS_LOUDNESS_BOOST
236 )
237 return target + boost
238
239 def _loudness_gain(self, queue_id: str, loudness: float | None) -> float | None:
240 """Return the dB to lift the clip by, or None when it should air untouched."""
241 if loudness is None or (wanted := self._wanted_loudness(queue_id)) is None:
242 return None
243 # the reference is taken behind speechnorm, which lands close to the target on its
244 # own, so this trim is small and runs in either direction
245 return wanted - loudness
246
247 def _tts_language(self, host_language: str | None = None) -> str | None:
248 """
249 Return the host's language, or the server locale, as a hyphenated language code.
250
251 :param host_language: The host's configured language override, if any.
252 """
253 if override := (host_language or "").strip():
254 return override.replace("_", "-")
255 return resolve_tts_language(self.mass)
256
257 def _find_clip_item(self, clip_id: str) -> QueueItem | None:
258 """Return the queue item holding the given clip, or None when no queue holds it."""
259 for queue_id in self._candidate_queue_ids(clip_id):
260 if (item := self._find_clip_in_queue(clip_id, queue_id)) is not None:
261 return item
262 return None
263
264 def _candidate_queue_ids(self, clip_id: str) -> list[str]:
265 """
266 Return the queue ids to search for a clip, the most likely one first.
267
268 The owning session knows its queue, but the session registry is empty after a
269 restart while the clip lives on in the persisted queue, so every queue stays a
270 candidate. Clip ids carry a uuid4-based session id, so a hit is unambiguous.
271 """
272 queue_ids = [queue.queue_id for queue in self.mass.player_queues.all()]
273 session = self._sessions.get(clip_id.rpartition("_")[0])
274 if session is not None and session.queue_id in queue_ids:
275 queue_ids.remove(session.queue_id)
276 queue_ids.insert(0, session.queue_id)
277 return queue_ids
278
279 def _find_clip_in_queue(self, clip_id: str, queue_id: str) -> QueueItem | None:
280 """Return the queue item holding the given clip, paging through the queue."""
281 page_size = 500
282 offset = 0
283 while True:
284 page = self.mass.player_queues.items(queue_id, limit=page_size, offset=offset)
285 if not page:
286 return None
287 for item in page:
288 if item.media_item is not None and item.media_item.item_id == clip_id:
289 return item
290 if len(page) < page_size:
291 return None
292 offset += page_size
293
294 async def _generate_script(self, queue_item: QueueItem, prompt: str, clip_id: str) -> str:
295 """Resolve the deferred placeholders and generate the spoken script."""
296 attributes = queue_item.extra_attributes
297 deferred = await self._resolve_deferred_placeholders(prompt)
298 resolved = prompt
299 for key, value in deferred.items():
300 resolved = resolved.replace(key, value)
301 host = self._hosts.get(str(attributes.get(ATTR_HOST_ID) or "")) or {}
302 instructions = str(host.get("instructions") or "")
303 language = str(host.get("language") or "")
304 max_chars = int(attributes.get(ATTR_MAX_CHARS) or 0)
305 web_mode = str(attributes.get(ATTR_WEB_SEARCH_MODE) or "disabled")
306 try:
307 text = cast(
308 "str",
309 await self._generate_text(
310 instructions=instructions,
311 prompt=resolved,
312 web_mode=web_mode,
313 language=language,
314 ),
315 )
316 except Exception as err:
317 self.logger.warning(
318 "AI Radio clip %s (%s) failed to generate: %s", clip_id, queue_item.name, err
319 )
320 self._record_skip(queue_item, f"generation failed: {err}")
321 raise MediaNotFoundError(f"AI Radio clip {clip_id} failed to generate") from err
322 if max_chars > 0:
323 text = soft_limit_text(text, max_chars=max_chars)
324 self.logger.debug(
325 "AI Radio clip %s (%s) rendered: %d chars", clip_id, queue_item.name, len(text)
326 )
327 return text
328
329 async def _resolve_deferred_placeholders(self, prompt: str) -> dict[str, str]:
330 """Return freshly resolved values for the placeholders deferred until airtime."""
331 values = dict.fromkeys(DEFERRED_PLACEHOLDERS, "")
332 values["<timestamp>"] = self._configured_now().strftime("%Y-%m-%d %H:%M %Z")
333 # weather is the only deferred placeholder that costs a network round-trip, so it is
334 # only fetched when the prompt actually references it
335 weather_tokens = ("<weather_hourly>", "<weather_daily>")
336 if any(token in prompt for token in weather_tokens):
337 values.update(await self._prepare_weather_tokens())
338 return values
339
340 async def _mint_clip_media(
341 self, queue_item: QueueItem, text: str, clip_id: str
342 ) -> tuple[str, StreamType, AudioFormat, int | None, float | None]:
343 """Convert the script to playable audio via the configured TTS engine."""
344 host = self._hosts.get(str(queue_item.extra_attributes.get(ATTR_HOST_ID) or "")) or {}
345 engine_uid = str(host.get("tts_engine") or "") or None
346 language = self._tts_language(str(host.get("language") or ""))
347 options = host.get("options") or {}
348 try:
349 path, stream_type, audio_format = await self._render_tts_media(
350 text, engine_uid, language, options
351 )
352 # the probe is the first fetch, so a failed render surfaces here and not in playback
353 duration = await self._probe_duration(path)
354 except Exception as err:
355 self.logger.warning("AI Radio clip %s failed TTS: %s", clip_id, err)
356 self._record_skip(queue_item, f"TTS failed: {err}")
357 raise MediaNotFoundError(f"AI Radio clip {clip_id} failed TTS") from err
358 # measuring costs a fetch and a decode on the just-in-time render path, so it only
359 # runs where the reading has somewhere to go
360 loudness = (
361 await self._reference_loudness(engine_uid, language, options, path, duration)
362 if self._wanted_loudness(queue_item.queue_id) is not None
363 else None
364 )
365 return path, stream_type, audio_format, duration, loudness
366
367 async def _reference_loudness(
368 self,
369 engine_uid: str | None,
370 language: str | None,
371 options: dict[str, Any],
372 path: str,
373 duration: int | None,
374 ) -> float | None:
375 """Return the loudness in LUFS to level this clip against, or None when unknown."""
376 if not hasattr(self, "_engine_loudness"):
377 self._engine_loudness = {}
378 # engine, language and options together decide which voice speaks, and clips from one
379 # voice land within a dB of each other, so measuring one of them is enough
380 key = (engine_uid or "", language or "", json.dumps(options, sort_keys=True, default=str))
381 if (cached := self._engine_loudness.get(key)) is not None:
382 return cached
383 if (loudness := await self._measure_loudness(path)) is None:
384 return None
385 if (duration or 0) >= MIN_LOUDNESS_REFERENCE_SECONDS:
386 self._engine_loudness[key] = loudness
387 return loudness
388
389 async def _measure_loudness(self, path: str) -> float | None:
390 """Return the integrated loudness of the given audio in LUFS, or None when it fails."""
391 try:
392 returncode, output = await check_output(
393 "ffmpeg",
394 "-hide_banner",
395 "-nostats",
396 "-i",
397 path,
398 # measure behind speechnorm: it is what the gain is applied on top of, and it
399 # levels the clip itself, so the reading has to come from its output or the
400 # gain corrects for a level that no longer reaches it
401 "-af",
402 f"{TTS_SPEECHNORM_FILTER},loudnorm=print_format=json",
403 "-f",
404 "null",
405 "-",
406 timeout=LOUDNESS_MEASURE_TIMEOUT,
407 )
408 except (OSError, TimeoutError) as err:
409 self.logger.debug("Could not measure AI Radio clip loudness: %s", err)
410 return None
411 if returncode != 0:
412 self.logger.debug("Could not measure AI Radio clip loudness: ffmpeg failed")
413 return None
414 return parse_loudnorm(output)
415
416 async def _render_tts_media(
417 self,
418 text: str,
419 engine_uid: str | None = None,
420 language: str | None = None,
421 options: dict[str, Any] | None = None,
422 ) -> tuple[str, StreamType, AudioFormat]:
423 """Ask the TTS engine for audio and return the path, stream type and format to play it."""
424 engine = await self._get_tts_engine(engine_uid)
425 stream_details = await query_tts_engine_with_language_fallback(
426 engine, text, language, logger=self.logger, options=options
427 )
428 path, stream_type = await resolve_tts_stream_path(engine, stream_details)
429 audio_format = stream_details.audio_format
430 if audio_format.content_type == ContentType.UNKNOWN:
431 audio_format = AudioFormat(content_type=ContentType.MP3)
432 return path, stream_type, audio_format
433
434 async def _probe_duration(self, path: str) -> int | None:
435 """Return the clip duration in seconds, or None when it cannot be determined."""
436 try:
437 tags = await async_parse_tags(path, require_duration=True)
438 except (InvalidDataError, OSError) as err:
439 if any(marker in str(err) for marker in TTS_SERVER_ERROR_MARKERS):
440 # the engine reports no reason of its own (Home Assistant answers a failed
441 # render with an empty 500), so the probe's message is the only clue there is
442 raise MusicAssistantError(
443 f"{err}. Does your TTS provider have enough credit? "
444 "Check the logs of your TTS provider for the reason."
445 ) from err
446 self.logger.warning("Could not determine AI Radio clip duration: %s", err)
447 return None
448 return int(tags.duration) if tags.duration else None
449
450 def _record_skip(self, queue_item: QueueItem, error: str) -> None:
451 """Record a skipped clip on its owning session."""
452 session_id = str(queue_item.extra_attributes.get(ATTR_SESSION_ID) or "")
453 if (session := self._sessions.get(session_id)) is None:
454 return
455 session.skipped_sections += 1
456 session.last_render_error = error
457