Compare commits

...
4 changed files with 500 additions and 70 deletions
+44
View File
@@ -1829,3 +1829,47 @@ def analyze_chart_semantics(viz_type: str | None, config: Any) -> ChartSemantics
anomalies=[], # Would need actual data analysis to populate
statistical_summary={}, # Would need actual data analysis to populate
)
def preserve_previous_adhoc_filters(
new_form_data: dict[str, Any], previous_form_data: dict[str, Any]
) -> None:
"""Preserve saved filters without dropping mapper-generated bindings."""
previous_filters = previous_form_data.get("adhoc_filters")
if not isinstance(previous_filters, list) or not previous_filters:
return
generated_filters = new_form_data.get("adhoc_filters", [])
previous_binding = previous_form_data.get(MCP_DASHBOARD_TIME_FILTER_SUBJECT)
new_binding = new_form_data.get(MCP_DASHBOARD_TIME_FILTER_SUBJECT)
merged_filters = [
filter_
for filter_ in previous_filters
if not (
previous_binding
and previous_binding != new_binding
and isinstance(filter_, dict)
and filter_.get("operator") == "TEMPORAL_RANGE"
and filter_.get("subject") == previous_binding
and filter_.get("comparator") == NO_TIME_RANGE
)
]
for generated_filter in generated_filters:
if not isinstance(generated_filter, dict):
if generated_filter not in merged_filters:
merged_filters.append(generated_filter)
continue
is_same_filter = any(
isinstance(previous_filter, dict)
and previous_filter.get("clause") == generated_filter.get("clause")
and previous_filter.get("expressionType")
== generated_filter.get("expressionType")
and previous_filter.get("subject") == generated_filter.get("subject")
and previous_filter.get("operator") == generated_filter.get("operator")
for previous_filter in merged_filters
)
if not is_same_filter:
merged_filters.append(generated_filter)
new_form_data["adhoc_filters"] = merged_filters
+202 -20
View File
@@ -42,10 +42,12 @@ from superset.mcp_service.chart.chart_utils import (
map_config_to_form_data,
merge_interactive_pivot_ui_config,
merge_table_column_config,
preserve_previous_adhoc_filters,
)
from superset.mcp_service.chart.compile import validate_and_compile
from superset.mcp_service.chart.schemas import (
AccessibilityMetadata,
ChartConfig,
ColumnRef,
GenerateChartResponse,
PerformanceMetadata,
@@ -196,9 +198,11 @@ def _append_table_columns(
def _merge_replacement_config(
existing_form_data: dict[str, Any],
new_form_data: dict[str, Any],
parsed_config: Any,
parsed_config: ChartConfig,
) -> dict[str, Any]:
"""Merge a replacement config, honoring an explicit empty filter list."""
"""Merge same-type config, honoring explicit filters and type changes."""
if existing_form_data.get("viz_type") != new_form_data.get("viz_type"):
return dict(new_form_data)
merged = {
**{
key: value
@@ -207,15 +211,197 @@ def _merge_replacement_config(
},
**new_form_data,
}
if getattr(parsed_config, "filters", None) == []:
fields_set = parsed_config.model_fields_set
if "filters" in fields_set and getattr(parsed_config, "filters", None) == []:
merged.pop("adhoc_filters", None)
if "group_by" in fields_set and getattr(parsed_config, "group_by", None) == []:
merged.pop("groupby", None)
if (
"group_by_secondary" in fields_set
and getattr(parsed_config, "group_by_secondary", None) == []
):
merged.pop("groupby_b", None)
if "sort_by" in fields_set and getattr(parsed_config, "sort_by", None) == []:
merged.pop("order_by_cols", None)
return merged
def _valid_dataset_reference(
value: Any,
columns: set[str],
metrics: set[str],
*,
allow_metric: bool = False,
) -> bool:
if not isinstance(value, str):
return True
normalized = value.casefold()
return normalized in columns or (allow_metric and normalized in metrics)
def _inherited_columns_match_dataset(
existing_form_data: dict[str, Any],
new_form_data: dict[str, Any],
columns: set[str],
metrics: set[str],
) -> bool:
for key in ("groupby", "groupby_b", "all_columns", "columns"):
if key in new_form_data:
continue
values = existing_form_data.get(key)
if isinstance(values, list) and not all(
_valid_dataset_reference(value, columns, metrics) for value in values
):
return False
for key in ("x_axis", "granularity_sqla"):
if key not in new_form_data and not _valid_dataset_reference(
existing_form_data.get(key), columns, metrics
):
return False
return True
def _inherited_metrics_match_dataset(
existing_form_data: dict[str, Any],
new_form_data: dict[str, Any],
columns: set[str],
metrics: set[str],
) -> bool:
if "metrics" not in new_form_data:
for metric in existing_form_data.get("metrics") or []:
if isinstance(metric, str) and not _valid_dataset_reference(
metric, columns, metrics, allow_metric=True
):
return False
if isinstance(metric, dict):
column = metric.get("column")
if isinstance(column, dict) and not _valid_dataset_reference(
column.get("column_name"), columns, metrics
):
return False
return True
def _inherited_sort_matches_dataset(
order_by_cols: Any, columns: set[str], metrics: set[str]
) -> bool:
for order_by in order_by_cols or []:
try:
column = json.loads(order_by)[0]
except (TypeError, ValueError, IndexError):
return False
if not _valid_dataset_reference(column, columns, metrics, allow_metric=True):
return False
return True
def _inherited_filters_match_dataset(
filters: Any, columns: set[str], metrics: set[str]
) -> bool:
for filter_ in filters or []:
if not isinstance(filter_, dict):
return False
if filter_.get("expressionType") not in (None, "SIMPLE"):
return False
subject = filter_.get("subject") or filter_.get("col")
allow_metric = str(filter_.get("clause", "WHERE")).upper() == "HAVING"
if not _valid_dataset_reference(
subject, columns, metrics, allow_metric=allow_metric
):
return False
return True
def _inherited_state_matches_dataset(
existing_form_data: dict[str, Any],
new_form_data: dict[str, Any],
parsed_config: ChartConfig,
dataset_id: int,
) -> bool:
"""Return whether carried-over query fields are valid for a new dataset."""
fields_set = parsed_config.model_fields_set
inherited_sort = "sort_by" not in fields_set and existing_form_data.get(
"order_by_cols"
)
inherited_filters = "filters" not in fields_set and existing_form_data.get(
"adhoc_filters"
)
inherited_query_fields = any(
key not in new_form_data and existing_form_data.get(key)
for key in (
"groupby",
"groupby_b",
"all_columns",
"columns",
"x_axis",
"granularity_sqla",
"metrics",
)
)
if not (inherited_query_fields or inherited_sort or inherited_filters):
return True
from superset.daos.dataset import DatasetDAO
from superset.mcp_service.chart.validation.dataset_validator import (
build_dataset_context_from_orm,
)
context = build_dataset_context_from_orm(DatasetDAO.find_by_id(dataset_id))
if context is None:
return False
columns = {column["name"].casefold() for column in context.available_columns}
metrics = {metric["name"].casefold() for metric in context.available_metrics}
return (
_inherited_columns_match_dataset(
existing_form_data, new_form_data, columns, metrics
)
and _inherited_metrics_match_dataset(
existing_form_data, new_form_data, columns, metrics
)
and (
not inherited_sort
or _inherited_sort_matches_dataset(inherited_sort, columns, metrics)
)
and (
not inherited_filters
or _inherited_filters_match_dataset(inherited_filters, columns, metrics)
)
)
def _build_replacement_form_data(
existing_form_data: dict[str, Any],
parsed_config: ChartConfig,
effective_dataset_id: int | None,
replacement_dataset_id: int | None = None,
) -> dict[str, Any]:
"""Map and merge a replacement config for both preview and save paths."""
new_form_data = map_config_to_form_data(
parsed_config, dataset_id=effective_dataset_id
)
new_form_data.pop("_mcp_warnings", None)
if replacement_dataset_id is not None and not _inherited_state_matches_dataset(
existing_form_data,
new_form_data,
parsed_config,
replacement_dataset_id,
):
existing_form_data = {}
if "filters" not in parsed_config.model_fields_set:
preserve_previous_adhoc_filters(new_form_data, existing_form_data)
merge_table_column_config(existing_form_data, new_form_data)
merge_interactive_pivot_ui_config(existing_form_data, new_form_data)
merged = _merge_replacement_config(existing_form_data, new_form_data, parsed_config)
if replacement_dataset_id is not None:
merged["datasource"] = f"{replacement_dataset_id}__table"
return merged
def _build_update_payload(
request: UpdateChartRequest,
chart: Any,
parsed_config: Any = None,
parsed_config: ChartConfig | None = None,
) -> dict[str, Any] | GenerateChartResponse:
"""Build the update payload for a chart update.
@@ -230,12 +416,13 @@ def _build_update_payload(
)
if parsed_config is not None:
new_form_data = map_config_to_form_data(
parsed_config, dataset_id=effective_dataset_id
existing_form_data = _get_existing_form_data(chart)
new_form_data = _build_replacement_form_data(
existing_form_data,
parsed_config,
effective_dataset_id,
replacement_dataset_id=request.dataset_id,
)
new_form_data.pop("_mcp_warnings", None)
merge_table_column_config(_get_existing_form_data(chart), new_form_data)
merge_interactive_pivot_ui_config(_get_existing_form_data(chart), new_form_data)
chart_name = (
request.chart_name
@@ -294,7 +481,7 @@ def _build_update_payload(
def _build_preview_form_data(
request: UpdateChartRequest,
chart: Any,
parsed_config: Any = None,
parsed_config: ChartConfig | None = None,
) -> dict[str, Any] | GenerateChartResponse:
"""Merge the existing chart's form_data with the requested changes.
@@ -312,16 +499,11 @@ def _build_preview_form_data(
)
if parsed_config is not None:
new_form_data = map_config_to_form_data(
parsed_config, dataset_id=effective_dataset_id
)
new_form_data.pop("_mcp_warnings", None)
merge_table_column_config(existing_form_data, new_form_data)
merge_interactive_pivot_ui_config(existing_form_data, new_form_data)
# In the preview, an explicit filters list, including [], replaces saved
# filters. An omitted filters field preserves them through the shallow merge.
merged = _merge_replacement_config(
existing_form_data, new_form_data, parsed_config
merged = _build_replacement_form_data(
existing_form_data,
parsed_config,
effective_dataset_id,
replacement_dataset_id=request.dataset_id,
)
elif request.add_columns is not None:
patched = _append_table_columns(existing_form_data, request.add_columns)
@@ -38,10 +38,9 @@ from superset.mcp_service.chart.chart_utils import (
generate_chart_name,
generate_explore_link,
map_config_to_form_data,
MCP_DASHBOARD_TIME_FILTER_SUBJECT,
merge_interactive_pivot_ui_config,
merge_table_column_config,
NO_TIME_RANGE,
preserve_previous_adhoc_filters as _preserve_previous_adhoc_filters,
)
from superset.mcp_service.chart.compile import validate_and_compile
from superset.mcp_service.chart.preview_utils import (
@@ -106,50 +105,6 @@ def _get_previous_form_data(form_data_key: str) -> dict[str, Any] | None:
return None
def _preserve_previous_adhoc_filters(
new_form_data: dict[str, Any], previous_form_data: dict[str, Any]
) -> None:
"""Preserve cached filters without dropping mapper-generated bindings."""
previous_filters = previous_form_data.get("adhoc_filters")
if not isinstance(previous_filters, list) or not previous_filters:
return
generated_filters = new_form_data.get("adhoc_filters", [])
previous_binding = previous_form_data.get(MCP_DASHBOARD_TIME_FILTER_SUBJECT)
new_binding = new_form_data.get(MCP_DASHBOARD_TIME_FILTER_SUBJECT)
merged_filters = [
filter_
for filter_ in previous_filters
if not (
previous_binding
and previous_binding != new_binding
and isinstance(filter_, dict)
and filter_.get("operator") == "TEMPORAL_RANGE"
and filter_.get("subject") == previous_binding
and filter_.get("comparator") == NO_TIME_RANGE
)
]
for generated_filter in generated_filters:
if not isinstance(generated_filter, dict):
if generated_filter not in merged_filters:
merged_filters.append(generated_filter)
continue
is_same_filter = any(
isinstance(previous_filter, dict)
and previous_filter.get("clause") == generated_filter.get("clause")
and previous_filter.get("expressionType")
== generated_filter.get("expressionType")
and previous_filter.get("subject") == generated_filter.get("subject")
and previous_filter.get("operator") == generated_filter.get("operator")
for previous_filter in merged_filters
)
if not is_same_filter:
merged_filters.append(generated_filter)
new_form_data["adhoc_filters"] = merged_filters
@tool(
tags=["mutate"],
class_permission_name="Chart",
@@ -36,6 +36,7 @@ from superset.mcp_service.chart.schemas import (
FilterConfig,
GenerateChartResponse,
LegendConfig,
MixedTimeseriesChartConfig,
TableChartConfig,
UpdateChartRequest,
XYChartConfig,
@@ -43,6 +44,7 @@ from superset.mcp_service.chart.schemas import (
from superset.mcp_service.chart.tool.update_chart import (
_build_preview_form_data,
_build_update_payload,
_inherited_state_matches_dataset,
)
from superset.utils import json
@@ -738,6 +740,248 @@ class TestBuildUpdatePayload:
# query_context must be cleared so get_chart_data uses updated params
assert result["query_context"] is None
def test_config_update_preserves_unrelated_mixed_timeseries_settings(self):
"""Save payload retains settings outside the simplified config schema."""
config = MixedTimeseriesChartConfig(
x=ColumnRef(name="ds"),
y=[ColumnRef(name="new_primary", aggregate="SUM")],
y_secondary=[ColumnRef(name="new_secondary", aggregate="SUM")],
)
request = UpdateChartRequest(identifier=1, config=config)
chart = Mock(
id=1,
datasource_id=7,
slice_name="Year-over-year metrics",
params=json.dumps(
{
"viz_type": "mixed_timeseries",
"time_compare": ["1 year ago"],
"comparison_type_b": "percentage",
"y_axis_format": ",.2f",
}
),
)
result = _build_update_payload(request, chart, parsed_config=config)
assert isinstance(result, dict)
saved_form_data = json.loads(result["params"])
assert saved_form_data["time_compare"] == ["1 year ago"]
assert saved_form_data["comparison_type_b"] == "percentage"
assert saved_form_data["y_axis_format"] == ",.2f"
@patch(
"superset.mcp_service.chart.chart_utils.is_column_truly_temporal",
return_value=True,
)
def test_temporal_update_preserves_omitted_non_temporal_filters(
self, unused_temporal_mock
) -> None:
"""Generated time bindings do not replace omitted saved predicates."""
country_filter = {
"clause": "WHERE",
"comparator": "US",
"expressionType": "SIMPLE",
"operator": "==",
"subject": "country",
}
temporal_filter = {
"clause": "WHERE",
"comparator": "No filter",
"expressionType": "SIMPLE",
"operator": "TEMPORAL_RANGE",
"subject": "ds",
}
config = XYChartConfig(
x=ColumnRef(name="ds"),
y=[ColumnRef(name="revenue", aggregate="SUM")],
)
request = UpdateChartRequest(identifier=1, config=config)
chart = Mock(
id=1,
datasource_id=7,
slice_name="Revenue",
params=json.dumps(
{
"viz_type": "echarts_timeseries_line",
"adhoc_filters": [country_filter, temporal_filter],
"_mcp_dashboard_time_filter_subject": "ds",
}
),
)
saved = _build_update_payload(request, chart, parsed_config=config)
preview = _build_preview_form_data(request, chart, parsed_config=config)
assert isinstance(saved, dict)
assert isinstance(preview, dict)
saved_filters = json.loads(saved["params"])["adhoc_filters"]
assert saved_filters == [country_filter, temporal_filter]
assert preview["adhoc_filters"] == saved_filters
def test_explicit_empty_group_by_clears_save_and_preview(self) -> None:
config = XYChartConfig.model_validate(
{
"x": {"name": "ds"},
"y": [{"name": "revenue", "aggregate": "SUM"}],
"groupby": [],
}
)
request = UpdateChartRequest(identifier=1, config=config)
chart = Mock(
id=1,
datasource_id=7,
slice_name="Revenue by region",
params=json.dumps(
{
"viz_type": "echarts_timeseries_line",
"groupby": ["region"],
}
),
)
saved = _build_update_payload(request, chart, parsed_config=config)
preview = _build_preview_form_data(request, chart, parsed_config=config)
assert isinstance(saved, dict)
assert isinstance(preview, dict)
assert "groupby" not in json.loads(saved["params"])
assert "groupby" not in preview
def test_explicit_empty_sort_by_clears_save_and_preview(self) -> None:
config = TableChartConfig.model_validate(
{
"columns": [{"name": "region"}],
"order_by_cols": [],
}
)
request = UpdateChartRequest(identifier=1, config=config)
chart = Mock(
id=1,
datasource_id=7,
slice_name="Regions",
params=json.dumps(
{
"viz_type": "table",
"order_by_cols": ['["region", false]'],
}
),
)
saved = _build_update_payload(request, chart, parsed_config=config)
preview = _build_preview_form_data(request, chart, parsed_config=config)
assert isinstance(saved, dict)
assert isinstance(preview, dict)
assert "order_by_cols" not in json.loads(saved["params"])
assert "order_by_cols" not in preview
@patch.object(update_chart_module, "_inherited_state_matches_dataset")
def test_dataset_rebind_only_preserves_compatible_state(
self, mock_state_matches
) -> None:
config = TableChartConfig(columns=[ColumnRef(name="revenue")])
request = UpdateChartRequest(identifier=1, config=config, dataset_id=9)
chart = Mock(
id=1,
datasource_id=7,
slice_name="Revenue",
params=json.dumps(
{
"viz_type": "table",
"groupby": ["removed_column"],
"custom_flag": True,
}
),
)
mock_state_matches.return_value = False
incompatible = _build_update_payload(request, chart, parsed_config=config)
mock_state_matches.return_value = True
compatible = _build_update_payload(request, chart, parsed_config=config)
assert isinstance(incompatible, dict)
assert isinstance(compatible, dict)
incompatible_params = json.loads(incompatible["params"])
compatible_params = json.loads(compatible["params"])
assert "groupby" not in incompatible_params
assert "custom_flag" not in incompatible_params
assert compatible_params["groupby"] == ["removed_column"]
assert compatible_params["custom_flag"] is True
assert incompatible_params["datasource"] == "9__table"
assert compatible_params["datasource"] == "9__table"
@patch(
"superset.mcp_service.chart.validation.dataset_validator."
"build_dataset_context_from_orm"
)
@patch("superset.daos.dataset.DatasetDAO.find_by_id")
def test_dataset_rebind_compatibility_uses_inherited_references(
self, mock_find_dataset, mock_build_context
) -> None:
mock_find_dataset.return_value = Mock()
mock_build_context.return_value = Mock(
available_columns=[{"name": "revenue"}, {"name": "region"}],
available_metrics=[],
)
config = TableChartConfig(columns=[ColumnRef(name="revenue")])
new_form_data = {"viz_type": "table", "all_columns": ["revenue"]}
assert _inherited_state_matches_dataset(
{
"viz_type": "table",
"groupby": ["region"],
"adhoc_filters": [
{
"expressionType": "SIMPLE",
"subject": "region",
"operator": "==",
"comparator": "US",
}
],
},
new_form_data,
config,
9,
)
assert not _inherited_state_matches_dataset(
{
"viz_type": "table",
"groupby": ["removed_column"],
},
new_form_data,
config,
9,
)
def test_config_update_does_not_merge_settings_from_another_viz_type(self):
"""Changing visualization types drops stale query-defining settings."""
config = XYChartConfig(
x=ColumnRef(name="ds"),
y=[ColumnRef(name="revenue", aggregate="SUM")],
)
request = UpdateChartRequest(identifier=1, config=config)
chart = Mock(
id=1,
datasource_id=7,
slice_name="Raw records",
params=json.dumps(
{
"viz_type": "table",
"query_mode": "raw",
"all_columns": ["ds", "revenue"],
}
),
)
result = _build_update_payload(request, chart, parsed_config=config)
assert isinstance(result, dict)
saved_form_data = json.loads(result["params"])
assert saved_form_data["viz_type"] == "echarts_timeseries_line"
assert "query_mode" not in saved_form_data
assert "all_columns" not in saved_form_data
def test_add_columns_preserves_existing_columns_and_metrics(self):
"""An additive update does not require reconstructing the table."""
request = UpdateChartRequest(
@@ -1159,7 +1403,9 @@ class TestUpdateChartPreviewFirst:
mock_chart.slice_name = "Existing Chart"
mock_chart.viz_type = "table"
mock_chart.uuid = "abc-123"
mock_chart.params = '{"viz_type": "table", "datasource": "10__table"}'
mock_chart.params = (
'{"viz_type": "table", "datasource": "10__table", "custom_flag": true}'
)
mock_find_by_id.return_value = mock_chart
mock_check_access.return_value = DatasetValidationResult(
@@ -1195,6 +1441,7 @@ class TestUpdateChartPreviewFirst:
# Ensure the chart was NOT persisted
mock_update_cmd_cls.assert_not_called()
mock_create_preview.assert_called_once()
assert mock_create_preview.call_args.args[1]["custom_flag"] is True
@patch.object(update_chart_module, "_create_preview_url", new_callable=Mock)
@patch(
@@ -1251,7 +1498,7 @@ class TestBuildPreviewFormData:
chart.id = 42
chart.datasource_id = 7
chart.slice_name = "Existing"
chart.params = '{"viz_type": "line", "custom_flag": true}'
chart.params = '{"viz_type": "table", "custom_flag": true}'
result = _build_preview_form_data(request, chart, parsed_config=config)
@@ -1444,7 +1691,7 @@ class TestUpdateChartSaveWithConfig:
mock_chart.slice_name = "Pre-save"
mock_chart.viz_type = "table"
mock_chart.uuid = "uuid-77"
mock_chart.params = '{"viz_type": "table"}'
mock_chart.params = '{"viz_type": "table", "custom_flag": true}'
mock_find_by_id.return_value = mock_chart
mock_check_access.return_value = DatasetValidationResult(
@@ -1484,6 +1731,7 @@ class TestUpdateChartSaveWithConfig:
# Verify query_context is cleared so get_chart_data uses updated params
payload = mock_update_cmd_cls.call_args[0][1]
assert payload["query_context"] is None
assert json.loads(payload["params"])["custom_flag"] is True
# Verify form_data is returned in the response
form_data = result.structured_content["form_data"]
@@ -2098,7 +2346,7 @@ class TestBuildUpdatePayloadDatasetId:
columns=[ColumnRef(name="col1")],
)
request = UpdateChartRequest(identifier=1, config=config, dataset_id=99)
chart = Mock()
chart = Mock(params='{"viz_type":"table","datasource":"10__table"}')
chart.datasource_id = 10
chart.slice_name = "Old Name"
@@ -2109,6 +2357,7 @@ class TestBuildUpdatePayloadDatasetId:
assert result["datasource_type"] == "table"
assert "params" in result
assert "viz_type" in result
assert json.loads(result["params"])["datasource"] == "99__table"
def test_config_without_dataset_does_not_include_datasource(self):
"""When dataset_id is None, payload must NOT include datasource_id."""