diff --git a/superset/mcp_service/chart/chart_utils.py b/superset/mcp_service/chart/chart_utils.py index a9934d7727b1..fa994f2662ba 100644 --- a/superset/mcp_service/chart/chart_utils.py +++ b/superset/mcp_service/chart/chart_utils.py @@ -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] # An explicitly empty collection clears the control rather than # falling through to the inherited value. for config_field, form_data_field in ( diff --git a/superset/mcp_service/chart/plugins/big_number.py b/superset/mcp_service/chart/plugins/big_number.py index 9a861883d22e..1368262d4dbe 100644 --- a/superset/mcp_service/chart/plugins/big_number.py +++ b/superset/mcp_service/chart/plugins/big_number.py @@ -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"): diff --git a/superset/mcp_service/chart/plugins/box_plot.py b/superset/mcp_service/chart/plugins/box_plot.py index 5d803fe1daad..e8cefb610137 100644 --- a/superset/mcp_service/chart/plugins/box_plot.py +++ b/superset/mcp_service/chart/plugins/box_plot.py @@ -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"): diff --git a/superset/mcp_service/chart/plugins/handlebars.py b/superset/mcp_service/chart/plugins/handlebars.py index d5ae6f5f9584..340264b1be46 100644 --- a/superset/mcp_service/chart/plugins/handlebars.py +++ b/superset/mcp_service/chart/plugins/handlebars.py @@ -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): diff --git a/superset/mcp_service/chart/plugins/histogram.py b/superset/mcp_service/chart/plugins/histogram.py index e4b43369501e..e96b96974a98 100644 --- a/superset/mcp_service/chart/plugins/histogram.py +++ b/superset/mcp_service/chart/plugins/histogram.py @@ -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"): diff --git a/superset/mcp_service/chart/plugins/interactive_pivot.py b/superset/mcp_service/chart/plugins/interactive_pivot.py index d8037d762bb7..b0d9d8ff11d7 100644 --- a/superset/mcp_service/chart/plugins/interactive_pivot.py +++ b/superset/mcp_service/chart/plugins/interactive_pivot.py @@ -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( diff --git a/superset/mcp_service/chart/plugins/mixed_timeseries.py b/superset/mcp_service/chart/plugins/mixed_timeseries.py index 8da78735ee39..a354f30b9503 100644 --- a/superset/mcp_service/chart/plugins/mixed_timeseries.py +++ b/superset/mcp_service/chart/plugins/mixed_timeseries.py @@ -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): diff --git a/superset/mcp_service/chart/plugins/pie.py b/superset/mcp_service/chart/plugins/pie.py index a16b23337a50..93aceee0bb64 100644 --- a/superset/mcp_service/chart/plugins/pie.py +++ b/superset/mcp_service/chart/plugins/pie.py @@ -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"] diff --git a/superset/mcp_service/chart/plugins/pivot_table.py b/superset/mcp_service/chart/plugins/pivot_table.py index 392047a0766e..2b90a75c0baa 100644 --- a/superset/mcp_service/chart/plugins/pivot_table.py +++ b/superset/mcp_service/chart/plugins/pivot_table.py @@ -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): diff --git a/superset/mcp_service/chart/plugins/table.py b/superset/mcp_service/chart/plugins/table.py index bd6831973104..09968a52b956 100644 --- a/superset/mcp_service/chart/plugins/table.py +++ b/superset/mcp_service/chart/plugins/table.py @@ -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] = {} diff --git a/superset/mcp_service/chart/plugins/treemap.py b/superset/mcp_service/chart/plugins/treemap.py index e566377dcb80..170c7d5d4761 100644 --- a/superset/mcp_service/chart/plugins/treemap.py +++ b/superset/mcp_service/chart/plugins/treemap.py @@ -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"): diff --git a/superset/mcp_service/chart/plugins/waterfall.py b/superset/mcp_service/chart/plugins/waterfall.py index b2c9658bb52f..50bef7084586 100644 --- a/superset/mcp_service/chart/plugins/waterfall.py +++ b/superset/mcp_service/chart/plugins/waterfall.py @@ -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) diff --git a/superset/mcp_service/chart/plugins/xy.py b/superset/mcp_service/chart/plugins/xy.py index cedd88ad0038..879222047825 100755 --- a/superset/mcp_service/chart/plugins/xy.py +++ b/superset/mcp_service/chart/plugins/xy.py @@ -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 diff --git a/tests/unit_tests/mcp_service/chart/test_chart_utils.py b/tests/unit_tests/mcp_service/chart/test_chart_utils.py index e2638b055e6e..c9321b1ae913 100644 --- a/tests/unit_tests/mcp_service/chart/test_chart_utils.py +++ b/tests/unit_tests/mcp_service/chart/test_chart_utils.py @@ -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, @@ -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.""" 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 be90f3b461a9..310d8f2e6a0d 100755 --- 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 @@ -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, @@ -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."""