Skip to content
Draft
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
1 change: 0 additions & 1 deletion pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -388,7 +388,6 @@ ignore_missing_imports = true
# - python3 -m tools.mypy_helpers.find_easiest_modules
[[tool.mypy.overrides]]
module = [
"sentry.snuba.metrics.query_builder",
"sentry.testutils.cases",
]
disable_error_code = [
Expand Down
117 changes: 78 additions & 39 deletions src/sentry/snuba/metrics/query_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,7 +3,7 @@
from collections.abc import Mapping, Sequence
from datetime import datetime, timedelta
from enum import Enum
from typing import Any, TypedDict, overload
from typing import Any, TypedDict, cast, overload

import sentry_sdk
from snuba_sdk import (
Expand Down Expand Up @@ -43,6 +43,7 @@
from sentry.snuba.metrics.fields.base import (
COMPOSITE_ENTITY_CONSTITUENT_ALIAS,
MetricExpressionBase,
MetricOperationParams,
generate_bottom_up_dependency_tree_for_metrics,
org_id_from_projects,
)
Expand Down Expand Up @@ -71,6 +72,7 @@
DerivedMetricParseException,
MetricDoesNotExistException,
MetricEntity,
MetricOperationType,
get_num_intervals,
get_timestamp_column_name,
require_rhs_condition_resolution,
Expand Down Expand Up @@ -128,15 +130,24 @@ def parse_public_field(field: str) -> MetricField:
matches = PUBLIC_EXPRESSION_REGEX.match(field)

if matches is not None:
operation = matches[1]
metric_name = matches[2]
operation = cast(MetricOperationType | None, matches[1])
metric_name = cast(str, matches[2])
else:
operation = None
metric_name = field

return MetricField(operation, get_mri(metric_name))


def _extract_scalar_metric_params(
params: dict[str, None | str | int | float | Sequence[tuple[str | int, ...]]] | None,
) -> MetricOperationParams | None:
# Runtime metric operations may legitimately rely on richer values (for example tuple sequences),
# but MetricOperationParams is currently typed as scalar-only. Keep runtime behavior unchanged and
# narrow only for static type checking.
return cast(MetricOperationParams | None, params)


def transform_null_transaction_to_unparameterized(use_case_id, org_id, alias=None):
"""
This function transforms any null tag.transaction to '<< unparameterized >>' so that it can be handled
Expand Down Expand Up @@ -747,11 +758,14 @@ def translate_meta_results(
continue
elif alias_type == AliasMetaType.GROUP_BY_METRIC_FIELD:
metric_groupby_field = alias_to_metric_group_by_field[record["name"]]
assert isinstance(metric_groupby_field.field, MetricField)
defined_parent_meta_type = get_metric_object_from_metric_field(
metric_groupby_field.field
).get_meta_type()

record["type"] = defined_parent_meta_type
record["type"] = (
record["type"] if defined_parent_meta_type is None else defined_parent_meta_type
)
elif alias_type == AliasMetaType.TAG:
record["type"] = "string"
elif alias_type == AliasMetaType.DATASET_COLUMN or alias_type == AliasMetaType.TIME_COLUMN:
Expand Down Expand Up @@ -812,7 +826,7 @@ def generate_snql_for_action_by_fields(

@staticmethod
def generate_snql_for_action_by_fields(
metric_action_by_field: MetricActionByField,
metric_action_by_field: MetricGroupByField | MetricOrderByField,
use_case_id: UseCaseID,
org_id: int,
projects: Sequence[Project],
Expand All @@ -824,16 +838,24 @@ def generate_snql_for_action_by_fields(
the snql generation starts to diverge significantly.
"""

is_group_by = isinstance(metric_action_by_field, MetricGroupByField)
is_order_by = isinstance(metric_action_by_field, MetricOrderByField)
if not is_group_by and not is_order_by:
group_by_field = (
metric_action_by_field
if isinstance(metric_action_by_field, MetricGroupByField)
else None
)
order_by_field = (
metric_action_by_field
if isinstance(metric_action_by_field, MetricOrderByField)
else None
)
if group_by_field is None and order_by_field is None:
raise InvalidParams("The metric action must either be an order by or group by.")

if isinstance(metric_action_by_field.field, str):
# This transformation is currently supported only for group by because OrderBy doesn't support the Function type.
if is_group_by and metric_action_by_field.field == "transaction":
if group_by_field is not None and metric_action_by_field.field == "transaction":
return transform_null_transaction_to_unparameterized(
use_case_id, org_id, metric_action_by_field.alias
use_case_id, org_id, group_by_field.alias
)

# Handles the case when we are trying to group or order by `project` for example, but we want
Expand All @@ -845,7 +867,7 @@ def generate_snql_for_action_by_fields(
else:
# The support for tags in the order by is disabled for now because there is no need to have it. If the
# need arise, we will implement it.
if is_group_by:
if group_by_field is not None:
assert isinstance(metric_action_by_field.field, str)
column_name = resolve_tag_key(use_case_id, org_id, metric_action_by_field.field)
else:
Expand All @@ -856,38 +878,39 @@ def generate_snql_for_action_by_fields(
exp = (
AliasedExpression(
exp=Column(name=column_name),
alias=metric_action_by_field.alias,
alias=group_by_field.alias,
)
if is_group_by and not is_column
if group_by_field is not None and not is_column
else Column(name=column_name)
)

if is_order_by:
if order_by_field is not None:
# We return a list in order to use the "extend" method and reduce the number of changes across
# the codebase.
exp = [OrderBy(exp=exp, direction=metric_action_by_field.direction)]
exp = [OrderBy(exp=exp, direction=order_by_field.direction)]

return exp
elif isinstance(metric_action_by_field.field, MetricField):
try:
metric_expression = metric_object_factory(
metric_action_by_field.field.op, metric_action_by_field.field.metric_mri
)
params = _extract_scalar_metric_params(metric_action_by_field.field.params)

if is_group_by:
if group_by_field is not None:
return metric_expression.generate_groupby_statements(
use_case_id=use_case_id,
alias=metric_action_by_field.field.alias,
params=metric_action_by_field.field.params,
params=params,
projects=projects,
)[0]
elif is_order_by:
elif order_by_field is not None:
return metric_expression.generate_orderby_clause(
use_case_id=use_case_id,
alias=metric_action_by_field.field.alias,
params=metric_action_by_field.field.params,
params=params,
projects=projects,
direction=metric_action_by_field.direction,
direction=order_by_field.direction,
)
else:
raise NotImplementedError(
Expand All @@ -898,7 +921,7 @@ def generate_snql_for_action_by_fields(
raise InvalidParams(f"Cannot resolve {metric_action_by_field.field} into SnQL")
else:
raise NotImplementedError(
f"Unsupported {'group by' if is_group_by else 'order by' if is_order_by else 'None'} field: {metric_action_by_field.field} needs to be either a MetricField or a string"
f"Unsupported {'group by' if group_by_field else 'order by' if order_by_field else 'None'} field: {metric_action_by_field.field} needs to be either a MetricField or a string"
)

def _build_where(self) -> list[BooleanCondition | Condition]:
Expand All @@ -921,20 +944,25 @@ def _build_where(self) -> list[BooleanCondition | Condition]:
condition.lhs.op, condition.lhs.metric_mri
)
try:
rhs = condition.rhs
if condition.lhs.op is not None and require_rhs_condition_resolution(
condition.lhs.op
):
if not isinstance(condition.rhs, str):
raise InvalidParams(
f"Cannot resolve non-string RHS for metric condition {condition.lhs}"
)
rhs = resolve_tag_value(self._use_case_id, self._org_id, condition.rhs)
metric_condition_filters.append(
Condition(
lhs=metric_expression.generate_where_statements(
use_case_id=self._use_case_id,
params=condition.lhs.params,
params=_extract_scalar_metric_params(condition.lhs.params),
projects=self._projects,
alias=condition.lhs.alias,
)[0],
op=condition.op,
rhs=(
resolve_tag_value(self._use_case_id, self._org_id, condition.rhs)
if require_rhs_condition_resolution(condition.lhs.op)
else condition.rhs
),
rhs=rhs,
)
)
except IndexError:
Expand Down Expand Up @@ -1069,6 +1097,8 @@ def __build_totals_and_series_queries(
series_limit = self._metrics_query.max_limit

if self._use_case_id in [UseCaseID.TRANSACTIONS, UseCaseID.SPANS]:
if self._metrics_query.interval is None:
raise InvalidParams("Interval is required for discover metrics series queries")
time_groupby_column = self.__generate_time_groupby_column_for_discover_queries(
self._metrics_query.interval
)
Expand Down Expand Up @@ -1099,17 +1129,19 @@ def __generate_time_groupby_column_for_discover_queries(interval: int) -> Functi

def __update_query_dicts_with_component_entities(
self,
component_entities: dict[MetricEntity, Sequence[str]],
metric_mri_to_obj_dict: dict[tuple[str | None, str, str], MetricExpressionBase],
fields_in_entities: dict[MetricEntity, list[tuple[str | None, str, str]]],
component_entities: Mapping[MetricEntity | None, Sequence[str]],
metric_mri_to_obj_dict: dict[tuple[MetricOperationType | None, str, str], MetricExpressionBase],
fields_in_entities: dict[MetricEntity, list[tuple[MetricOperationType | None, str, str]]],
parent_alias,
) -> dict[tuple[str | None, str, str], MetricExpressionBase]:
) -> dict[tuple[MetricOperationType | None, str, str], MetricExpressionBase]:
# At this point in time, we are only supporting raw metrics in the metrics attribute of
# any instance of DerivedMetric, and so in this case the op will always be None
# ToDo(ahmed): In future PR, we might want to allow for dependency metrics to also have an
# an aggregate and in this case, we would need to parse the op here
op = None
op: MetricOperationType | None = None
for entity, metric_mris in component_entities.items():
if entity is None:
continue
for metric_mri in metric_mris:
# The constituents of an instance of CompositeEntityDerivedMetric will have a reference to their parent
# alias so that we are able to distinguish the constituents in case we have naming collisions that could
Expand All @@ -1128,8 +1160,10 @@ def __update_query_dicts_with_component_entities(
return metric_mri_to_obj_dict

def get_snuba_queries(self):
metric_mri_to_obj_dict: dict[tuple[str | None, str, str], MetricExpressionBase] = {}
fields_in_entities: dict[MetricEntity, list[tuple[str | None, str, str]]] = {}
metric_mri_to_obj_dict: dict[
tuple[MetricOperationType | None, str, str], MetricExpressionBase
] = {}
fields_in_entities: dict[MetricEntity, list[tuple[MetricOperationType | None, str, str]]] = {}

for select_field in self._metrics_query.select:
metric_field_obj = metric_object_factory(select_field.op, select_field.metric_mri)
Expand Down Expand Up @@ -1190,13 +1224,14 @@ def get_snuba_queries(self):
for field in fields:
metric_field_obj = metric_mri_to_obj_dict[field]
try:
params = self._alias_to_metric_field[field[2]].params
params = _extract_scalar_metric_params(self._alias_to_metric_field[field[2]].params)
except KeyError:
params = None

# In order to support on demand metrics which require an interval (e.g. epm),
# we want to pass the interval down via params so we can pass it to the associated snql_factory
params = {"interval": self._metrics_query.interval, **(params or {})}
if self._metrics_query.interval is not None:
params = {"interval": self._metrics_query.interval, **(params or {})}
select += metric_field_obj.generate_select_statements(
projects=self._projects,
use_case_id=self._use_case_id,
Expand Down Expand Up @@ -1255,7 +1290,7 @@ def __init__(
self,
organization_id: int,
metrics_query: DeprecatingMetricsQuery,
fields_in_entities: dict[MetricEntity, list[tuple[str | None, str, str]]],
fields_in_entities: dict[MetricEntity, list[tuple[MetricOperationType | None, str, str]]],
intervals: list[datetime],
results,
use_case_id: UseCaseID,
Expand Down Expand Up @@ -1432,7 +1467,9 @@ def resolve_tag_value(value: int | str | None) -> str | None:
metric_obj = metric_object_factory(op=op, metric_mri=metric_mri)
if totals is not None:
try:
params = self._alias_to_metric_field[alias].params
params = _extract_scalar_metric_params(
self._alias_to_metric_field[alias].params
)
except KeyError:
params = None
totals[alias] = metric_obj.run_post_query_function(
Expand All @@ -1447,7 +1484,9 @@ def resolve_tag_value(value: int | str | None) -> str | None:
[metric_obj.generate_default_null_values()] * len(self._intervals),
)
try:
params = self._alias_to_metric_field[alias].params
params = _extract_scalar_metric_params(
self._alias_to_metric_field[alias].params
)
except KeyError:
params = None
series[alias][idx] = metric_obj.run_post_query_function(
Expand Down
Loading