diff --git a/pyproject.toml b/pyproject.toml index e6111078b337..b561af879a83 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -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 = [ diff --git a/src/sentry/snuba/metrics/query_builder.py b/src/sentry/snuba/metrics/query_builder.py index 066d553cc6f3..236aa0cf097f 100644 --- a/src/sentry/snuba/metrics/query_builder.py +++ b/src/sentry/snuba/metrics/query_builder.py @@ -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 ( @@ -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, ) @@ -71,6 +72,7 @@ DerivedMetricParseException, MetricDoesNotExistException, MetricEntity, + MetricOperationType, get_num_intervals, get_timestamp_column_name, require_rhs_condition_resolution, @@ -128,8 +130,8 @@ 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 @@ -137,6 +139,15 @@ def parse_public_field(field: str) -> MetricField: 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 @@ -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: @@ -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], @@ -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 @@ -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: @@ -856,16 +878,16 @@ 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): @@ -873,21 +895,22 @@ def generate_snql_for_action_by_fields( 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( @@ -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]: @@ -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: @@ -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 ) @@ -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 @@ -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) @@ -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, @@ -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, @@ -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( @@ -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(