mirror of
https://github.com/apache/superset.git
synced 2026-07-20 05:36:00 +00:00
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:
@@ -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(
|
||||
|
||||
@@ -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
|
||||
)
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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")
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user