fix(mcp): fix saved-metric name normalization across all chart plugins

Add _get_canonical_metric_name() to DatasetValidator that searches only
available_metrics, preventing a column with matching case-insensitive name
from shadowing a saved metric's canonical casing.

Update all 7 chart plugins (xy, table, pie, big_number, handlebars,
mixed_timeseries, pivot_table) to branch on saved_metric flag: saved
metrics now go through _get_canonical_metric_name while regular column
refs continue to use _get_canonical_column_name.

Fix pre_validate alias handling in xy and mixed_timeseries plugins to
accept Pydantic AliasChoices keys (metrics/x_axis/metrics_b) so payloads
using canonical Superset field names are not incorrectly rejected.

Add TestGetCanonicalMetricName, TestSavedMetricNormalizationCorrectness,
and TestPreValidateAliasHandling test classes covering the collision case
and alias acceptance.

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
Amin Ghadersohi
2026-05-21 22:37:48 +00:00
parent 83cd604b8e
commit d3f30fddbb
9 changed files with 319 additions and 21 deletions

View File

@@ -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(

View File

@@ -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
)

View File

@@ -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")

View File

@@ -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)

View File

@@ -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")

View File

@@ -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)

View File

@@ -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)

View File

@@ -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

View File

@@ -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