diff --git a/pyproject.toml b/pyproject.toml index 02b42987e724..c57d4909e6a7 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -383,7 +383,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 = [ @@ -1315,7 +1314,6 @@ module = [ "sentry.snuba.issue_platform", "sentry.snuba.metrics.fields.*", "sentry.snuba.metrics.mqb_query_transformer", - "sentry.snuba.metrics.query_builder", "sentry.snuba.metrics_enhanced_performance", "sentry.snuba.metrics_performance", "sentry.snuba.outcomes", diff --git a/src/sentry/snuba/metrics/query_builder.py b/src/sentry/snuba/metrics/query_builder.py index 025565831bab..dc1b49f09ecc 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, @@ -134,10 +136,22 @@ def parse_public_field(field: str) -> MetricField: operation = None metric_name = field - return MetricField(operation, get_mri(metric_name)) + # `MetricField` validates supported operations downstream, but mypy requires + # us to narrow this to the operation literal union. + return MetricField(cast(MetricOperationType | None, operation), get_mri(metric_name)) -def transform_null_transaction_to_unparameterized(use_case_id, org_id, alias=None): +def _as_metric_operation_params( + params: ( + dict[str, None | str | int | float | Sequence[tuple[str | int, ...]]] | None + ), +) -> MetricOperationParams | None: + return cast(MetricOperationParams | None, params) + + +def transform_null_transaction_to_unparameterized( + use_case_id: UseCaseID, org_id: int, alias: str | None = None +) -> Function: """ This function transforms any null tag.transaction to '<< unparameterized >>' so that it can be handled as such in any query using that tag value. @@ -508,11 +522,11 @@ class QueryDefinition: def __init__( self, - projects, - query_params, + projects: Sequence[Project], + query_params: Any, allow_mri: bool = False, - paginator_kwargs: dict | None = None, - ): + paginator_kwargs: Mapping[str, int] | None = None, + ) -> None: self._projects = projects paginator_kwargs = paginator_kwargs or {} @@ -561,7 +575,7 @@ def to_metrics_query(self) -> DeprecatingMetricsQuery: ) @staticmethod - def _parse_orderby(query_params, allow_mri: bool = False): + def _parse_orderby(query_params: Any, allow_mri: bool = False) -> list[MetricsOrderBy] | None: orderbys = query_params.getlist("orderBy", []) if not orderbys: return None @@ -579,19 +593,19 @@ def _parse_orderby(query_params, allow_mri: bool = False): return orderby_list @staticmethod - def _parse_limit(paginator_kwargs) -> Limit | None: + def _parse_limit(paginator_kwargs: Mapping[str, int]) -> Limit | None: if "limit" not in paginator_kwargs: return None return Limit(paginator_kwargs["limit"]) @staticmethod - def _parse_offset(paginator_kwargs) -> Offset | None: + def _parse_offset(paginator_kwargs: Mapping[str, int]) -> Offset | None: if "offset" not in paginator_kwargs: return None return Offset(paginator_kwargs["offset"]) -def get_date_range(params: Mapping) -> tuple[datetime, datetime, int]: +def get_date_range(params: Mapping[str, Any]) -> tuple[datetime, datetime, int]: """Get start, end, rollup for the given parameters. Apply a similar logic as `sessions_v2.get_constrained_date_range`, @@ -747,11 +761,16 @@ 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"Field {metric_groupby_field.field} needs to be a metric field for grouped metric metadata" + ) 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: @@ -809,7 +828,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], @@ -821,82 +840,70 @@ 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: - 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": - return transform_null_transaction_to_unparameterized( - use_case_id, org_id, metric_action_by_field.alias - ) - - # Handles the case when we are trying to group or order by `project` for example, but we want - # to translate it to `project_id` as that is what the metrics dataset understands. - if metric_action_by_field.field in FIELD_ALIAS_MAPPINGS: - column_name = FIELD_ALIAS_MAPPINGS[metric_action_by_field.field] - elif metric_action_by_field.field in FIELD_ALIAS_MAPPINGS.values(): - column_name = metric_action_by_field.field - 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: - assert isinstance(metric_action_by_field.field, str) - column_name = resolve_tag_key(use_case_id, org_id, metric_action_by_field.field) - else: - raise NotImplementedError( - f"Unsupported string field: {metric_action_by_field.field}" + if isinstance(metric_action_by_field, MetricGroupByField): + groupby_field = metric_action_by_field + if isinstance(groupby_field.field, str): + # This transformation is currently supported only for group by because + # OrderBy doesn't support the Function type. + if groupby_field.field == "transaction": + return transform_null_transaction_to_unparameterized( + use_case_id, org_id, groupby_field.alias ) - exp = ( - AliasedExpression( - exp=Column(name=column_name), - alias=metric_action_by_field.alias, - ) - if is_group_by and not is_column - else Column(name=column_name) - ) + # Handles the case when we are trying to group by `project` for + # example, but we want to translate it to `project_id` as that is + # what the metrics dataset understands. + if groupby_field.field in FIELD_ALIAS_MAPPINGS: + column_name = FIELD_ALIAS_MAPPINGS[groupby_field.field] + elif groupby_field.field in FIELD_ALIAS_MAPPINGS.values(): + column_name = groupby_field.field + else: + column_name = resolve_tag_key(use_case_id, org_id, groupby_field.field) - if is_order_by: - # 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)] + if is_column: + return Column(name=column_name) + return AliasedExpression(exp=Column(name=column_name), alias=groupby_field.alias) - 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 - ) - - 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, - projects=projects, - )[0] - elif is_order_by: - 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, - projects=projects, - direction=metric_action_by_field.direction, - ) - else: - raise NotImplementedError( - f"Unsupported metric field: {metric_action_by_field.field}" - ) - + metric_expression = metric_object_factory(groupby_field.field.op, groupby_field.field.metric_mri) + return metric_expression.generate_groupby_statements( + use_case_id=use_case_id, + alias=groupby_field.field.alias, + params=_as_metric_operation_params(groupby_field.field.params), + projects=projects, + )[0] except IndexError: - 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" + raise InvalidParams(f"Cannot resolve {groupby_field.field} into SnQL") + + orderby_field = metric_action_by_field + if isinstance(orderby_field.field, str): + # Handles the case when we are trying to order by `project` for example, + # but we want to translate it to `project_id` as that is what the metrics + # dataset understands. + if orderby_field.field in FIELD_ALIAS_MAPPINGS: + column_name = FIELD_ALIAS_MAPPINGS[orderby_field.field] + elif orderby_field.field in FIELD_ALIAS_MAPPINGS.values(): + column_name = orderby_field.field + 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. + raise NotImplementedError(f"Unsupported string field: {orderby_field.field}") + + # We return a list in order to use the "extend" method and reduce the number + # of changes across the codebase. + return [OrderBy(exp=Column(name=column_name), direction=orderby_field.direction)] + + try: + metric_expression = metric_object_factory(orderby_field.field.op, orderby_field.field.metric_mri) + return metric_expression.generate_orderby_clause( + use_case_id=use_case_id, + alias=orderby_field.field.alias, + params=_as_metric_operation_params(orderby_field.field.params), + projects=projects, + direction=orderby_field.direction, ) + except IndexError: + raise InvalidParams(f"Cannot resolve {orderby_field.field} into SnQL") def _build_where(self) -> list[BooleanCondition | Condition]: where: list[BooleanCondition | Condition] = [ @@ -922,20 +929,32 @@ def _build_where(self) -> list[BooleanCondition | Condition]: Condition( lhs=metric_expression.generate_where_statements( use_case_id=self._use_case_id, - params=condition.lhs.params, + params=_as_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) + if ( + condition.lhs.op is not None + and require_rhs_condition_resolution(condition.lhs.op) + and isinstance(condition.rhs, str) + ) else condition.rhs ), ) ) except IndexError: raise InvalidParams(f"Cannot resolve {condition.lhs} into SnQL") + if ( + condition.lhs.op is not None + and require_rhs_condition_resolution(condition.lhs.op) + and not isinstance(condition.rhs, str) + ): + raise InvalidParams( + f"Unsupported rhs type for metric condition field {condition}: {type(condition.rhs)!r}" + ) else: snuba_conditions.append(condition) @@ -1033,18 +1052,18 @@ def _build_having(self) -> list[BooleanCondition | Condition]: def __build_totals_and_series_queries( self, - entity, - select, - where, - having, - groupby, - orderby, - limit, - offset, - rollup, - intervals_len, - ): - rv = {} + entity: MetricEntity, + select: list[Column | AliasedExpression | Function], + where: list[BooleanCondition | Condition], + having: ConditionGroup | None, + groupby: list[Column] | None, + orderby: list[OrderBy] | None, + limit: Limit, + offset: Offset | None, + rollup: Granularity, + intervals_len: int, + ) -> dict[str, Query]: + rv: dict[str, Query] = {} totals_query = Query( match=Entity(entity), groupby=groupby, @@ -1066,8 +1085,11 @@ def __build_totals_and_series_queries( series_limit = self._metrics_query.max_limit if self._use_case_id in [UseCaseID.TRANSACTIONS, UseCaseID.SPANS]: + interval = self._metrics_query.interval + if interval is None: + interval = self._metrics_query.granularity.granularity time_groupby_column = self.__generate_time_groupby_column_for_discover_queries( - self._metrics_query.interval + interval ) else: time_groupby_column = Column(TS_COL_GROUP) @@ -1096,10 +1118,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]], + component_entities: Mapping[MetricEntity | None, 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, + parent_alias: str, ) -> dict[tuple[str | 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 @@ -1107,6 +1129,8 @@ def __update_query_dicts_with_component_entities( # an aggregate and in this case, we would need to parse the op here op = 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 @@ -1124,7 +1148,12 @@ def __update_query_dicts_with_component_entities( fields_in_entities.setdefault(entity, []).append(metric_key) return metric_mri_to_obj_dict - def get_snuba_queries(self): + def get_snuba_queries( + self, + ) -> tuple[ + dict[MetricEntity, dict[str, Query]], + dict[MetricEntity, list[tuple[str | None, str, str]]], + ]: metric_mri_to_obj_dict: dict[tuple[str | None, str, str], MetricExpressionBase] = {} fields_in_entities: dict[MetricEntity, list[tuple[str | None, str, str]]] = {} @@ -1182,6 +1211,8 @@ def get_snuba_queries(self): queries_dict = {} for entity, fields in fields_in_entities.items(): + if self._metrics_query.limit is None: + raise InvalidParams("A limit is required for metrics queries") select = [] metric_ids_set = set() for field in fields: @@ -1193,12 +1224,16 @@ def get_snuba_queries(self): # 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 {})} + metric_operation_params: dict[str, str | int | float] = {} + if self._metrics_query.interval is not None: + metric_operation_params["interval"] = self._metrics_query.interval + if params is not None: + metric_operation_params.update(_as_metric_operation_params(params) or {}) select += metric_field_obj.generate_select_statements( projects=self._projects, use_case_id=self._use_case_id, alias=field[2], - params=params, + params=metric_operation_params or None, ) metric_ids_set |= metric_field_obj.generate_metric_ids( self._projects, self._use_case_id @@ -1254,9 +1289,9 @@ def __init__( metrics_query: DeprecatingMetricsQuery, fields_in_entities: dict[MetricEntity, list[tuple[str | None, str, str]]], intervals: list[datetime], - results, + results: Mapping[str, Any], use_case_id: UseCaseID, - ): + ) -> None: self._organization_id = organization_id self._intervals = intervals self._results = results @@ -1270,7 +1305,7 @@ def __init__( } # This is a set of all the `(op, metric_mri, alias)` combinations passed in the metrics_query - self._metrics_query_fields_set = { + self._metrics_query_fields_set: set[tuple[MetricOperationType | None, str, str]] = { (field.op, field.metric_mri, field.alias) for field in metrics_query.select } # This is a set of all queryable `(op, metric_mri)` combinations. Queryable can mean it @@ -1278,11 +1313,15 @@ def __init__( # SingularEntityDerivedMetric or the instances of SingularEntityDerivedMetric that are # the constituents necessary to calculate instances of CompositeEntityDerivedMetric but # are not necessarily requested in the query definition - self._fields_in_entities_set = { - elem for fields_in_entity in fields_in_entities.values() for elem in fields_in_entity + self._fields_in_entities_set: set[tuple[MetricOperationType | None, str, str]] = { + (cast(MetricOperationType | None, op), metric_mri, alias) + for fields_in_entity in fields_in_entities.values() + for op, metric_mri, alias in fields_in_entity } - self._set_of_constituent_queries = self._fields_in_entities_set.union( + self._set_of_constituent_queries: set[tuple[MetricOperationType | None, str, str]] = ( + self._fields_in_entities_set.union( self._metrics_query_fields_set + ) ) # This basically generate a dependency tree for all instances of `MetricFieldBase` so @@ -1294,7 +1333,9 @@ def __init__( self._timestamp_index = {timestamp: index for index, timestamp in enumerate(intervals)} - def _extract_data(self, data, groups: dict[tuple[tuple[str, str], ...], _SeriesTotals]) -> None: + def _extract_data( + self, data: dict[str, Any], groups: dict[tuple[tuple[str, str], ...], _SeriesTotals] + ) -> None: group_key_aliases = ( {metric_groupby_obj.alias for metric_groupby_obj in self._metrics_query.groupby} if self._metrics_query.groupby @@ -1354,7 +1395,7 @@ def _extract_data(self, data, groups: dict[tuple[tuple[str, str], ...], _SeriesT if series[series_index] == default_null_value: series[series_index] = cleaned_value - def translate_result_groups(self): + def translate_result_groups(self) -> list[_BySeriesTotals]: groups_d: dict[tuple[tuple[str, str], ...], _SeriesTotals] = {} for _, subresults in self._results.items(): for k in "totals", "series": @@ -1433,7 +1474,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=_as_metric_operation_params(params), alias=alias ) if series is not None: @@ -1448,7 +1489,7 @@ 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=_as_metric_operation_params(params), idx=idx, alias=alias ) # Remove the extra fields added due to the constituent metrics that were added