Coverage for api/controllers/search_controller.py: 81%

219 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 06:14 +0000

1from __future__ import annotations 

2 

3import re 

4from math import ceil 

5from typing import TYPE_CHECKING 

6 

7from django.conf import settings 

8from django.core.cache import cache 

9 

10import structlog 

11from decouple import config 

12from elasticsearch.exceptions import NotFoundError 

13from elasticsearch_dsl import Q, Search 

14from elasticsearch_dsl.query import EMPTY_QUERY 

15from elasticsearch_dsl.response import Hit, Response 

16from redis.exceptions import ConnectionError 

17 

18import api.models as models 

19from api.constants.media_types import OriginIndex, SearchIndex 

20from api.constants.search import SearchStrategy 

21from api.constants.sorting import INDEXED_ON 

22from api.controllers.elasticsearch.helpers import ( 

23 ELASTICSEARCH_MAX_RESULT_WINDOW, 

24 get_es_response, 

25 get_query_slice, 

26 get_raw_es_response, 

27) 

28from api.utils import tallies 

29from api.utils.check_dead_links import check_dead_links 

30from api.utils.dead_link_mask import get_query_hash 

31from api.utils.search_context import SearchContext 

32 

33 

34# Using TYPE_CHECKING to avoid circular imports when importing types 

35if TYPE_CHECKING: 35 ↛ 36line 35 didn't jump to line 36 because the condition on line 35 was never true

36 from api.serializers.media_serializers import MediaSearchRequestSerializer 

37 

38logger = structlog.get_logger(__name__) 

39 

40 

41NESTING_THRESHOLD = config("POST_PROCESS_NESTING_THRESHOLD", cast=int, default=5) 

42SOURCE_CACHE_TIMEOUT = 60 * 60 * 4 # 4 hours 

43FILTER_CACHE_TIMEOUT = 30 

44FILTERED_SOURCES_CACHE_KEY = "filtered_sources" 

45FILTERED_SOURCES_CACHE_VERSION = 1 

46DEFAULT_BOOST = 10000 

47DEFAULT_SEARCH_FIELDS = ["title", "description", "tags.name"] 

48DEFAULT_SQS_FLAGS = "AND|NOT|PHRASE|WHITESPACE" 

49UNUSED_SQS_FLAGS = [ 

50 ("PRECEDENCE", r"\(.*\)"), 

51 ("ESCAPE", r"\\"), 

52 ("FUZZY|SLOP", r"~\d"), 

53 ("PREFIX", r"\*"), 

54] 

55 

56 

57def _quote_escape(query_string): 

58 """Ignore any unmatched quotes in the query supplied by the user.""" 

59 

60 num_quotes = query_string.count('"') 

61 if num_quotes % 2 == 1: 

62 return query_string.replace('"', '\\"') 

63 else: 

64 return query_string 

65 

66 

67def _post_process_results( 

68 s, start, end, page_size, search_results, filter_dead, nesting=0 

69) -> list[Hit] | None: 

70 """ 

71 Perform some steps on results fetched from the backend. 

72 

73 After fetching the search results from the back end, iterate through the 

74 results, perform image validation, and route certain thumbnails through our 

75 proxy. 

76 

77 Keeps making new query requests until it is able to fill the page size. 

78 

79 :param s: The Elasticsearch Search object. 

80 :param start: The start of the result slice. 

81 :param end: The end of the result slice. 

82 :param search_results: The Elasticsearch response object containing search 

83 results. 

84 :param filter_dead: Whether images should be validated. 

85 :param nesting: the level of nesting at which this function is being called 

86 :return: List of results. 

87 """ 

88 

89 if nesting > NESTING_THRESHOLD: 89 ↛ 90line 89 didn't jump to line 90 because the condition on line 89 was never true

90 logger.info( 

91 "Nesting threshold breached", 

92 nesting=nesting, 

93 start=start, 

94 end=end, 

95 page_size=page_size, 

96 ) 

97 

98 results = list(search_results) 

