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
59 changes: 59 additions & 0 deletions solrorbit/workload/loader.py
Original file line number Diff line number Diff line change
Expand Up @@ -1008,6 +1008,11 @@ def check_one_of_each_name_present(self, obj):
DEFAULT_N = 5000
DEFAULT_ALPHA = 1
DEFAULT_QUERY_RANDOMIZATION_INFO = QueryRandomizationInfo("range", [["gte", "gt"], ["lte", "lt"]], ["format"])
SOLR_RANGE_TERM_PATTERN = re.compile(r"(?P<field>[A-Za-z_][A-Za-z0-9_.]*):"
r"(?P<lower_bracket>[\[{])(?P<lower>[^\s\[\]{}]+)"
r"\s+TO\s+"
r"(?P<upper>[^\s\[\]{}]+)(?P<upper_bracket>[\]}])")

def __init__(self, cfg):
self.randomization_enabled = cfg.opts("workload", "randomization.enabled", mandatory=False, default_value=False)
self.rf = float(cfg.opts("workload", "randomization.repeat_frequency", mandatory=False, default_value=self.DEFAULT_RF))
Expand Down Expand Up @@ -1076,6 +1081,56 @@ def extract_fields_helper(self, root, current_path, query_randomization_info):
# leaf node
return []

def solr_string_paths(self, body):
if isinstance(body.get("query"), str):
yield ("query",)
filters = body.get("filter")
if isinstance(filters, str):
yield ("filter",)
elif isinstance(filters, list):
for i, filter_clause in enumerate(filters):
if isinstance(filter_clause, str):
yield ("filter", i)

def solr_range_term_states_both_bounds(self, match):
return match.group("lower") != "*" and match.group("upper") != "*"

def extract_solr_range_terms(self, body):
fields_and_paths = []
for path in self.solr_string_paths(body):
for match in self.SOLR_RANGE_TERM_PATTERN.finditer(self.get_dict_from_previous_path(body, path)):
if self.solr_range_term_states_both_bounds(match):
fields_and_paths.append((match.group("field"), path))
return fields_and_paths

def set_solr_range_terms(self, params, fields_and_paths, new_values, query_randomization_info):
bound_names = [parameter_name_options[0] for parameter_name_options in query_randomization_info.parameter_name_options_list]
if len(bound_names) != 2:
return params
lower_name, upper_name = bound_names
new_values_by_path = {}
for field_and_path, new_value in zip(fields_and_paths, new_values):
new_values_by_path.setdefault(field_and_path[1], []).append(new_value)

for path, path_new_values in new_values_by_path.items():
remaining = iter(path_new_values)

def replace(match, remaining=remaining):
if not self.solr_range_term_states_both_bounds(match):
return match.group(0)
new_value = next(remaining, None)
if new_value is None:
return match.group(0)
return "{}:{}{} TO {}{}".format(match.group("field"),
match.group("lower_bracket"),
new_value[lower_name],
new_value[upper_name],
match.group("upper_bracket"))

parent = self.get_dict_from_previous_path(params["body"], path[:-1])
parent[path[-1]] = self.SOLR_RANGE_TERM_PATTERN.sub(replace, parent[path[-1]])
return params

def extract_fields_and_paths(self, params, query_randomization_info):
# Search for fields used in range queries, and the paths to those fields
# Return pairs of (field, path_to_field)
Expand All @@ -1087,11 +1142,15 @@ def extract_fields_and_paths(self, params, query_randomization_info):
raise exceptions.SystemSetupError(
f"Cannot extract range query fields from these params: {params}\n, missing params[\"body\"][\"query\"]\n"
f"Make sure the operation in operations/default.json is well-formed")
if isinstance(root, str):
return self.extract_solr_range_terms(params["body"])
fields_and_paths = self.extract_fields_helper(root, [], query_randomization_info)
return fields_and_paths

def set_range(self, params, fields_and_paths, new_values, query_randomization_info):
assert len(fields_and_paths) == len(new_values)
if isinstance(params["body"].get("query"), str):
return self.set_solr_range_terms(params, fields_and_paths, new_values, query_randomization_info)
for field_and_path, new_value in zip(fields_and_paths, new_values):
field = field_and_path[0]
path = field_and_path[1]
Expand Down
104 changes: 104 additions & 0 deletions tests/workload/loader_test.py
Original file line number Diff line number Diff line change
Expand Up @@ -2159,6 +2159,110 @@ def test_range_finding_function(self):
geo_point_expected = [("location", ["geo_bounding_box"])]
self.assertEqual(geo_point_result, geo_point_expected)

def test_range_finding_function_for_string_queries(self):
cfg = config.Config()
processor = loader.QueryRandomizerWorkloadProcessor(cfg)
default_info = loader.QueryRandomizerWorkloadProcessor.DEFAULT_QUERY_RANDOMIZATION_INFO

query_range = {
"name": "range",
"operation-type": "search",
"body": {
"query": "total_amount:[5 TO 15}"
}
}
self.assertEqual(processor.extract_fields_and_paths(query_range, default_info),
[("total_amount", ("query",))])

