mirror of
https://github.com/apache/superset.git
synced 2026-09-09 00:34:49 +00:00
Compare commits
5
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
6a4de207fb | ||
|
|
a5530a92da | ||
|
|
22a7997bf6 | ||
|
|
b8c37f4270 | ||
|
|
93f49347f8 |
@@ -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
|
||||
|
||||
@@ -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."""
|
||||
|
||||
Reference in New Issue
Block a user