99 

100 if filter_dead: 

101 query_hash = get_query_hash(s) 

102 check_dead_links(query_hash, start, results) 

103 

104 if len(results) == 0: 

105 # first page is all dead links 

106 return None 

107 

108 if len(results) < page_size: 

109 """ 

110 The variables in this function get updated in an interesting way. 

111 Here is an example of that for a typical query. Note that ``end`` 

112 increases but start stays the same. This has the effect of slowly 

113 increasing the size of the query we send to Elasticsearch with the 

114 goal of backfilling the results until we have enough valid (live) 

115 results to fulfill the requested page size. 

116 

117 ``` 

118 page_size: 20 

119 page: 1 

120 

121 start: 0 

122 end: 40 (DEAD_LINK_RATIO applied) 

123 

124 end gets updated to end + end/2 = 60 

125 

126 end = 90 

127 end = 90 + 45 

128 ``` 

129 """ 

130 if end >= search_results.hits.total.value: 

131 # Total available hits already exhausted in previous iteration 

132 return results 

133 

134 end += int(end / 2) 

135 query_size = start + end 

136 if query_size > ELASTICSEARCH_MAX_RESULT_WINDOW: 136 ↛ 137line 136 didn't jump to line 137 because the condition on line 136 was never true

137 return results 

138 

139 # subtract start to account for the records skipped 

140 # and which should not count towards the total 

141 # available hits for the query 

142 total_available_hits = search_results.hits.total.value - start 

143 if query_size > total_available_hits: 143 ↛ 148line 143 didn't jump to line 148 because the condition on line 143 was never true

144 # Clamp the query size to last available hit. On the next 

145 # iteration, if results are still insufficient, the check 

146 # to compare previous_query_size and total_available_hits 

147 # will prevent further query attempts 

148 end = search_results.hits.total.value 

149 

150 s = s[start:end] 

151 search_response = get_es_response(s, es_query="postprocess_search") 

152 

153 return _post_process_results( 

154 s, start, end, page_size, search_response, filter_dead, nesting + 1 

155 ) 

156 

157 return results[:page_size] 

158 

159 

160def get_excluded_sources_query() -> Q | None: 

161 """ 

162 Hide data sources from the catalog dynamically. 

163 To exclude a source, set ``filter_content`` to ``True`` in the 

164 ``ContentSource`` model in Django admin. 

165 The list of ``source_identifier``s is cached in Redis with 

166 `:FILTERED_SOURCES_CACHE_VERSION:FILTERED_SOURCES_CACHE_KEY` key. 

167 """ 

168 

169 try: 

170 filtered_sources = cache.get( 

171 key=FILTERED_SOURCES_CACHE_KEY, version=FILTERED_SOURCES_CACHE_VERSION 

172 ) 

173 logger.info(f"Filtered sources from cache: {filtered_sources}") 

174 except ConnectionError: 

175 logger.warning("Redis connect failed, cannot get cached filtered sources.") 

176 filtered_sources = None 

177 

178 if not filtered_sources: 178 ↛ 195line 178 didn't jump to line 195 because the condition on line 178 was always true

179 filtered_sources = list( 

180 models.ContentSource.objects.filter(filter_content=True).values_list( 

181 "source_identifier", flat=True 

182 ) 

183 ) 

184 

185 try: 

186 cache.set( 

187 key=FILTERED_SOURCES_CACHE_KEY, 

188 version=FILTERED_SOURCES_CACHE_VERSION, 

189 timeout=FILTER_CACHE_TIMEOUT, 

190 value=filtered_sources, 

191 ) 

192 except ConnectionError: 

193 logger.warning("Redis connect failed, cannot cache filtered sources.") 

194 

195 if filtered_sources: 195 ↛ 196line 195 didn't jump to line 196 because the condition on line 195 was never true

196 return Q("terms", source=filtered_sources) 

197 return None 

198 

199 

200def get_index( 

201 exact_index: bool, 

202 origin_index: OriginIndex, 

203 search_params: MediaSearchRequestSerializer, 

204) -> SearchIndex: 

