Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions superset/mcp_service/chart/chart_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -694,6 +694,11 @@ def merge_chart_form_data( # noqa: C901
if "filters" not in fields_set:
preserve_previous_adhoc_filters(new_form_data, existing_form_data)
merged = {**existing_form_data, **new_form_data}
# Preserve the shared color/limit controls when omitted. Chart-specific
# presentation defaults retain their existing mapper behavior.
for field in ("color_scheme", "row_limit"):
if field not in fields_set and field in existing_form_data:
merged[field] = existing_form_data[field]
Comment thread
dennisimoo marked this conversation as resolved.
# An explicitly empty collection clears the control rather than
# falling through to the inherited value.
for config_field, form_data_field in (
Expand Down
2 changes: 1 addition & 1 deletion superset/mcp_service/chart/plugins/big_number.py
Original file line number Diff line number Diff line change
Expand Up @@ -206,7 +206,7 @@ def resolve_viz_type(self, config: Any) -> str:
return "big_number_total"

def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any:
config_dict = config.model_dump()
config_dict = config.model_dump(exclude_unset=True)

if config_dict.get("metric"):
if config_dict["metric"].get("sql_expression"):
Expand Down
2 changes: 1 addition & 1 deletion superset/mcp_service/chart/plugins/box_plot.py
Original file line number Diff line number Diff line change
Expand Up @@ -109,7 +109,7 @@ def resolve_viz_type(self, config: Any) -> str:
return "box_plot"

def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any:
config_dict = config.model_dump()
config_dict = config.model_dump(exclude_unset=True)

for metric in config_dict.get("metrics") or []:
if metric.get("sql_expression"):
Expand Down
2 changes: 1 addition & 1 deletion superset/mcp_service/chart/plugins/handlebars.py
Original file line number Diff line number Diff line change
Expand Up @@ -153,7 +153,7 @@ def resolve_viz_type(self, config: Any) -> str:
return "handlebars"

def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any:
config_dict = config.model_dump()
config_dict = config.model_dump(exclude_unset=True)

def _norm_list(key: str) -> None:
if config_dict.get(key):
Expand Down
2 changes: 1 addition & 1 deletion superset/mcp_service/chart/plugins/histogram.py
Original file line number Diff line number Diff line change
Expand Up @@ -151,7 +151,7 @@ def resolve_viz_type(self, config: Any) -> str:
return "histogram_v2"

def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any:
config_dict = config.model_dump()
config_dict = config.model_dump(exclude_unset=True)

column = config_dict.get("column")
if column and not column.get("sql_expression"):
Expand Down
2 changes: 1 addition & 1 deletion superset/mcp_service/chart/plugins/interactive_pivot.py
Original file line number Diff line number Diff line change
Expand Up @@ -186,7 +186,7 @@ def resolve_viz_type(self, config: Any) -> str:
def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any:
if not isinstance(config, InteractivePivotChartConfig):
return config
config_dict = config.model_dump()
config_dict = config.model_dump(exclude_unset=True)

if temporal_column := config_dict.get("temporal_column"):
config_dict["temporal_column"] = DatasetValidator.get_canonical_column_name(
Expand Down
2 changes: 1 addition & 1 deletion superset/mcp_service/chart/plugins/mixed_timeseries.py
Original file line number Diff line number Diff line change
Expand Up @@ -122,7 +122,7 @@ def resolve_viz_type(self, config: Any) -> str:
return "mixed_timeseries"

def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any:
config_dict = config.model_dump()
config_dict = config.model_dump(exclude_unset=True)

def _norm_single(key: str) -> None:
if config_dict.get(key):
Expand Down
2 changes: 1 addition & 1 deletion superset/mcp_service/chart/plugins/pie.py
Original file line number Diff line number Diff line change
Expand Up @@ -96,7 +96,7 @@ def resolve_viz_type(self, config: Any) -> str:
return "pie"

def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any:
config_dict = config.model_dump()
config_dict = config.model_dump(exclude_unset=True)

if config_dict.get("dimension"):
dim = config_dict["dimension"]
Expand Down
2 changes: 1 addition & 1 deletion superset/mcp_service/chart/plugins/pivot_table.py
Original file line number Diff line number Diff line change
Expand Up @@ -122,7 +122,7 @@ def resolve_viz_type(self, config: Any) -> str:
return "pivot_table_v2"

def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any:
config_dict = config.model_dump()
config_dict = config.model_dump(exclude_unset=True)

def _norm_col_list(key: str) -> None:
if config_dict.get(key):
Expand Down
2 changes: 1 addition & 1 deletion superset/mcp_service/chart/plugins/table.py
Original file line number Diff line number Diff line change
Expand Up @@ -107,7 +107,7 @@ def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any:
# Preserve which nested column formatting fields were explicitly supplied.
# Round-tripping them through model_dump/model_validate would materialize
# omitted optional fields as None, turning a partial update into a clear.
config_dict = config.model_dump(exclude={"column_config"})
config_dict = config.model_dump(exclude={"column_config"}, exclude_unset=True)
get_canonical = DatasetValidator.get_canonical_column_name
get_canonical_metric = DatasetValidator.get_canonical_metric_name
raw_column_names: dict[str, str] = {}
Expand Down
2 changes: 1 addition & 1 deletion superset/mcp_service/chart/plugins/treemap.py
Original file line number Diff line number Diff line change
Expand Up @@ -100,7 +100,7 @@ def resolve_viz_type(self, config: Any) -> str:
return "treemap_v2"

def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any:
config_dict = config.model_dump()
config_dict = config.model_dump(exclude_unset=True)

for col in config_dict.get("groupby") or []:
if not col.get("sql_expression") and not col.get("saved_metric"):
Expand Down
2 changes: 1 addition & 1 deletion superset/mcp_service/chart/plugins/waterfall.py
Original file line number Diff line number Diff line change
Expand Up @@ -142,7 +142,7 @@ def resolve_viz_type(self, config: Any) -> str:
return "waterfall"

def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any:
config_dict = config.model_dump()
config_dict = config.model_dump(exclude_unset=True)

for key in ("x_axis", "breakdown"):
col = config_dict.get(key)
Expand Down
2 changes: 1 addition & 1 deletion superset/mcp_service/chart/plugins/xy.py
Original file line number Diff line number Diff line change
Expand Up @@ -111,7 +111,7 @@ def to_form_data(
return map_xy_config(config, dataset_id=dataset_id)

def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any:
config_dict = config.model_dump()
config_dict = config.model_dump(exclude_unset=True)
get_canonical = DatasetValidator.get_canonical_column_name
get_canonical_metric = DatasetValidator.get_canonical_metric_name

Expand Down
73 changes: 73 additions & 0 deletions tests/unit_tests/mcp_service/chart/test_chart_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -36,8 +36,10 @@
is_column_truly_temporal,
map_config_to_form_data,
map_filter_operator,
map_pie_config,
map_table_config,
map_xy_config,
merge_chart_form_data,
merge_interactive_pivot_ui_config,
merge_table_column_config,
validate_chart_dataset,
Expand All @@ -47,13 +49,84 @@
ColumnRef,
FilterConfig,
LegendConfig,
PieChartConfig,
SortByConfig,
TableChartConfig,
XYChartConfig,
)
from superset.mcp_service.chart.validation.dataset_validator import DatasetValidator
from superset.mcp_service.common.error_schemas import DatasetContext
from superset.utils.core import ColumnSpec, FilterOperator, GenericDataType


@pytest.mark.parametrize(
"updates,expected",
[
({}, {"color_scheme": "lyftColors", "row_limit": 42}),
(
{"color_scheme": "googleCategory10c", "row_limit": 200},
{"color_scheme": "googleCategory10c", "row_limit": 200},
),
(
{"color_scheme": None, "row_limit": 100},
{"color_scheme": "supersetColors", "row_limit": 100},
),
],
)
def test_merge_chart_preserves_omitted_defaults(
updates: dict[str, Any], expected: dict[str, Any]
) -> None:
config = PieChartConfig(
dimension=ColumnRef(name="product"),
metric=ColumnRef(name="revenue", aggregate="SUM"),
**updates,
)
config = DatasetValidator.normalize_column_names(
config,
dataset_id=1,
dataset_context=DatasetContext(
id=1,
table_name="sales",
database_name="db",
available_columns=[{"name": "Product"}, {"name": "Revenue"}],
available_metrics=[],
),
)
assert config.dimension.name == "Product"
assert config.metric.name == "Revenue"
new_form_data = map_pie_config(config)
existing = {
"viz_type": new_form_data["viz_type"],
"color_scheme": "lyftColors",
"row_limit": 42,
}
merged = merge_chart_form_data(existing, new_form_data, config)
assert {key: merged[key] for key in expected} == expected
assert merged["metric"] == new_form_data["metric"]


@pytest.mark.parametrize("dataset_rebind", [False, True])
def test_merge_chart_defaults_on_viz_change_or_dataset_rebind(
dataset_rebind: bool,
) -> None:
config = PieChartConfig(
dimension=ColumnRef(name="product"),
metric=ColumnRef(name="revenue", aggregate="SUM"),
)
new_form_data = map_pie_config(config)
existing = {
"viz_type": new_form_data["viz_type"] if dataset_rebind else "table",
"color_scheme": "lyftColors",
"row_limit": 42,
}
assert (
merge_chart_form_data(
existing, new_form_data, config, dataset_rebind=dataset_rebind
)
== new_form_data
)


class TestGetTableChartTypeLabel:
"""Test user-facing labels for table-family chart types."""

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -26,8 +26,11 @@
from unittest.mock import patch

import pytest
from pydantic import TypeAdapter

from superset.extensions import feature_flag_manager
from superset.mcp_service.chart.schemas import (
ChartConfig,
ColumnRef,
FilterConfig,
GenerateChartRequest,
Expand Down Expand Up @@ -61,6 +64,63 @@ def mock_dataset_context() -> DatasetContext:
)


@pytest.mark.parametrize(
"chart_type,roles",
[
("pie", {"dimension": "column", "metric": "metric"}),
("xy", {"x": "column", "y": "metrics"}),
("table", {"columns": "columns"}),
("big_number", {"metric": "metric"}),
("gauge", {"metric": "metric"}),
("histogram", {"column": "column"}),
("box_plot", {"metrics": "metrics", "distribute_across": "columns"}),
("treemap_v2", {"groupby": "columns", "metric": "metric"}),
("pivot_table", {"rows": "columns", "metrics": "metrics"}),
("interactive_pivot", {"rows": "columns", "metrics": "metrics"}),
("mixed_timeseries", {"x": "column", "y": "metrics", "y_secondary": "metrics"}),
("waterfall", {"x_axis": "column", "metric": "metric"}),
("handlebars", {"metrics": "metrics"}),
],
)
@pytest.mark.parametrize("updates", [{}, {"filters": []}, {"filters": None}])
def test_normalization_preserves_explicit_fields(
chart_type: str,
roles: dict[str, str],
updates: dict[str, Any],
mock_dataset_context: DatasetContext,
) -> None:
"""Column normalization must not turn omitted controls into explicit updates."""
column = {"name": "orderdate"}
metric = {"name": "sales", "aggregate": "SUM"}
values = {
"column": column,
"columns": [column],
"metric": metric,
"metrics": [metric],
}
data = {
"chart_type": chart_type,
**{key: values[role] for key, role in roles.items()},
}
if chart_type == "handlebars":
data["handlebars_template"] = "{{#each data}}{{Sales}}{{/each}}"
config = TypeAdapter(ChartConfig).validate_python({**data, **updates})
original = config.model_dump()
with patch.object(feature_flag_manager, "is_feature_enabled", return_value=True):
normalized = DatasetValidator.normalize_column_names(
config, dataset_id=18, dataset_context=mock_dataset_context
)
assert normalized.model_fields_set == config.model_fields_set
assert normalized.filters == config.filters
assert config.model_dump() == original
for key, role in roles.items():
refs = getattr(normalized, key)
if not isinstance(refs, list):
refs = [refs]
expected_name = "Sales" if role.startswith("metric") else "OrderDate"
assert all(ref.name == expected_name for ref in refs)


class TestGetCanonicalColumnName:
"""Test get_canonical_column_name static method."""

Expand Down
Loading