/
/
1"""Manage MediaItems of type Podcast."""
2
3from __future__ import annotations
4
5from collections.abc import AsyncGenerator
6from typing import TYPE_CHECKING, Any, cast
7
8from music_assistant_models.auth import Scope
9from music_assistant_models.enums import MediaType, ProviderFeature
10from music_assistant_models.errors import MediaNotFoundError, ProviderUnavailableError
11from music_assistant_models.helpers import create_safe_string
12from music_assistant_models.media_items import (
13 Podcast,
14 PodcastEpisode,
15 PodcastSummary,
16 ProviderMapping,
17 UniqueList,
18)
19
20from music_assistant.constants import DB_TABLE_PLAYLOG, DB_TABLE_PODCASTS
21from music_assistant.controllers.webserver.helpers.auth_middleware import get_current_user
22from music_assistant.helpers.audio import get_probed_duration
23from music_assistant.helpers.compare import (
24 compare_media_item,
25 compare_podcast,
26 loose_compare_strings,
27)
28from music_assistant.helpers.database import UNSET
29from music_assistant.helpers.json import serialize_to_json
30from music_assistant.models.music_provider import MusicProvider
31
32from .base import MediaControllerBase
33
34if TYPE_CHECKING:
35 from collections.abc import Mapping
36
37 from music_assistant_models.auth import User
38
39 from music_assistant import MusicAssistant
40
41
42class PodcastsController(MediaControllerBase[Podcast]):
43 """Controller managing MediaItems of type Podcast."""
44
45 db_table = DB_TABLE_PODCASTS
46 media_type = MediaType.PODCAST
47 item_cls = Podcast
48 summary_item_cls = PodcastSummary
49
50 def __init__(self, mass: MusicAssistant) -> None:
51 """Initialize class."""
52 super().__init__(mass)
53 # register (extra) api handlers
54 api_base = self.api_base
55 self.mass.register_api_command(
56 f"music/{api_base}/podcast_episodes", self.episodes, required_scope=Scope.LIBRARY_READ
57 )
58 self.mass.register_api_command(
59 f"music/{api_base}/podcast_episode", self.episode, required_scope=Scope.LIBRARY_READ
60 )
61 self.mass.register_api_command(
62 f"music/{api_base}/podcast_versions", self.versions, required_scope=Scope.LIBRARY_READ
63 )
64
65 @property
66 def summary_query(self) -> tuple[str, dict[str, Any]]:
67 """Return the slim SELECT query used for podcast summary listings."""
68 query = f"""
69 SELECT
70 {self._summary_base_columns()},
71 podcasts.version,
72 podcasts.publisher,
73 podcasts.total_episodes,
74 {self._provider_mappings_query()} AS provider_mappings
75 FROM podcasts"""
76 return query, {}
77
78 async def library_items(
79 self,
80 favorite: bool | None = None,
81 search: str | None = None,
82 limit: int = 500,
83 offset: int = 0,
84 order_by: str = "sort_name",
85 provider: str | list[str] | None = None,
86 genre: int | list[int] | None = None,
87 played_only: bool = False,
88 *,
89 summary: bool = True,
90 **kwargs: Any,
91 ) -> list[Podcast]:
92 """
93 Get in-database podcasts.
94
95 :param favorite: Filter by favorite status.
96 :param search: Filter by search query.
97 :param limit: Maximum number of items to return.
98 :param offset: Number of items to skip.
99 :param order_by: Order by field (e.g. 'sort_name', 'timestamp_added').
100 :param provider: Filter by provider instance ID (single string or list).
101 :param genre: Filter by genre id(s).
102 :param summary: When True (default), return slim summary items containing only the
103 fields needed for a list view. Set to False to get fully hydrated items.
104 """
105 result = await self.get_library_items_by_query(
106 favorite=favorite,
107 search=search,
108 genre_ids=genre,
109 limit=limit,
110 offset=offset,
111 order_by=order_by,
112 provider_filter=self._ensure_provider_filter(provider),
113 played_only=played_only,
114 in_library_only=True,
115 summary=summary,
116 )
117 if search and len(result) < 25 and not offset:
118 # append publisher items to result
119 extra_query_parts: list[str] = [
120 "WHERE podcasts.publisher LIKE :search",
121 ]
122 extra_query_params: dict[str, Any] = {
123 "search": f"%{search}%",
124 }
125 return result + await self.get_library_items_by_query(
126 favorite=favorite,
127 search=None,
128 genre_ids=genre,
129 limit=limit,
130 order_by=order_by,
131 provider_filter=self._ensure_provider_filter(provider),
132 extra_query_parts=extra_query_parts,
133 extra_query_params=extra_query_params,
134 in_library_only=True,
135 summary=summary,
136 )
137 return result
138
139 async def episodes(
140 self,
141 item_id: str,
142 provider_instance_id_or_domain: str,
143 ) -> AsyncGenerator[PodcastEpisode]:
144 """Return podcast episodes for the given provider podcast id."""
145 # always check if we have a library item for this podcast
146 if provider_instance_id_or_domain == "library":
147 library_podcast = await self.get_library_item(item_id)
148 if not library_podcast:
149 raise MediaNotFoundError(f"Podcast {item_id} not found in library")
150 provider_instance_id_or_domain, item_id = self._select_provider_id(library_podcast)
151 # podcast episodes are not stored in the db/library
152 # so we always need to fetch them from the provider
153 async for episode in self._get_provider_podcast_episodes(
154 item_id, provider_instance_id_or_domain
155 ):
156 yield episode
157
158 async def episode(
159 self,
160 item_id: str,
161 provider_instance_id_or_domain: str,
162 ) -> PodcastEpisode:
163 """Return single podcast episode by the given provider podcast id."""
164 prov = self.mass.get_provider(provider_instance_id_or_domain)
165 if not isinstance(prov, MusicProvider):
166 raise ProviderUnavailableError("Provider not found")
167 episode = await prov.get_podcast_episode(item_id)
168 await self._restore_probed_duration(episode)
169 return episode
170
171 async def versions(
172 self,
173 item_id: str,
174 provider_instance_id_or_domain: str,
175 ) -> UniqueList[Podcast]:
176 """Return all versions of an podcast we can find on all providers."""
177 podcast = await self.get_provider_item(item_id, provider_instance_id_or_domain)
178 search_query = podcast.name
179 result: UniqueList[Podcast] = UniqueList()
180 for provider_id in self.mass.music.get_unique_providers():
181 provider = self.mass.get_provider(provider_id)
182 if not isinstance(provider, MusicProvider):
183 continue
184 if not self.mass.music.library_supported(provider, MediaType.PODCAST):
185 continue
186 result.extend(
187 prov_item
188 for prov_item in await self.search(search_query, provider_id)
189 if loose_compare_strings(podcast.name, prov_item.name)
190 # make sure that the 'base' version is NOT included
191 and not podcast.provider_mappings.intersection(prov_item.provider_mappings)
192 )
193 return result
194
195 async def match_provider(
196 self, db_podcast: Podcast, provider: MusicProvider, strict: bool = True
197 ) -> list[ProviderMapping]:
198 """
199 Try to find match on (streaming) provider for the provided (database) podcast.
200
201 This is used to link objects of different providers/qualities together.
202 """
203 self.logger.debug(
204 "Trying to match podcast %s on provider %s",
205 db_podcast.name,
206 provider.name,
207 )
208 matches: list[ProviderMapping] = []
209 search_str = db_podcast.name
210 search_result = await self.search(search_str, provider.instance_id)
211 for search_result_item in search_result:
212 if not search_result_item.available:
213 continue
214 if not compare_media_item(db_podcast, search_result_item, strict=strict):
215 continue
216 # we must fetch the full podcast version, search results can be simplified objects
217 prov_podcast = await self.get_provider_item(
218 search_result_item.item_id,
219 search_result_item.provider,
220 fallback=search_result_item,
221 )
222 if compare_podcast(db_podcast, prov_podcast, strict=strict):
223 # 100% match
224 matches.extend(prov_podcast.provider_mappings)
225 if not matches:
226 self.logger.debug(
227 "Could not find match for Podcast %s on provider %s",
228 db_podcast.name,
229 provider.name,
230 )
231 return matches
232
233 async def match_providers(self, db_podcast: Podcast) -> None:
234 """
235 Try to find match on all (streaming) providers for the provided (database) podcast.
236
237 This is used to link objects of different providers/qualities together.
238 """
239 if db_podcast.provider != "library":
240 return # Matching only supported for database items
241
242 # try to find match on all providers
243 cur_provider_domains = {x.provider_domain for x in db_podcast.provider_mappings}
244 for provider in self.mass.music.providers:
245 if provider.domain in cur_provider_domains:
246 continue
247 if ProviderFeature.SEARCH not in provider.supported_features:
248 continue
249 if not self.mass.music.library_supported(provider, MediaType.PODCAST):
250 continue
251 if not provider.is_streaming_provider:
252 # matching on unique providers is pointless as they push (all) their content to MA
253 continue
254 if match := await self.match_provider(db_podcast, provider):
255 # 100% match, we update the db with the additional provider mapping(s)
256 await self.add_provider_mappings(db_podcast.item_id, match)
257 cur_provider_domains.add(provider.domain)
258
259 async def _add_library_item(self, item: Podcast, overwrite_existing: bool = False) -> int:
260 """Add a new record to the database."""
261 db_id = await self.mass.music.database.insert(
262 self.db_table,
263 {
264 "name": item.name,
265 "sort_name": item.sort_name,
266 "version": item.version,
267 "favorite": item.favorite,
268 "metadata": serialize_to_json(item.metadata),
269 "publisher": item.publisher,
270 "total_episodes": item.total_episodes or 0,
271 "search_name": create_safe_string(item.name, True, True),
272 "search_sort_name": create_safe_string(item.sort_name or "", True, True),
273 "timestamp_added": int(item.date_added.timestamp()) if item.date_added else UNSET,
274 },
275 )
276 # update/set external id lookup table
277 await self.set_external_ids(db_id, item.external_ids)
278 # update/set provider_mappings table
279 await self.set_provider_mappings(db_id, item.provider_mappings)
280 self.logger.debug("added %s to database (id: %s)", item.name, db_id)
281 return db_id
282
283 async def _update_library_item(
284 self, item_id: str | int, update: Podcast, overwrite: bool = False
285 ) -> None:
286 """Update existing record in the database."""
287 db_id = int(item_id) # ensure integer
288 cur_item = await self.get_library_item(db_id)
289 metadata = update.metadata if overwrite else cur_item.metadata.update(update.metadata)
290 if not overwrite and update.metadata.images is not None:
291 # podcasts have no image picker, so keep the cover in sync with the
292 # provider instead of accumulating merged entries
293 metadata.images = update.metadata.images
294 cur_item.external_ids.update(update.external_ids)
295 name = update.name if overwrite else cur_item.name
296 sort_name = update.sort_name if overwrite else cur_item.sort_name or update.sort_name
297 await self.mass.music.database.update(
298 self.db_table,
299 {"item_id": db_id},
300 {
301 "name": name,
302 "sort_name": sort_name,
303 "version": update.version if overwrite else cur_item.version or update.version,
304 "metadata": serialize_to_json(metadata),
305 "publisher": cur_item.publisher or update.publisher,
306 "total_episodes": cur_item.total_episodes or update.total_episodes or 0,
307 "search_name": create_safe_string(name, True, True),
308 "search_sort_name": create_safe_string(sort_name or "", True, True),
309 "timestamp_added": int(update.date_added.timestamp())
310 if update.date_added
311 else UNSET,
312 },
313 )
314 # update/set external id lookup table
315 await self.set_external_ids(
316 db_id, update.external_ids if overwrite else cur_item.external_ids
317 )
318 # update/set provider_mappings table
319 provider_mappings = (
320 update.provider_mappings
321 if overwrite
322 else {*update.provider_mappings, *cur_item.provider_mappings}
323 )
324 await self.set_provider_mappings(db_id, provider_mappings, overwrite)
325 self.logger.debug("updated %s in database: (id %s)", update.name, db_id)
326
327 async def _get_provider_podcast_episodes(
328 self, item_id: str, provider_instance_id_or_domain: str
329 ) -> AsyncGenerator[PodcastEpisode]:
330 """Return podcast episodes for the given provider podcast id."""
331 prov = self.mass.get_provider(provider_instance_id_or_domain)
332 if not isinstance(prov, MusicProvider):
333 return
334
335 # Get user who initiated the query. Querying the userid as well is most useful
336 # in a multi-user environment where a single instance provider is used.
337 user: User | None = None
338 if session_user := get_current_user():
339 # this is the active session user that triggered the action
340 user = session_user
341 elif provider_user := await self.mass.music._get_user_for_provider(
342 provider_mappings_or_instance_id=provider_instance_id_or_domain
343 ):
344 # based on configured provider filter we can try to find a user
345 user = provider_user
346
347 # fetched in one query on first use instead of one per episode: a podcast can have
348 # thousands of them
349 resume_rows: dict[str, Mapping[str, Any]] | None = None
350
351 async def load_resume_rows() -> dict[str, Mapping[str, Any]]:
352 match: dict[str, Any] = {
353 "provider": prov.instance_id,
354 "media_type": MediaType.PODCAST_EPISODE,
355 }
356 if user is not None:
357 match["userid"] = user.user_id
358 # limit=0 lifts get_rows' 500 row default, which combined with the ascending sort
359 # would drop the newest rows - the part-played episodes this lookup is for. That
360 # sort also picks the newest row per item_id in the map below, where without a
361 # userid filter several users can hold one
362 rows = await self.mass.music.database.get_rows(
363 DB_TABLE_PLAYLOG, match=match, order_by="timestamp", limit=0
364 )
365 return {row["item_id"]: row for row in rows}
366
367 async def set_resume_position(episode: PodcastEpisode) -> None:
368 nonlocal resume_rows
369 if episode.fully_played is not None or episode.resume_position_ms:
370 # provider supports resume info, we can skip
371 return
372 # for providers that do not natively support providing resume info,
373 # we fallback to the playlog db table
374 if resume_rows is None:
375 resume_rows = await load_resume_rows()
376 resume_info_db_row = resume_rows.get(episode.item_id)
377 if resume_info_db_row is None:
378 return
379 if resume_info_db_row["seconds_played"]:
380 episode.resume_position_ms = int(resume_info_db_row["seconds_played"] * 1000)
381 if resume_info_db_row["fully_played"] is not None:
382 episode.fully_played = bool(resume_info_db_row["fully_played"])
383
384 # grab the episodes from the provider. Providers cache their own listing, so resume
385 # info is applied here to keep per-user progress out of those caches
386 async for item in prov.get_podcast_episodes(item_id):
387 await set_resume_position(item)
388 await self._restore_probed_duration(item)
389 yield item
390
391 async def _restore_probed_duration(self, episode: PodcastEpisode) -> None:
392 """
393 Fill in the duration determined during an earlier playback, for feeds that omit it.
394
395 :param episode: The episode to fill the duration of, left untouched when it has one.
396 """
397 if episode.duration or not (uri := episode.uri):
398 return
399 if probed_duration := await get_probed_duration(self.mass, uri):
400 episode.duration = probed_duration
401
402 def _parse_summary_row(self, db_row: Mapping[str, Any]) -> PodcastSummary:
403 """Parse a raw summary db row into a PodcastSummary object."""
404 item = cast("PodcastSummary", super()._parse_summary_row(db_row))
405 item.version = db_row["version"] or ""
406 item.publisher = db_row["publisher"]
407 item.total_episodes = db_row["total_episodes"]
408 return item
409