205 if exact_index: 205 ↛ 206line 205 didn't jump to line 206 because the condition on line 205 was never true

206 return origin_index 

207 

208 include_sensitive_results = search_params.validated_data.get( 

209 "include_sensitive_results", False 

210 ) 

211 if settings.ENABLE_FILTERED_INDEX_QUERIES and not include_sensitive_results: 

212 return f"{origin_index}-filtered" 

213 return origin_index 

214 

215 

216def create_search_filter_queries( 

217 search_params: MediaSearchRequestSerializer, 

218) -> dict[str, list[Q]]: 

219 """ 

220 Create a list of Elasticsearch queries for filtering search results. 

221 The filter values are given in the request query string. 

222 We use ES filters (`filter`, `must_not`) because we don't need to 

223 compute the relevance score and the queries are cached for better 

224 performance. 

225 """ 

226 queries = {"filter": [], "must_not": []} 

227 # Apply term filters. Each tuple pairs a filter's parameter name in the API 

228 # with its corresponding field in Elasticsearch. "None" means that the 

229 # names are identical. 

230 query_filters = { 

231 "filter": [ 

232 ("extension", None), 

233 ("category", None), 

234 ("source", None), 

235 ("license", None), 

236 ("license_type", "license"), 

237 # Audio-specific filters 

238 ("length", None), 

239 # Image-specific filters 

240 ("aspect_ratio", None), 

241 ("size", None), 

242 ], 

243 "must_not": [ 

244 ("excluded_source", "source"), 

245 ], 

246 } 

247 for behaviour, filters in query_filters.items(): 

248 for serializer_field, es_field in filters: 

249 if not (arguments := search_params.data.get(serializer_field)): 

250 continue 

251 arguments = arguments.split(",") 

252 parameter = es_field or serializer_field 

253 queries[behaviour].append(Q("terms", **{parameter: arguments})) 

254 return queries 

255 

256 

257def create_ranking_queries( 

258 search_params: MediaSearchRequestSerializer, 

259) -> list[Q]: 

260 queries = [Q("rank_feature", field="standardized_popularity", boost=DEFAULT_BOOST)] 

261 if search_params.data["unstable__authority"]: 

262 boost = int(search_params.data["unstable__authority_boost"] * DEFAULT_BOOST) 

263 authority_query = Q("rank_feature", field="authority_boost", boost=boost) 

264 queries.append(authority_query) 

265 return queries 

266 

267 

268def build_search_query( 

269 search_params: MediaSearchRequestSerializer, 

270) -> Q: 

271 # Apply filters from the url query search parameters. 

272 url_queries = create_search_filter_queries(search_params) 

273 search_queries = { 

274 "filter": url_queries["filter"], 

275 "must_not": url_queries["must_not"], 

276 "must": [], 

277 "should": [], 

278 } 

279 

280 # Exclude mature content 

281 if not search_params.validated_data["include_sensitive_results"]: 

282 search_queries["must_not"].append(Q("term", mature=True)) 

283 # Exclude dynamically disabled sources (see Redis cache) 

284 if excluded_sources_query := get_excluded_sources_query(): 284 ↛ 285line 284 didn't jump to line 285 because the condition on line 284 was never true

285 search_queries["must_not"].append(excluded_sources_query) 

286 

287 # Search either by generic multimatch or by "advanced search" with 

288 # individual field-level queries specified. 

289 if "q" in search_params.data: 

290 query = _quote_escape(search_params.data["q"]) 

291 log_query_features(query, query_name="q") 

292 

293 base_query_kwargs = { 

294 "query": query, 

295 "flags": DEFAULT_SQS_FLAGS, 

296 "fields": DEFAULT_SEARCH_FIELDS, 

297 "default_operator": "AND", 

298 } 

299 

300 if '"' in query: 

301 base_query_kwargs["quote_field_suffix"] = ".raw" 

302 