filter_range = {
"name": "distance_amount_facet",
"operation-type": "search",
"body": {
"query": "*:*",
"filter": ["trip_distance:[0 TO 50}"],
"limit": 0
}
}
self.assertEqual(processor.extract_fields_and_paths(filter_range, default_info),
[("trip_distance", ("filter", 0))])

several_terms = {
"name": "several",
"operation-type": "search",
"body": {
"query": "*:*",
"filter": ["trip_distance:[0 TO 50} AND total_amount:{5 TO 100]", "passenger_count:2"]
}
}
self.assertEqual(processor.extract_fields_and_paths(several_terms, default_info),
[("trip_distance", ("filter", 0)), ("total_amount", ("filter", 0))])

no_range = {"name": "match-all", "operation-type": "search", "body": {"query": "*:*"}}
self.assertEqual(processor.extract_fields_and_paths(no_range, default_info), [])

def test_set_range_keeps_the_brackets_of_a_string_query(self):
cfg = config.Config()
processor = loader.QueryRandomizerWorkloadProcessor(cfg)
default_info = loader.QueryRandomizerWorkloadProcessor.DEFAULT_QUERY_RANDOMIZATION_INFO
params = {
"body": {
"query": "*:*",
"filter": ["trip_distance:[0 TO 50} AND total_amount:{5 TO 100]"]
}
}
fields_and_paths = processor.extract_fields_and_paths(params, default_info)
result = processor.set_range(params, fields_and_paths,
[{"gte": 3, "lte": 7}, {"gte": 10.5, "lte": 20.25}], default_info)
self.assertEqual(result["body"]["filter"][0], "trip_distance:[3 TO 7} AND total_amount:{10.5 TO 20.25]")
self.assertEqual(result["body"]["query"], "*:*")

def test_get_randomized_values_for_string_queries(self):
cfg = config.Config()
cfg.add(config.Scope.application, "workload", "randomization.repeat_frequency", 0.0)
processor = loader.QueryRandomizerWorkloadProcessor(cfg)
default_info = loader.QueryRandomizerWorkloadProcessor.DEFAULT_QUERY_RANDOMIZATION_INFO
new_value = {"gte": "2015-01-05T00:00:00Z", "lte": "2015-01-09T00:00:00Z", "format": "yyyy-MM-dd"}
params = {
"index": "nyc_taxis",
"body": {
"query": "dropoff_datetime:[2015-01-01T00:00:00Z TO 2015-01-22T00:00:00Z}",
"limit": 0
}
}
result = processor.get_randomized_values(None, params, default_info,
op_name="date_histogram_facet",
get_standard_value=lambda op_name, field, index: new_value,
get_standard_value_source=lambda op_name, field: lambda: new_value)
self.assertEqual(result["body"]["query"], "dropoff_datetime:[2015-01-05T00:00:00Z TO 2015-01-09T00:00:00Z}")
self.assertEqual(result["body"]["limit"], 0)

def test_a_string_range_that_leaves_a_bound_open_is_not_randomized(self):
cfg = config.Config()
processor = loader.QueryRandomizerWorkloadProcessor(cfg)
default_info = loader.QueryRandomizerWorkloadProcessor.DEFAULT_QUERY_RANDOMIZATION_INFO
for query in ("total_amount:[* TO 15}", "total_amount:[5 TO *]", "total_amount:[* TO *]"):
params = {"body": {"query": query}}
self.assertEqual(processor.extract_fields_and_paths(params, default_info), [])
self.assertEqual(processor.set_range(params, [], [], default_info)["body"]["query"], query)

def test_a_term_that_is_not_randomized_does_not_take_another_terms_value(self):
cfg = config.Config()
processor = loader.QueryRandomizerWorkloadProcessor(cfg)
default_info = loader.QueryRandomizerWorkloadProcessor.DEFAULT_QUERY_RANDOMIZATION_INFO
params = {"body": {"query": "total_amount:[* TO 15} AND trip_distance:[1 TO 9}"}}
fields_and_paths = processor.extract_fields_and_paths(params, default_info)
self.assertEqual(fields_and_paths, [("trip_distance", ("query",))])
result = processor.set_range(params, fields_and_paths, [{"gte": 3, "lte": 7}], default_info)
self.assertEqual(result["body"]["query"], "total_amount:[* TO 15} AND trip_distance:[3 TO 7}")

def test_a_value_source_that_omits_a_bound_is_not_silently_ignored(self):
cfg = config.Config()
processor = loader.QueryRandomizerWorkloadProcessor(cfg)
default_info = loader.QueryRandomizerWorkloadProcessor.DEFAULT_QUERY_RANDOMIZATION_INFO
params = {"body": {"query": "total_amount:[5 TO 15}"}}
fields_and_paths = processor.extract_fields_and_paths(params, default_info)
with self.assertRaises(KeyError):
processor.set_range(params, fields_and_paths, [{"gt": 3, "lt": 7}], default_info)

def test_get_randomized_values(self):
helper = self.StandardValueHelper()
Expand Down