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