303 search_queries["must"].append(Q("simple_query_string", **base_query_kwargs)) 

304 # Boost exact matches on the title 

305 exact_match_boost = Q("match_phrase", title={"query": query, "boost": 10000}) 

306 search_queries["should"].append(exact_match_boost) 

307 else: 

308 for field, field_name in [ 

309 ("creator", "creator"), 

310 ("title", "title"), 

311 ("tags", "tags.name"), 

312 ]: 

313 if field_value := search_params.data.get(field): 

314 log_query_features(field_value, query_name="field") 

315 search_queries["must"].append( 

316 Q( 

317 "simple_query_string", 

318 flags=DEFAULT_SQS_FLAGS, 

319 query=_quote_escape(field_value), 

320 fields=[field_name], 

321 ) 

322 ) 

323 

324 if settings.USE_RANK_FEATURES: 324 ↛ 330line 324 didn't jump to line 330 because the condition on line 324 was always true

325 search_queries["should"].extend(create_ranking_queries(search_params)) 

326 

327 # If there are no `must` query clauses, only the results that match 

328 # the `should` clause are returned. To avoid this, we add an empty 

329 # query clause to the `must` list. 

330 if not search_queries["must"]: 

331 search_queries["must"].append(EMPTY_QUERY) 

332 

333 return Q( 

334 "bool", 

335 filter=search_queries["filter"], 

336 must_not=search_queries["must_not"], 

337 must=search_queries["must"], 

338 should=search_queries["should"], 

339 ) 

340 

341 

342def log_query_features(query: str, query_name) -> None: 

343 query_flags = [] 

344 for flag, pattern in UNUSED_SQS_FLAGS: 

345 if bool(re.search(pattern, query)): 

346 query_flags.append(flag) 

347 if query_flags: 

348 logger.info( 

349 { 

350 "log_message": "Special features present in query", 

351 "query_name": query_name, 

352 "query": query, 

353 "flags": query_flags, 

354 } 

355 ) 

356 

357 

358def build_collection_query( 

359 search_params: MediaSearchRequestSerializer, 

360): 

361 """ 

362 Build the query to retrieve items in a collection. 

363 :param search_params: the validated search parameters. 

364 :return: the search client with the query applied. 

365 """ 

366 search_query = {"filter": [], "must": [], "should": [], "must_not": []} 

367 # Apply the term filters. Each tuple pairs a filter's parameter name in the API 

368 # with its corresponding field in Elasticsearch. "None" means that the 

369 # names are identical. 

370 filters = [ 

371 ("tag", "tags.name.keyword"), 

372 ("source", None), 

373 ("creator", "creator.keyword"), 

374 ] 

375 for serializer_field, es_field in filters: 

376 if argument := search_params.validated_data.get(serializer_field): 

377 parameter = es_field or serializer_field 

378 search_query["filter"].append({"term": {parameter: argument}}) 

379 

380 # Exclude mature content and disabled sources 

381 include_sensitive_by_params = search_params.validated_data.get( 

382 "include_sensitive_results", False 

383 ) 

384 if not include_sensitive_by_params: 

385 search_query["must_not"].append({"term": {"mature": True}}) 

386 

387 if excluded_sources_query := get_excluded_sources_query(): 

388 search_query["must_not"].append(excluded_sources_query) 

389 

390 return Q("bool", **search_query) 

391 

392 

393query_builders = { 

394 "search": build_search_query, 

395 "collection": build_collection_query, 

396} 

397 

398 

399def query_media( 

400 search_params: MediaSearchRequestSerializer, 

401 origin_index: OriginIndex, 

402 exact_index: bool, 

403 page_size: int, 

404 ip: int, 

405 filter_dead: bool, 

406 page: int = 1, 

407) -> tuple[list[Hit], int, int, dict]: 

