diff --git a/pyproject.toml b/pyproject.toml index fb62fbc4bf51..d1bc768e62ff 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..6647ca36cca7 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, @@ -127,8 +129,9 @@ def parse_field(field: str, allow_mri: bool = False) -> MetricField: def parse_public_field(field: str) -> MetricField: matches = PUBLIC_EXPRESSION_REGEX.match(field) + operation: MetricOperationType | None if matches is not None: - operation = matches[1] + operation = cast(MetricOperationType, matches[1]) metric_name = matches[2] else: operation = None @@ -179,6 +182,12 @@ def _refers_to_column(expression: Column | Function) -> str | None: FUNCTION_ALLOWLIST = ("and", "or", "equals", "in", "tuple", "has", "match", "team_key_transaction") +def _coerce_metric_operation_params( + params: dict[str, None | str | int | float | Sequence[tuple[str | int, ...]]] | None, +) -> MetricOperationParams | None: + return cast(MetricOperationParams | None, params) + + def resolve_tags( use_case_id: UseCaseID, org_id: int, @@ -747,11 +756,15 @@ def translate_meta_results( continue elif alias_type == AliasMetaType.GROUP_BY_METRIC_FIELD: metric_groupby_field = alias_to_metric_group_by_field[record["name"]] + if not isinstance(metric_groupby_field.field, MetricField): + raise InvalidParams(f"Group by field {record['name']} was not a metric field") + defined_parent_meta_type = get_metric_object_from_metric_field( metric_groupby_field.field ).get_meta_type() - record["type"] = defined_parent_meta_type + if defined_parent_meta_type is not None: + record["type"] = defined_parent_meta_type elif alias_type == AliasMetaType.TAG: record["type"] = "string" elif alias_type == AliasMetaType.DATASET_COLUMN or alias_type == AliasMetaType.TIME_COLUMN: @@ -831,7 +844,7 @@ def generate_snql_for_action_by_fields( 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 isinstance(metric_action_by_field, MetricGroupByField) and metric_action_by_field.field == "transaction": return transform_null_transaction_to_unparameterized( use_case_id, org_id, metric_action_by_field.alias ) @@ -853,39 +866,42 @@ def generate_snql_for_action_by_fields( f"Unsupported string field: {metric_action_by_field.field}" ) - exp = ( + snuba_expression: Column | AliasedExpression = ( AliasedExpression( exp=Column(name=column_name), alias=metric_action_by_field.alias, ) - if is_group_by and not is_column + if isinstance(metric_action_by_field, MetricGroupByField) and not is_column else Column(name=column_name) ) if is_order_by: + if not isinstance(metric_action_by_field, MetricOrderByField): + raise InvalidParams("Order by field is not typed as MetricOrderByField") # 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)] + return [OrderBy(exp=snuba_expression, direction=metric_action_by_field.direction)] - return exp + return snuba_expression 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 = _coerce_metric_operation_params(metric_action_by_field.field.params) if is_group_by: 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 isinstance(metric_action_by_field, MetricOrderByField): 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, ) @@ -921,20 +937,24 @@ def _build_where(self) -> list[BooleanCondition | Condition]: condition.lhs.op, condition.lhs.metric_mri ) try: + resolved_rhs: int | float | str = 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("Resolved metric condition rhs should be a string") + resolved_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=_coerce_metric_operation_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=resolved_rhs, ) ) except IndexError: @@ -1068,9 +1088,13 @@ def __build_totals_and_series_queries( if self._metrics_query.max_limit: series_limit = self._metrics_query.max_limit + interval = self._metrics_query.interval + if interval is None: + interval = self._metrics_query.granularity.granularity + if self._use_case_id in [UseCaseID.TRANSACTIONS, UseCaseID.SPANS]: time_groupby_column = self.__generate_time_groupby_column_for_discover_queries( - self._metrics_query.interval + interval ) else: time_groupby_column = Column(TS_COL_GROUP) @@ -1100,10 +1124,10 @@ 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]]], - parent_alias, - ) -> dict[tuple[str | None, str, str], MetricExpressionBase]: + 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: str, + ) -> 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 @@ -1128,8 +1152,8 @@ 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) @@ -1152,8 +1176,10 @@ def get_snuba_queries(self): # lists of metric_mris as values representing all the entities and # metric_mris combination that this metric_object is composed of, or rather # the instances of SingleEntityDerivedMetric that it is composed of + if any(component_entity is None for component_entity in component_entities): + raise DerivedMetricParseException("Entity parsed is in incorrect format") metric_mri_to_obj_dict = self.__update_query_dicts_with_component_entities( - component_entities=component_entities, + component_entities=cast(dict[MetricEntity, Sequence[str]], component_entities), metric_mri_to_obj_dict=metric_mri_to_obj_dict, fields_in_entities=fields_in_entities, parent_alias=select_field.alias, @@ -1201,7 +1227,7 @@ def get_snuba_queries(self): projects=self._projects, use_case_id=self._use_case_id, alias=field[2], - params=params, + params=cast(MetricOperationParams, params), ) metric_ids_set |= metric_field_obj.generate_metric_ids( self._projects, self._use_case_id @@ -1255,7 +1281,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, @@ -1436,7 +1462,7 @@ def resolve_tag_value(value: int | str | None) -> str | None: except KeyError: params = None totals[alias] = metric_obj.run_post_query_function( - totals, params=params, alias=alias + totals, params=_coerce_metric_operation_params(params), alias=alias ) if series is not None: @@ -1451,7 +1477,10 @@ def resolve_tag_value(value: int | str | None) -> str | None: except KeyError: params = None series[alias][idx] = metric_obj.run_post_query_function( - series, params=params, idx=idx, alias=alias + series, + params=_coerce_metric_operation_params(params), + idx=idx, + alias=alias, ) # Remove the extra fields added due to the constituent metrics that were added