Coverage for api/views/media_views.py: 88%
150 statements
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 06:14 +0000
« prev ^ index » next coverage.py v7.15.2, created at 2026-10-07 06:14 +0000
1from typing import Union
3from rest_framework import status
4from rest_framework.decorators import action
5from rest_framework.exceptions import APIException, NotFound
6from rest_framework.response import Response
7from rest_framework.viewsets import ReadOnlyModelViewSet
9import structlog
10from adrf.generics import GenericAPIView as AsyncAPIView
11from adrf.viewsets import ViewSetMixin as AsyncViewSetMixin
13from api.constants.media_types import MediaType
14from api.controllers import search_controller
15from api.controllers.elasticsearch.related import related_media
16from api.models import ContentSource
17from api.models.base import OpenLedgerModel
18from api.models.media import AbstractMedia
19from api.serializers import media_serializers
20from api.serializers.source_serializers import SourceSerializer
21from api.utils import image_proxy
22from api.utils.pagination import StandardPagination
23from api.utils.search_context import SearchContext
24from api.utils.throttle import (
25 AnonThumbnailRateThrottle,
26 OAuth2IdThumbnailRateThrottle,
27 OpenverseReferrerAnonThumbnailRateThrottle,
28)
31logger = structlog.get_logger(__name__)
33MediaListRequestSerializer = Union[
34 media_serializers.PaginatedRequestSerializer,
35 media_serializers.MediaSearchRequestSerializer,
36]
39class InvalidSource(APIException):
40 status_code = 400
41 default_detail = "Invalid source."
42 default_code = "invalid_source"
45class MediaViewSet(AsyncViewSetMixin, AsyncAPIView, ReadOnlyModelViewSet):
46 view_is_async = True
48 lookup_field = "identifier"
49 lookup_value_converter = "uuid"
51 pagination_class = StandardPagination
53 # Populate these in the corresponding subclass
54 model_class: type[AbstractMedia] = None
55 addon_model_class: type[OpenLedgerModel] = None
56 media_type: MediaType | None = None
57 query_serializer_class = None
58 default_index = None
60 def __init__(self, *args, **kwargs):
61 super().__init__(*args, **kwargs)
62 required_fields = [
63 self.model_class,
64 self.media_type,
65 self.query_serializer_class,
66 self.default_index,
67 ]
68 if any(val is None for val in required_fields): 68 ↛ 69line 68 didn't jump to line 69 because the condition on line 68 was never true
69 msg = "Viewset fields are not completely populated."
70 raise ValueError(msg)
72 def get_queryset(self):
73 # The alternative to a sub-query would be using `extra` to do a join
74 # to the content source table and filtering `filter_content`. However,
75 # that assumes that a content source entry exists, which is not necessarily
76 # the case. We often don't add a content source until after works from
77 # new source are available in the API, and sometimes not even then.
78 # Search returns results with sources that do not have a ContentSource
79 # table entry. Therefore, to maintain that assumption, a subquery is the only
80 # workable approach, as Django's `extra` does not provide any facility for
81 # handling null relations on the join.
82 return self.model_class.objects.exclude(
83 source__in=ContentSource.objects.filter(filter_content=True).values_list(
84 "source_identifier"
85 )
86 )
88 def get_serializer_context(self):
89 context = super().get_serializer_context()
90 req_serializer = self._get_request_serializer(self.request)
91 context.update({"validated_data": req_serializer.validated_data})
92 return context
94 def _get_request_serializer(self, request):
95 req_serializer = self.query_serializer_class(
96 data=request.query_params,
97 context={"request": request, "media_type": self.media_type},
98 )
99 req_serializer.is_valid(raise_exception=True)
100 return req_serializer
102 def get_db_results(
103 self,
104 results,
105 include_addons=False,
106 ) -> tuple[list[AbstractMedia], list[OpenLedgerModel]]:
107 """
108 Map ES hits to ORM model instances.
110 ORM instances have all necessary info needed for serializers whereas ES
111 hits only contain the subset of fields needed for indexing and search.
112 This function issues one query to the DB, using the ``identifier`` field
113 which is both unique and indexed, so it's quite performant.
115 :param results: the list of ES hits
116 :param include_addons: whether to include add-ons with results
117 :return: the corresponding list of ORM model instances
118 """
120 identifiers = []
121 hits = []
122 for hit in results:
123 identifiers.append(hit.identifier)
124 hits.append(hit)
126 results = list(self.get_queryset().filter(identifier__in=identifiers))
127 results.sort(key=lambda x: identifiers.index(str(x.identifier)))
128 for result, hit in zip(results, hits):
129 result.fields_matched = getattr(hit.meta, "highlight", None)
131 if include_addons and self.addon_model_class:
132 addons = list(self.addon_model_class.objects.filter(pk__in=identifiers))
133 else:
134 addons = []
136 return (results, addons)
138 # Standard actions
140 def retrieve(self, request, *_, **__):
141 instance = self.get_object()
142 search_context = SearchContext.build(
143 [str(instance.identifier)], self.default_index
144 ).asdict()
145 serializer_context = search_context | self.get_serializer_context()
146 serializer = self.get_serializer(instance, context=serializer_context)
148 return Response(serializer.data)
150 def list(self, request, *_, **__):
151 params = self._get_request_serializer(request)
152 return self.get_media_results(request, params)
154 def _validate_source(self, source):
155 valid_sources = search_controller.get_sources(self.media_type)
156 if source not in valid_sources:
157 valid_string = ", ".join([f"'{k}'" for k in valid_sources.keys()])
158 raise InvalidSource(
159 detail=f"Invalid source '{source}'. Valid sources are: {valid_string}.",
160 )
162 def include_addons(self, serializer):
163 """
164 Whether to include objects of the addon model when mapping hits to
165 objects of the media model.
167 If the media type has an addon model, this method should be overridden
168 in the subclass to return ``True`` based on serializer input.
170 :param serializer: the validated serializer instance
171 :return: whether to include addon model objects
172 """
174 return False
176 def get_media_results(
177 self,
178 request,
179 params: MediaListRequestSerializer,
180 ):
181 page_size = self.paginator.page_size = params.data["page_size"]
182 page = self.paginator.page = params.data["page"]
183 self.paginator.warnings = params.context["warnings"]
185 hashed_ip = hash(self._get_user_ip(request))
186 filter_dead = params.validated_data.get("filter_dead", True)
188 if pref_index := params.validated_data.get("index"): 188 ↛ 189line 188 didn't jump to line 189 because the condition on line 188 was never true
189 logger.info(f"Using preferred index {pref_index} for media.")
190 search_index = pref_index
191 exact_index = True
192 else:
193 logger.info("Using default index for media.")
194 search_index = self.default_index
195 exact_index = False
197 try:
198 (
199 results,
200 num_pages,
201 num_results,
202 search_context,
203 ) = search_controller.query_media(
204 params,
205 search_index,
206 exact_index,
207 page_size,
208 hashed_ip,
209 filter_dead,
210 page,
211 )
212 self.paginator.page_count = params.clamp_page_count(num_pages)
213 self.paginator.result_count = params.clamp_result_count(num_results)
214 except ValueError as e:
215 raise APIException(getattr(e, "message", str(e)))
217 include_addons = self.include_addons(params)
218 results, addons = self.get_db_results(results, include_addons)
219 serializer_context = (
220 search_context
221 | self.get_serializer_context()
222 | {"addons": {addon.audio_identifier: addon for addon in addons}}
223 )
225 serializer = self.get_serializer(results, many=True, context=serializer_context)
226 return self.get_paginated_response(serializer.data)
228 # Extra actions
230 @action(detail=False, serializer_class=SourceSerializer, pagination_class=None)
231 def stats(self, *_, **__):
232 source_counts = search_controller.get_sources(self.default_index)
233 context = self.get_serializer_context() | {
234 "source_counts": source_counts,
235 }
237 sources = ContentSource.objects.filter(
238 media_type=self.default_index, filter_content=False
239 )
240 serializer = self.get_serializer(sources, many=True, context=context)
241 return Response(serializer.data)
243 @action(detail=True)
244 def related(self, request, identifier=None, *_, **__):
245 try:
246 results = related_media(
247 uuid=identifier,
248 index=self.default_index,
249 filter_dead=True,
250 )
251 self.paginator.page_count = 1
252 # `page_size` refers to the maximum number of related images to return.
253 self.paginator.page_size = 10
254 # `result_count` is hard-coded and is equal to the page size.
255 self.paginator.result_count = 10
256 except ValueError as e:
257 raise APIException(getattr(e, "message", str(e)))
258 # If there are no hits in the search controller
259 except IndexError:
260 raise NotFound
262 serializer_context = self.get_serializer_context()
264 results, _ = self.get_db_results(results)
266 serializer = self.get_serializer(results, many=True, context=serializer_context)
267 return self.get_paginated_response(serializer.data)
269 def report(self, request, identifier):
270 serializer = self.get_serializer(data=request.data | {"identifier": identifier})
271 serializer.is_valid(raise_exception=True)
272 serializer.save()
274 return Response(data=serializer.data, status=status.HTTP_201_CREATED)
276 async def get_image_proxy_media_info(self) -> image_proxy.MediaInfo:
277 raise NotImplementedError(
278 "Subclasses must implement `get_image_proxy_media_info`"
279 )
281 thumbnail_action = action(
282 detail=True,
283 url_path="thumb",
284 url_name="thumb",
285 serializer_class=media_serializers.MediaThumbnailRequestSerializer,
286 throttle_classes=[
287 AnonThumbnailRateThrottle,
288 OpenverseReferrerAnonThumbnailRateThrottle,
289 OAuth2IdThumbnailRateThrottle,
290 ],
291 )
293 async def thumbnail(self, request):
294 serializer = self.get_serializer(data=request.query_params)
295 serializer.is_valid(raise_exception=True)
297 media_info = await self.get_image_proxy_media_info()
299 return await image_proxy.get(
300 media_info,
301 request_config=image_proxy.RequestConfig(
302 accept_header=request.headers.get("Accept", "image/*"),
303 **serializer.validated_data,
304 ),
305 )
307 # Helper functions
309 @staticmethod
310 def _get_user_ip(request):
311 """
312 Read request headers to find the correct IP address.
314 It is assumed that X-Forwarded-For has been sanitized by the load balancer and
315 thus cannot be rewritten by malicious users.
317 :param request: a Django request object
318 :return: an IP address
319 """
321 x_forwarded_for = request.META.get("HTTP_X_FORWARDED_FOR")
322 if x_forwarded_for: 322 ↛ 323line 322 didn't jump to line 323 because the condition on line 322 was never true
323 ip = x_forwarded_for.split(",")[0]
324 else:
325 ip = request.META.get("REMOTE_ADDR")
326 return ip