408 """ 

409 Build the search or collection query, execute it and return 

410 paginated result. 

411 For queries with `collection` parameter, returns media filtered 

412 by the `tag`, `source` or `source`/`creator` combination, ordered 

413 by the time when they were added to Openverse. 

414 For other queries, performs a ranked paginated search 

415 from the set of keywords and, optionally, filters. 

416 

417 :param search_params: Search query params, see :class: `MediaSearchRequestSerializer`. 

418 :param origin_index: The Elasticsearch index to search (e.g. 'image') 

419 :param exact_index: whether to skip all modifications to the index name 

420 :param page_size: The number of results to return per page. 

421 :param ip: The user's hashed IP. Hashed IPs are used to anonymously but 

422 uniquely identify users exclusively for ensuring query consistency across 

423 Elasticsearch shards. 

424 :param filter_dead: Whether dead links should be removed. 

425 :param page: The results page number. 

426 :return: Tuple with a list of Hits from elasticsearch, the total count of 

427 pages, the number of results, and the ``SearchContext`` as a dict. 

428 """ 

429 index = get_index(exact_index, origin_index, search_params) 

430 

431 strategy: SearchStrategy = ( 

432 "collection" if search_params.validated_data.get("collection") else "search" 

433 ) 

434 

435 query = query_builders[strategy](search_params) 

436 

437 s = Search(index=index).query(query) 

438 

439 if strategy == "search": 439 ↛ 449line 439 didn't jump to line 449 because the condition on line 439 was always true

440 # Use highlighting to determine which fields contribute to the selection of 

441 # top results. 

442 s = s.highlight(*DEFAULT_SEARCH_FIELDS) 

443 s = s.highlight_options(order="score") 

444 s.extra(track_scores=True) 

445 

446 # Route users to the same Elasticsearch worker node to reduce 

447 # pagination inconsistencies and increase cache hits. 

448 # TODO: Re-add 7s request_timeout when ES stability is restored 

449 s = s.params(preference=str(ip)) 

450 

451 # Sort by `created_on` if the parameter is set or if `strategy` is `collection`. 

452 sort_by = search_params.validated_data.get("sort_by") 

453 if strategy == "collection" or sort_by == INDEXED_ON: 

454 sort_dir = search_params.validated_data.get("sort_dir", "desc") 

455 s = s.sort({"created_on": {"order": sort_dir}}) 

456 

457 # Execute paginated search and tally results 

458 page_count, result_count, results = execute_search( 

459 s, page, page_size, filter_dead, index, es_query=strategy 

460 ) 

461 

462 result_ids = [result.identifier for result in results] 

463 search_context = SearchContext.build(result_ids, origin_index) 

464 

465 return results, page_count, result_count, search_context.asdict() 

466 

467 

468def tally_results( 

469 index: SearchIndex, results: list[Hit] | None, page: int, page_size: int 

470) -> None: 

471 """ 

472 Tally the number of the results from each provider in the results 

473 for the search query. 

474 """ 

475 results_to_tally = results or [] 

476 max_result_depth = page * page_size 

477 if max_result_depth <= 80: 477 ↛ 480line 477 didn't jump to line 480 because the condition on line 477 was always true

478 # Applies when `page_size * page` could land "evenly" on 80 

479 should_tally = True 

480 elif max_result_depth - page_size < 80: 

481 # Applies when `page_size * page` could land beyond 80, but still 

482 # encompass some results on _this page_ that are at or below the 80th 

483 # position. For example: page=7 page_size=12 result depth=84. 

484 # While max_result_depth exceeds 80, we still want to count 

485 # the first eight results in `results` that are below or at the 80th 

486 # position for the query. 

487 should_tally = True 

488 results_to_tally = results_to_tally[: 80 - (max_result_depth - page_size)] 

489 else: 

490 should_tally = False 

491 

492 if results and should_tally: 

493 # We ignore tallies for deep results because they're not likely to 

494 # be as important for search relevancy for most users at this point 

495 # 80 is chosen because it represents the first four pages of the 

496 # default page count of 20 (20 * 4) which is how our own frontend 

497 # makes requests and displays results. Because that is the only 

498 # place we can actually conceivably measure relevancy down the 

