-
Notifications
You must be signed in to change notification settings - Fork 304
fix(dbt): emit expr '1' for COUNT(*) metrics #432
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from all commits
94085af
664e3a1
08a693d
45504e0
ab8513f
e9b0a69
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -19,7 +19,7 @@ | |
| from collections import defaultdict | ||
| from dataclasses import dataclass | ||
| from itertools import combinations | ||
| from typing import Dict, List, Optional, Sequence, Tuple | ||
| from typing import Dict, FrozenSet, List, Optional, Sequence, Tuple | ||
|
|
||
| from ossie import ( | ||
| OssieDataset, | ||
|
|
@@ -33,6 +33,7 @@ | |
| OssieRelationship, | ||
| ) | ||
| from ossie_dbt.converter_issues import ConverterIssue, ConverterIssueType, ConverterResult | ||
| from ossie_dbt.expression_utils import ROW_COUNT_EXPR | ||
| from ossie_dbt.filter_utils import _collect_filter_sql, _merge_filter_sqls | ||
|
|
||
| from metricflow_semantic_interfaces.enum_extension import assert_values_exhausted | ||
|
|
@@ -81,14 +82,28 @@ class AmbiguousDerivedReferenceError(Exception): | |
|
|
||
|
|
||
| class MSIToOssieConverter: | ||
| """Converts an MSI SemanticManifest into an Ossie Document.""" | ||
| """Converts an MSI SemanticManifest into an Ossie Document. | ||
|
|
||
| Holds no per-call state on the instance, so one instance may be reused, including concurrently | ||
| from multiple threads, across any number of ``convert()`` calls. | ||
| """ | ||
|
|
||
| def __init__(self, dialect: OssieDialect = OssieDialect.ANSI_SQL) -> None: | ||
| self._dialect = dialect | ||
|
|
||
| def convert( | ||
| self, manifest: PydanticSemanticManifest, ossie_model_name: str = "semantic_model" | ||
| ) -> ConverterResult[OssieDocument]: | ||
| # The transformer rewrites COUNT to SUM (leaving expr '1' as SUM(1)), which loses the dataset a row | ||
| # count belongs to. Remember these metrics so they come back as COUNT(<dataset>.*). | ||
| row_count_metrics: FrozenSet[str] = frozenset( | ||
| metric.name | ||
| for metric in manifest.metrics | ||
| if metric.type is MetricType.SIMPLE | ||
| and metric.type_params.metric_aggregation_params is not None | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. This only looks at
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Yeah, the pre-scan only looks at metric_aggregation_params so it misses that shape entirely. Will fold it into the fix for the comment below. |
||
| and metric.type_params.metric_aggregation_params.agg is AggregationType.COUNT | ||
| and metric.type_params.expr == ROW_COUNT_EXPR | ||
| ) | ||
|
Comment on lines
+99
to
+106
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I believe this pre-scan never matches on the CLI path. Detection needs to work on a manifest that has already been transformed, since that is what the CLI and the README example pass in.
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Confirmed, and worse than I thought. Ran a manifest through parse_manifest_from_dbt_generated_manifest like the CLI does, then into convert() by then agg is already SUM, pre-scan finds nothing. Every real conversion just gives SUM(1), same as before this PR. SUM(1) and COUNT(*) mean the same thing regardless of which one was written, so maybe no pre-scan is needed at all, any SUM with expr == '1' is a row count. Would also let me drop the row_count_metrics state from last round. Match your thinking, or should they stay distinguishable? |
||
| manifest = PydanticSemanticManifestTransformer.transform(manifest) | ||
| issues: List[ConverterIssue] = [] | ||
|
|
||
|
|
@@ -118,7 +133,7 @@ def convert( | |
| ConverterIssue(issue_type=ConverterIssueType.CUMULATIVE_SEMANTICS_LOSS, element_name=metric.name) | ||
| ) | ||
| try: | ||
| expr = self._resolve_metric_expression(metric, metric_index, expression_cache) | ||
| expr = self._resolve_metric_expression(metric, metric_index, expression_cache, row_count_metrics) | ||
| except AmbiguousDerivedReferenceError: | ||
| # Every other unsupported shape drops one metric and records an issue; | ||
| # an ambiguous reference is no reason to fail the whole conversion. | ||
|
|
@@ -232,6 +247,7 @@ def _resolve_metric_expression( | |
| metric: Metric, | ||
| metric_index: Dict[str, Metric], | ||
| cache: Dict[Tuple[str, Optional[str]], str], | ||
| row_count_metrics: FrozenSet[str], | ||
| parent_filter: Optional[str] = None, | ||
| ) -> str: | ||
| """Recursively resolve a metric to a fully-inlined SQL expression string.""" | ||
|
|
@@ -243,13 +259,13 @@ def _resolve_metric_expression( | |
| return cache[cache_key] | ||
|
|
||
| if metric.type is MetricType.SIMPLE: | ||
| expr = self._resolve_simple(metric, combined_filter) | ||
| expr = self._resolve_simple(metric, row_count_metrics, combined_filter) | ||
| elif metric.type is MetricType.CUMULATIVE: | ||
| expr = self._resolve_cumulative(metric, metric_index, cache, combined_filter) | ||
| expr = self._resolve_cumulative(metric, metric_index, cache, row_count_metrics, combined_filter) | ||
| elif metric.type is MetricType.RATIO: | ||
| expr = self._resolve_ratio(metric, metric_index, cache, combined_filter) | ||
| expr = self._resolve_ratio(metric, metric_index, cache, row_count_metrics, combined_filter) | ||
| elif metric.type is MetricType.DERIVED: | ||
| expr = self._resolve_derived(metric, metric_index, cache, combined_filter) | ||
| expr = self._resolve_derived(metric, metric_index, cache, row_count_metrics, combined_filter) | ||
| elif metric.type is MetricType.CONVERSION: | ||
| # CONVERSION metrics are skipped in convert(); this branch should never be reached. | ||
| raise RuntimeError(f"Unexpected CONVERSION metric in expression resolver: metric_name={metric.name!r}") | ||
|
|
@@ -262,6 +278,7 @@ def _resolve_metric_expression( | |
| def _resolve_simple( | ||
| self, | ||
| metric: Metric, | ||
| row_count_metrics: FrozenSet[str], | ||
| filter_sql: Optional[str] = None, | ||
| ) -> str: | ||
| """Resolve a SIMPLE metric using metric_aggregation_params (always set after transformation).""" | ||
|
|
@@ -270,6 +287,9 @@ def _resolve_simple( | |
| raise ValueError( | ||
| f"SIMPLE metric has no metric_aggregation_params after transformation: metric_name={metric.name!r}" | ||
| ) | ||
| # With a filter the count is emitted as SUM(CASE WHEN <filter> THEN 1 END), which has no `dataset.*` form. | ||
| if metric.name in row_count_metrics and not filter_sql: | ||
| return f"COUNT({agg_params_obj.semantic_model}.*)" | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. I'm not sure
The previous output,
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. greed, want to settle this first. Leaning toward dropping it and going back to SUM(1), since the dataset tracking never worked on the real path anyway. Happy to open a dev@ thread if this should go through the spec instead. |
||
| col = metric.type_params.expr if metric.type_params.expr is not None else metric.name | ||
| col = self._qualify_col(col, agg_params_obj.semantic_model) | ||
| return self._build_agg_expression(agg_params_obj.agg, col, agg_params_obj.agg_params, filter_sql) | ||
|
|
@@ -291,6 +311,7 @@ def _resolve_cumulative( | |
| metric: Metric, | ||
| metric_index: Dict[str, Metric], | ||
| cache: Dict[Tuple[str, Optional[str]], str], | ||
| row_count_metrics: FrozenSet[str], | ||
| filter_sql: Optional[str] = None, | ||
| ) -> str: | ||
| """Resolve a CUMULATIVE metric to its base aggregation expression. | ||
|
|
@@ -308,6 +329,7 @@ def _resolve_cumulative( | |
| self._lookup_metric(metric_index, sub_input.name, f"CUMULATIVE metric '{metric.name}'"), | ||
| metric_index, | ||
| cache, | ||
| row_count_metrics, | ||
| sub_filter, | ||
| ) | ||
|
|
||
|
|
@@ -316,6 +338,7 @@ def _resolve_ratio( | |
| metric: Metric, | ||
| metric_index: Dict[str, Metric], | ||
| cache: Dict[Tuple[str, Optional[str]], str], | ||
| row_count_metrics: FrozenSet[str], | ||
| filter_sql: Optional[str] = None, | ||
| ) -> str: | ||
| """Resolve a RATIO metric as (numerator) / (denominator), both fully inlined.""" | ||
|
|
@@ -331,12 +354,14 @@ def _resolve_ratio( | |
| self._lookup_metric(metric_index, num_input.name, f"RATIO metric '{metric.name}' numerator"), | ||
| metric_index, | ||
| cache, | ||
| row_count_metrics, | ||
| num_filter, | ||
| ) | ||
| den_expr = self._resolve_metric_expression( | ||
| self._lookup_metric(metric_index, den_input.name, f"RATIO metric '{metric.name}' denominator"), | ||
| metric_index, | ||
| cache, | ||
| row_count_metrics, | ||
| den_filter, | ||
| ) | ||
| return f"({num_expr}) / ({den_expr})" | ||
|
|
@@ -346,6 +371,7 @@ def _resolve_derived( | |
| metric: Metric, | ||
| metric_index: Dict[str, Metric], | ||
| cache: Dict[Tuple[str, Optional[str]], str], | ||
| row_count_metrics: FrozenSet[str], | ||
| filter_sql: Optional[str] = None, | ||
| ) -> str: | ||
| """Resolve a DERIVED metric by substituting each input metric's expression into the expr string. | ||
|
|
@@ -375,7 +401,7 @@ def _resolve_derived( | |
| ref = input_metric.alias if input_metric.alias else input_metric.name | ||
| dep_metric = self._lookup_metric(metric_index, input_metric.name, f"DERIVED metric '{metric.name}'") | ||
| input_filter = _merge_filter_sqls(filter_sql, _collect_filter_sql(input_metric.filter)) | ||
| resolved = self._resolve_metric_expression(dep_metric, metric_index, cache, input_filter) | ||
| resolved = self._resolve_metric_expression(dep_metric, metric_index, cache, row_count_metrics, input_filter) | ||
| if dep_metric.type in (MetricType.DERIVED, MetricType.RATIO): | ||
| resolved = f"({resolved})" | ||
| distinct = resolutions.setdefault(ref, []) | ||
|
|
||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
Returning None here sends
COUNT(DISTINCT 1)to the raw fallback, which emitsexpr='COUNT(DISTINCT 1)'under agg SUM: a nested aggregate MetricFlow cannot run. On main it was a count_distinct metric with expr1.COUNT(DISTINCT orders.*)also moves from semantic_modelordersto the first dataset.The comment says these are not valid SQL, but both are. If they are unsupported, I'd rather drop them with an issue than fall back, and
test_unsupported_count_star_forms_fall_back_to_the_raw_expressioncurrently locks the fallback in.There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
You're right, both are valid SQL. Checked it, both now come out as agg=sum bound to the first dataset, worse than main. Will drop these with an issue instead of falling back.