Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion edxsearch/__init__.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,3 @@
""" Container module for testing / demoing search """

__version__ = '5.0.1'
__version__ = '5.0.2'
2 changes: 2 additions & 0 deletions search/api.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand All @@ -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
3 changes: 3 additions & 0 deletions search/elastic.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
27 changes: 24 additions & 3 deletions search/meilisearch.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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
Expand All @@ -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.
Expand All @@ -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):
"""
Expand Down
23 changes: 22 additions & 1 deletion search/tests/test_meilisearch.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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=[
Expand Down
Loading