499 # line, it is the only sensible, controlled space we can use to 

500 # check things like provider density for a set of queries. 

501 tallies.count_provider_occurrences(results_to_tally, index) 

502 

503 

504def execute_search( 

505 s: Search, 

506 page: int, 

507 page_size: int, 

508 filter_dead: bool, 

509 index: SearchIndex, 

510 es_query: str, 

511) -> tuple[int, int, list[Hit]]: 

512 """ 

513 Execute search for the given query slice, post-processes the results, 

514 and returns the results and result and page counts. 

515 """ 

516 start, end = get_query_slice(s, page_size, page, filter_dead) 

517 s = s[start:end] 

518 

519 search_response = get_es_response(s, es_query=es_query) 

520 

521 results: list[Hit] = ( 

522 _post_process_results(s, start, end, page_size, search_response, filter_dead) 

523 or [] 

524 ) 

525 result_count, page_count = _get_result_and_page_count( 

526 search_response, results, page_size, page 

527 ) 

528 tally_results(index, results, page, page_size) 

529 return page_count, result_count, results 

530 

531 

532def get_sources(index): 

533 """ 

534 Given an index, find all available data sources and return their counts. 

535 

536 :param index: An Elasticsearch index, such as `'image'`. 

537 :return: A dictionary mapping sources to the count of their images.` 

538 """ 

539 source_cache_name = "sources-" + index 

540 try: 

541 sources = cache.get(key=source_cache_name) 

542 except ConnectionError: 

543 logger.warning("Redis connect failed, cannot get cached sources.") 

544 sources = None 

545 

546 if not sources: 

547 # Don't increase `size` without reading this issue first: 

548 # https://github.com/elastic/elasticsearch/issues/18838 

549 size = 100 

550 body = { 

551 "size": 0, 

552 "aggs": { 

553 "unique_sources": { 

554 "terms": { 

555 "field": "source", 

556 "size": size, 

557 "order": {"_key": "desc"}, 

558 } 

559 }, 

560 }, 

561 } 

562 try: 

563 results = get_raw_es_response( 

564 index=index, 

565 body=body, 

566 request_cache=True, 

567 es_query="sources", 

568 ) 

569 buckets = results["aggregations"]["unique_sources"]["buckets"] 

570 except NotFoundError: 

571 buckets = [{"key": "none_found", "doc_count": 0}] 

572 sources = {bucket["key"]: bucket["doc_count"] for bucket in buckets} 

573 

574 try: 

575 cache.set( 

576 key=source_cache_name, 

577 timeout=SOURCE_CACHE_TIMEOUT, 

578 value=sources, 

579 ) 

580 except ConnectionError: 

581 logger.warning("Redis connect failed, cannot cache sources.") 

582 

583 sources = {source: int(count) for source, count in sources.items()} 

584 return sources 

585 

586 

587def _get_result_and_page_count( 

588 response_obj: Response, results: list[Hit] | None, page_size: int, page: int 

589) -> tuple[int, int]: 

590 """ 

591 Adjust page count because ES disallows deep pagination of ranked queries. 

592 

593 :param response_obj: The original Elasticsearch response object. 

594 :param results: The list of filtered result Hits. 

595 :return: Result and page count. 

596 """ 

597 if not results: 

598 return 0, 0 

599 

600 result_count = response_obj.hits.total.value 

601 page_count = ceil(result_count / page_size) 

602 

603 if len(results) < page_size: 

604 if page_count == 1: 604 ↛ 613line 604 didn't jump to line 613 because the condition on line 604 was always true

605 result_count = len(results) 

606 

607 # If we have fewer results than the requested page size and are 

608 # not on the first page that means that we've reached the end of 

609 # the query and can set the page_count to the currently requested 

610 # page. This means that the `page_count` can change during 

611 # pagination for the same query, but it's the only way to 

612 # communicate under the current v1 API that a query has been exhausted. 

613 page_count = page 

614 

615 return result_count, page_count