diff --git a/edxsearch/__init__.py b/edxsearch/__init__.py index e64c0723..c20df098 100644 --- a/edxsearch/__init__.py +++ b/edxsearch/__init__.py @@ -1,3 +1,3 @@ """ Container module for testing / demoing search """ -__version__ = '5.0.1' +__version__ = '5.0.2' diff --git a/search/api.py b/search/api.py index 87586cd1..4a1029ed 100644 --- a/search/api.py +++ b/search/api.py @@ -143,6 +143,7 @@ def course_discovery_search( field: search_fields[field] for field in search_fields if field in use_search_fields } + allowed_orgs = use_field_dictionary.copy().get("org") if field_dictionary: use_field_dictionary.update(field_dictionary) if enable_course_sorting_by_start_date: @@ -168,6 +169,7 @@ def course_discovery_search( aggregation_terms=course_discovery_aggregations(), sort_by=sort_by, is_multivalue=is_multivalue, + allowed_orgs=allowed_orgs, ) return results diff --git a/search/elastic.py b/search/elastic.py index 4607abd7..e9d03c80 100644 --- a/search/elastic.py +++ b/search/elastic.py @@ -718,6 +718,9 @@ def search(self, body = {"query": query} + # Strip "allowed_orgs" out of kwargs; it's not a valid Elasticsearch search parameter. + kwargs.pop("allowed_orgs", None) + is_multivalue = kwargs.pop("is_multivalue", False) if aggregation_terms: if is_multivalue: diff --git a/search/meilisearch.py b/search/meilisearch.py index 5264d74c..3e0999ab 100644 --- a/search/meilisearch.py +++ b/search/meilisearch.py @@ -181,6 +181,7 @@ def search( See meilisearch docs: https://www.meilisearch.com/docs/reference/api/search """ is_multivalue = kwargs.pop("is_multivalue", False) + allowed_orgs = kwargs.pop("allowed_orgs", []) opt_params = get_search_params( field_dictionary=field_dictionary, filter_dictionary=filter_dictionary, @@ -195,13 +196,20 @@ def search( meilisearch_results = self.meilisearch_index.search(query_string, opt_params) if is_multivalue: - self._expand_facet_distibutions(field_dictionary, query_string, opt_params, meilisearch_results) + self._expand_facet_distibutions( + field_dictionary, + allowed_orgs, + query_string, + opt_params, + meilisearch_results, + ) return process_results(meilisearch_results, self.index_name) def _expand_facet_distibutions( self, field_dictionary: dict, + allowed_orgs: list, query_string: str, opt_params: dict, meilisearch_results: dict @@ -212,12 +220,19 @@ def _expand_facet_distibutions( for facet in field_dictionary.keys(): expanded_facet_distribution = self._get_expanded_distribution( query_string, + allowed_orgs, facet, opt_params.get("filter", []), ) meilisearch_results.setdefault("facetDistribution", {})[facet] = expanded_facet_distribution - def _get_expanded_distribution(self, query: str, facet_to_exclude: str, filter_rules: list) -> dict: + def _get_expanded_distribution( + self, + query: str, + allowed_orgs: list, + facet_to_exclude: str, + filter_rules: list + ) -> dict: """ Run a secondary query excluding one facet to get its full distribution. Only return distribution data, without any actual results. @@ -228,7 +243,13 @@ def _get_expanded_distribution(self, query: str, facet_to_exclude: str, filter_r 'limit': 0, } result = self.meilisearch_index.search(query, secondary_opt_params) - return result.get("facetDistribution", {}).get(facet_to_exclude, {}) + facet_distribution = result.get("facetDistribution", {}).get(facet_to_exclude, {}) + if facet_to_exclude == "org" and allowed_orgs: + allowed = set(allowed_orgs) + facet_distribution = { + org: count for org, count in facet_distribution.items() if org in allowed + } + return facet_distribution def remove(self, doc_ids, **kwargs): """ diff --git a/search/tests/test_meilisearch.py b/search/tests/test_meilisearch.py index fc2a45dd..5e520627 100644 --- a/search/tests/test_meilisearch.py +++ b/search/tests/test_meilisearch.py @@ -459,8 +459,9 @@ def test_multivalue_search_expands_selected_facet_without_filtering(self): 'org = "EDX"', ] selected_facet = 'language' + allowed_orgs = ["EDX"] actual_distribution = engine._get_expanded_distribution( # pylint: disable=protected-access - '', selected_facet, original_filter + '', allowed_orgs, selected_facet, original_filter ) self.assertDictEqual(actual_distribution, multivalue_distribution) (query, opt_params), _ = engine.meilisearch_index.search.call_args # pylint: disable=unused-variable @@ -505,6 +506,26 @@ def test_multivalue_search_merges_expanded_facet_distributions(self): ) self.assertDictEqual(aggregations["org"]["terms"], {"EDX": 2}) + def test_get_expanded_distribution_filters_orgs_to_allowed(self): + engine = search.meilisearch.MeilisearchEngine(index='test_index') + engine.meilisearch_index.search = Mock(return_value={ + "facetDistribution": { + "org": {"EDX": 2, "MITx": 1, "HarvardX": 3}, # More orgs than allowed + } + }) + + allowed_orgs = ["EDX", "MITx"] + facet_distribution = engine._get_expanded_distribution( # pylint: disable=protected-access + query="", + allowed_orgs=allowed_orgs, + facet_to_exclude="org", + filter_rules=[], + ) + + # Only the allowed orgs are returned; HarvardX is filtered out + self.assertDictEqual(facet_distribution, {"EDX": 2, "MITx": 1}) + self.assertEqual(set(facet_distribution.keys()), set(allowed_orgs)) + def test_single_value_search_narrows_selected_facet(self): engine = search.meilisearch.MeilisearchEngine(index='test_index') engine.meilisearch_index.search = Mock(side_effect=[