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

1from typing import Union 

2 

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 

8 

9import structlog 

10from adrf.generics import GenericAPIView as AsyncAPIView 

11from adrf.viewsets import ViewSetMixin as AsyncViewSetMixin 

12 

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) 

29 

30 

31logger = structlog.get_logger(__name__) 

32 

33MediaListRequestSerializer = Union[ 

34 media_serializers.PaginatedRequestSerializer, 

35 media_serializers.MediaSearchRequestSerializer, 

36] 

37 

38 

39class InvalidSource(APIException): 

40 status_code = 400 

41 default_detail = "Invalid source." 

42 default_code = "invalid_source" 

43 

44 

45class MediaViewSet(AsyncViewSetMixin, AsyncAPIView, ReadOnlyModelViewSet): 

46 view_is_async = True 

47 

48 lookup_field = "identifier" 

49 lookup_value_converter = "uuid" 

50 

51 pagination_class = StandardPagination 

52 

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 

59 

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) 

71 

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 ) 

87 

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 

93 

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 

101 

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. 

109 

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. 

114 

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 """ 

119 

120 identifiers = [] 

121 hits = [] 

122 for hit in results: 

123 identifiers.append(hit.identifier) 

124 hits.append(hit) 

125 

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) 

130 

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 = [] 

135 

136 return (results, addons) 

137 

138 # Standard actions 

139 

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) 

147 

148 return Response(serializer.data) 

149 

150 def list(self, request, *_, **__): 

151 params = self._get_request_serializer(request) 

152 return self.get_media_results(request, params) 

153 

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 ) 

161 

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. 

166 

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. 

169 

170 :param serializer: the validated serializer instance 

171 :return: whether to include addon model objects 

172 """ 

173 

174 return False 

175 

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"] 

184 

185 hashed_ip = hash(self._get_user_ip(request)) 

186 filter_dead = params.validated_data.get("filter_dead", True) 

187 

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 

196 

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))) 

216 

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 ) 

224 

225 serializer = self.get_serializer(results, many=True, context=serializer_context) 

226 return self.get_paginated_response(serializer.data) 

227 

228 # Extra actions 

229 

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 } 

236 

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) 

242 

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 

261 

262 serializer_context = self.get_serializer_context() 

263 

264 results, _ = self.get_db_results(results) 

265 

266 serializer = self.get_serializer(results, many=True, context=serializer_context) 

267 return self.get_paginated_response(serializer.data) 

268 

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() 

273 

274 return Response(data=serializer.data, status=status.HTTP_201_CREATED) 

275 

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 ) 

280 

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 ) 

292 

293 async def thumbnail(self, request): 

294 serializer = self.get_serializer(data=request.query_params) 

295 serializer.is_valid(raise_exception=True) 

296 

297 media_info = await self.get_image_proxy_media_info() 

298 

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 ) 

306 

307 # Helper functions 

308 

309 @staticmethod 

310 def _get_user_ip(request): 

311 """ 

312 Read request headers to find the correct IP address. 

313 

314 It is assumed that X-Forwarded-For has been sanitized by the load balancer and 

315 thus cannot be rewritten by malicious users. 

316 

317 :param request: a Django request object 

318 :return: an IP address 

319 """ 

320 

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