diff --git a/superset/mcp_service/chart/plugins/big_number.py b/superset/mcp_service/chart/plugins/big_number.py index e542f8e75f0..0147ec0f02c 100755 --- a/superset/mcp_service/chart/plugins/big_number.py +++ b/superset/mcp_service/chart/plugins/big_number.py @@ -189,10 +189,19 @@ class BigNumberChartPlugin(BaseChartPlugin): def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any: config_dict = config.model_dump() - if config_dict.get("metric") and not config_dict["metric"].get("saved_metric"): - config_dict["metric"]["name"] = DatasetValidator._get_canonical_column_name( - config_dict["metric"]["name"], dataset_context - ) + if config_dict.get("metric"): + if config_dict["metric"].get("saved_metric"): + config_dict["metric"]["name"] = ( + DatasetValidator._get_canonical_metric_name( + config_dict["metric"]["name"], dataset_context + ) + ) + else: + config_dict["metric"]["name"] = ( + DatasetValidator._get_canonical_column_name( + config_dict["metric"]["name"], dataset_context + ) + ) if config_dict.get("temporal_column"): config_dict["temporal_column"] = ( DatasetValidator._get_canonical_column_name( diff --git a/superset/mcp_service/chart/plugins/handlebars.py b/superset/mcp_service/chart/plugins/handlebars.py index 53d78cd8b82..07b59d29e10 100755 --- a/superset/mcp_service/chart/plugins/handlebars.py +++ b/superset/mcp_service/chart/plugins/handlebars.py @@ -157,7 +157,11 @@ class HandlebarsChartPlugin(BaseChartPlugin): def _norm_list(key: str) -> None: if config_dict.get(key): for col in config_dict[key]: - if not col.get("saved_metric"): + if col.get("saved_metric"): + col["name"] = DatasetValidator._get_canonical_metric_name( + col["name"], dataset_context + ) + else: col["name"] = DatasetValidator._get_canonical_column_name( col["name"], dataset_context ) diff --git a/superset/mcp_service/chart/plugins/mixed_timeseries.py b/superset/mcp_service/chart/plugins/mixed_timeseries.py index 0cf7b82e80e..897a33c5817 100755 --- a/superset/mcp_service/chart/plugins/mixed_timeseries.py +++ b/superset/mcp_service/chart/plugins/mixed_timeseries.py @@ -47,11 +47,11 @@ class MixedTimeseriesChartPlugin(BaseChartPlugin): ) -> ChartGenerationError | None: missing_fields = [] - if "x" not in config: + if "x" not in config and "x_axis" not in config: missing_fields.append("'x' (X-axis temporal column)") - if "y" not in config: + if "y" not in config and "metrics" not in config: missing_fields.append("'y' (primary Y-axis metrics)") - if "y_secondary" not in config: + if "y_secondary" not in config and "metrics_b" not in config: missing_fields.append("'y_secondary' (secondary Y-axis metrics)") if missing_fields: @@ -132,9 +132,14 @@ class MixedTimeseriesChartPlugin(BaseChartPlugin): def _norm_list(key: str) -> None: if config_dict.get(key): for col in config_dict[key]: - col["name"] = DatasetValidator._get_canonical_column_name( - col["name"], dataset_context - ) + if col.get("saved_metric"): + col["name"] = DatasetValidator._get_canonical_metric_name( + col["name"], dataset_context + ) + else: + col["name"] = DatasetValidator._get_canonical_column_name( + col["name"], dataset_context + ) _norm_single("x") _norm_list("y") diff --git a/superset/mcp_service/chart/plugins/pie.py b/superset/mcp_service/chart/plugins/pie.py index 3d87fe7f05f..0e31d219927 100755 --- a/superset/mcp_service/chart/plugins/pie.py +++ b/superset/mcp_service/chart/plugins/pie.py @@ -103,10 +103,19 @@ class PieChartPlugin(BaseChartPlugin): config_dict["dimension"]["name"], dataset_context ) ) - if config_dict.get("metric") and not config_dict["metric"].get("saved_metric"): - config_dict["metric"]["name"] = DatasetValidator._get_canonical_column_name( - config_dict["metric"]["name"], dataset_context - ) + if config_dict.get("metric"): + if config_dict["metric"].get("saved_metric"): + config_dict["metric"]["name"] = ( + DatasetValidator._get_canonical_metric_name( + config_dict["metric"]["name"], dataset_context + ) + ) + else: + config_dict["metric"]["name"] = ( + DatasetValidator._get_canonical_column_name( + config_dict["metric"]["name"], dataset_context + ) + ) DatasetValidator._normalize_filters(config_dict, dataset_context) return PieChartConfig.model_validate(config_dict) diff --git a/superset/mcp_service/chart/plugins/pivot_table.py b/superset/mcp_service/chart/plugins/pivot_table.py index 038f8c79416..9dccc539e35 100755 --- a/superset/mcp_service/chart/plugins/pivot_table.py +++ b/superset/mcp_service/chart/plugins/pivot_table.py @@ -123,9 +123,14 @@ class PivotTableChartPlugin(BaseChartPlugin): def _norm_col_list(key: str) -> None: if config_dict.get(key): for col in config_dict[key]: - col["name"] = DatasetValidator._get_canonical_column_name( - col["name"], dataset_context - ) + if col.get("saved_metric"): + col["name"] = DatasetValidator._get_canonical_metric_name( + col["name"], dataset_context + ) + else: + col["name"] = DatasetValidator._get_canonical_column_name( + col["name"], dataset_context + ) _norm_col_list("rows") _norm_col_list("metrics") diff --git a/superset/mcp_service/chart/plugins/table.py b/superset/mcp_service/chart/plugins/table.py index 86f5dcaead2..40128237e9c 100755 --- a/superset/mcp_service/chart/plugins/table.py +++ b/superset/mcp_service/chart/plugins/table.py @@ -102,9 +102,13 @@ class TableChartPlugin(BaseChartPlugin): def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any: config_dict = config.model_dump() get_canonical = DatasetValidator._get_canonical_column_name + get_canonical_metric = DatasetValidator._get_canonical_metric_name for col in config_dict.get("columns") or []: - col["name"] = get_canonical(col["name"], dataset_context) + if col.get("saved_metric"): + col["name"] = get_canonical_metric(col["name"], dataset_context) + else: + col["name"] = get_canonical(col["name"], dataset_context) DatasetValidator._normalize_filters(config_dict, dataset_context) return TableChartConfig.model_validate(config_dict) diff --git a/superset/mcp_service/chart/plugins/xy.py b/superset/mcp_service/chart/plugins/xy.py index 076826f3f08..a021a01c965 100755 --- a/superset/mcp_service/chart/plugins/xy.py +++ b/superset/mcp_service/chart/plugins/xy.py @@ -58,7 +58,7 @@ class XYChartPlugin(BaseChartPlugin): config: dict[str, Any], ) -> ChartGenerationError | None: # x is optional — defaults to dataset's main_dttm_col in map_xy_config - if "y" not in config: + if "y" not in config and "metrics" not in config: return ChartGenerationError( error_type="missing_xy_fields", message="XY chart missing required field: 'y' (Y-axis metrics)", @@ -112,13 +112,17 @@ class XYChartPlugin(BaseChartPlugin): def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any: config_dict = config.model_dump() get_canonical = DatasetValidator._get_canonical_column_name + get_canonical_metric = DatasetValidator._get_canonical_metric_name if config_dict.get("x"): config_dict["x"]["name"] = get_canonical( config_dict["x"]["name"], dataset_context ) for y_col in config_dict.get("y") or []: - y_col["name"] = get_canonical(y_col["name"], dataset_context) + if y_col.get("saved_metric"): + y_col["name"] = get_canonical_metric(y_col["name"], dataset_context) + else: + y_col["name"] = get_canonical(y_col["name"], dataset_context) for gb_col in config_dict.get("group_by") or []: gb_col["name"] = get_canonical(gb_col["name"], dataset_context) diff --git a/superset/mcp_service/chart/validation/dataset_validator.py b/superset/mcp_service/chart/validation/dataset_validator.py index 4fd52920aec..e9f0544a4cd 100644 --- a/superset/mcp_service/chart/validation/dataset_validator.py +++ b/superset/mcp_service/chart/validation/dataset_validator.py @@ -342,6 +342,25 @@ class DatasetValidator: # Return original if not found (validation should catch this case) return column_name + @staticmethod + def _get_canonical_metric_name( + metric_name: str, dataset_context: DatasetContext + ) -> str: + """Return the canonical saved-metric name from available_metrics. + + Unlike _get_canonical_column_name, this only searches available_metrics + so that a same-named column with different casing cannot shadow the + metric's canonical name. Use this whenever saved_metric=True. + + Returns the original name when no metric matches (validation catches + the missing-metric case separately). + """ + metric_lower = metric_name.lower() + for metric in dataset_context.available_metrics: + if metric["name"].lower() == metric_lower: + return metric["name"] + return metric_name + @staticmethod def _normalize_filters( config_dict: Dict[str, Any], dataset_context: DatasetContext diff --git a/tests/unit_tests/mcp_service/chart/validation/test_column_name_normalization.py b/tests/unit_tests/mcp_service/chart/validation/test_column_name_normalization.py index a81f0864f26..c83729bb3eb 100644 --- a/tests/unit_tests/mcp_service/chart/validation/test_column_name_normalization.py +++ b/tests/unit_tests/mcp_service/chart/validation/test_column_name_normalization.py @@ -665,3 +665,242 @@ class TestValidateSavedMetrics: assert not is_valid assert error is not None assert error.error_code == "INVALID_SAVED_METRIC" + + +class TestGetCanonicalMetricName: + """Tests for _get_canonical_metric_name — metrics-only lookup.""" + + def test_exact_match(self, mock_dataset_context: DatasetContext) -> None: + result = DatasetValidator._get_canonical_metric_name( + "TotalRevenue", mock_dataset_context + ) + assert result == "TotalRevenue" + + def test_case_insensitive_match(self, mock_dataset_context: DatasetContext) -> None: + result = DatasetValidator._get_canonical_metric_name( + "totalrevenue", mock_dataset_context + ) + assert result == "TotalRevenue" + + def test_unknown_metric_returns_original( + self, mock_dataset_context: DatasetContext + ) -> None: + result = DatasetValidator._get_canonical_metric_name( + "no_such_metric", mock_dataset_context + ) + assert result == "no_such_metric" + + def test_column_name_not_matched( + self, mock_dataset_context: DatasetContext + ) -> None: + """A name that matches a column but not a metric returns the original.""" + result = DatasetValidator._get_canonical_metric_name( + "Sales", mock_dataset_context + ) + assert result == "Sales" + + +@pytest.fixture +def collision_dataset_context() -> DatasetContext: + """Dataset where a column and a metric share the same case-insensitive name + but have different casing — the scenario that exposed the saved-metric bug.""" + return DatasetContext( + id=99, + table_name="sales_data", + schema="public", + database_name="examples", + available_columns=[ + {"name": "totalrevenue", "type": "DECIMAL", "is_numeric": True}, + ], + available_metrics=[ + { + "name": "TotalRevenue", + "expression": "SUM(amount)", + "description": None, + }, + ], + ) + + +class TestSavedMetricNormalizationCorrectness: + """Saved metrics must resolve against available_metrics, not available_columns. + + When a column and a metric share the same case-insensitive name but have + different casing, _get_canonical_column_name (columns-first) returns the + column's casing. For saved_metric=True refs this is wrong — downstream + metric resolution is exact-name based and expects the metric's casing. + """ + + @patch.object(DatasetValidator, "_get_dataset_context") + def test_xy_saved_metric_uses_metric_casing( + self, + mock_get_context: Any, + collision_dataset_context: DatasetContext, + ) -> None: + mock_get_context.return_value = collision_dataset_context + + config = XYChartConfig( + chart_type="xy", + x=ColumnRef(name="totalrevenue"), + y=[ColumnRef(name="totalrevenue", saved_metric=True)], + ) + normalized = DatasetValidator.normalize_column_names(config, dataset_id=99) + + # x is a regular column ref — gets column casing + assert normalized.x is not None + assert normalized.x.name == "totalrevenue" + # y is a saved metric — must get metric casing, not column casing + assert normalized.y[0].name == "TotalRevenue" + + @patch.object(DatasetValidator, "_get_dataset_context") + def test_table_saved_metric_uses_metric_casing( + self, + mock_get_context: Any, + collision_dataset_context: DatasetContext, + ) -> None: + from superset.mcp_service.chart.schemas import TableChartConfig + + mock_get_context.return_value = collision_dataset_context + + config = TableChartConfig( + chart_type="table", + columns=[ + ColumnRef(name="totalrevenue"), + ColumnRef(name="totalrevenue", saved_metric=True), + ], + ) + normalized = DatasetValidator.normalize_column_names(config, dataset_id=99) + + assert normalized.columns[0].name == "totalrevenue" + assert normalized.columns[1].name == "TotalRevenue" + + @patch.object(DatasetValidator, "_get_dataset_context") + def test_pie_saved_metric_uses_metric_casing( + self, + mock_get_context: Any, + collision_dataset_context: DatasetContext, + ) -> None: + from superset.mcp_service.chart.schemas import PieChartConfig + + mock_get_context.return_value = collision_dataset_context + + config = PieChartConfig( + chart_type="pie", + dimension=ColumnRef(name="totalrevenue"), + metric=ColumnRef(name="totalrevenue", saved_metric=True), + ) + normalized = DatasetValidator.normalize_column_names(config, dataset_id=99) + + assert normalized.dimension.name == "totalrevenue" + assert normalized.metric.name == "TotalRevenue" + + @patch.object(DatasetValidator, "_get_dataset_context") + def test_big_number_saved_metric_uses_metric_casing( + self, + mock_get_context: Any, + collision_dataset_context: DatasetContext, + ) -> None: + from superset.mcp_service.chart.schemas import BigNumberChartConfig + + mock_get_context.return_value = collision_dataset_context + + config = BigNumberChartConfig( + chart_type="big_number", + metric=ColumnRef(name="totalrevenue", saved_metric=True), + ) + normalized = DatasetValidator.normalize_column_names(config, dataset_id=99) + + assert normalized.metric.name == "TotalRevenue" + + @patch.object(DatasetValidator, "_get_dataset_context") + def test_mixed_timeseries_saved_metrics_use_metric_casing( + self, + mock_get_context: Any, + collision_dataset_context: DatasetContext, + ) -> None: + from superset.mcp_service.chart.schemas import ( + ColumnRef, + MixedTimeseriesChartConfig, + ) + + context = DatasetContext( + id=99, + table_name="sales_data", + schema="public", + database_name="examples", + available_columns=[ + {"name": "ds", "type": "TIMESTAMP", "is_temporal": True}, + {"name": "totalrevenue", "type": "DECIMAL", "is_numeric": True}, + ], + available_metrics=[ + { + "name": "TotalRevenue", + "expression": "SUM(amount)", + "description": None, + }, + { + "name": "OrderCount", + "expression": "COUNT(*)", + "description": None, + }, + ], + ) + mock_get_context.return_value = context + + config = MixedTimeseriesChartConfig( + chart_type="mixed_timeseries", + x=ColumnRef(name="ds"), + y=[ColumnRef(name="totalrevenue", saved_metric=True)], + y_secondary=[ColumnRef(name="ordercount", saved_metric=True)], + ) + normalized = DatasetValidator.normalize_column_names(config, dataset_id=99) + + assert normalized.y[0].name == "TotalRevenue" + assert normalized.y_secondary[0].name == "OrderCount" + + +class TestPreValidateAliasHandling: + """pre_validate must accept schema field aliases, not just canonical names.""" + + def test_xy_pre_validate_accepts_metrics_alias(self) -> None: + from superset.mcp_service.chart.registry import get_registry + + plugin = get_registry().get("xy") + assert plugin is not None + + config_with_alias = { + "chart_type": "xy", + "metrics": [{"name": "revenue", "aggregate": "SUM"}], + } + error = plugin.pre_validate(config_with_alias) + assert error is None, f"pre_validate rejected 'metrics' alias: {error}" + + def test_mixed_timeseries_pre_validate_accepts_x_axis_alias(self) -> None: + from superset.mcp_service.chart.registry import get_registry + + plugin = get_registry().get("mixed_timeseries") + assert plugin is not None + + config_with_alias = { + "chart_type": "mixed_timeseries", + "x_axis": {"name": "ds"}, + "metrics": [{"name": "revenue", "aggregate": "SUM"}], + "metrics_b": [{"name": "orders", "aggregate": "COUNT"}], + } + error = plugin.pre_validate(config_with_alias) + assert error is None, f"pre_validate rejected aliases: {error}" + + def test_mixed_timeseries_pre_validate_still_rejects_truly_missing(self) -> None: + from superset.mcp_service.chart.registry import get_registry + + plugin = get_registry().get("mixed_timeseries") + assert plugin is not None + + config_missing_secondary = { + "chart_type": "mixed_timeseries", + "x": {"name": "ds"}, + "y": [{"name": "revenue", "aggregate": "SUM"}], + } + error = plugin.pre_validate(config_missing_secondary) + assert error is not None + assert "y_secondary" in error.message