Skip to content
Open
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
39 changes: 39 additions & 0 deletions superset/mcp_service/chart/chart_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,7 @@

import hashlib
import logging
from collections.abc import Mapping
from dataclasses import dataclass
from typing import Any, Dict, TYPE_CHECKING

Expand Down Expand Up @@ -570,10 +571,48 @@ def map_table_config(config: TableChartConfig) -> Dict[str, Any]:

form_data["row_limit"] = config.row_limit
add_color_scheme(form_data, config.color_scheme)
if config.column_config is not None:
form_data["column_config"] = {
label: column.model_dump(by_alias=True, exclude_none=True)
for label, column in config.column_config.items()
}

return form_data


def merge_table_column_config(
existing_form_data: Mapping[str, Any], new_form_data: Dict[str, Any]
) -> None:
"""Merge MCP table formatting without discarding UI-only settings.

An omitted ``column_config`` preserves the saved value, an explicit empty
mapping clears it, and a non-empty mapping updates only the supplied labels
and properties. The nested merge is important because the Superset UI stores
additional column settings that the MCP schema does not expose.
"""
if "column_config" not in new_form_data:
if "column_config" in existing_form_data:
new_form_data["column_config"] = existing_form_data["column_config"]
return

new_column_config = new_form_data["column_config"]
if not isinstance(new_column_config, dict) or not new_column_config:
return

existing_column_config = existing_form_data.get("column_config")
if not isinstance(existing_column_config, dict):
return

merged_column_config = dict(existing_column_config)
for label, settings in new_column_config.items():
existing_settings = existing_column_config.get(label)
if isinstance(existing_settings, dict) and isinstance(settings, dict):
merged_column_config[label] = {**existing_settings, **settings}
else:
merged_column_config[label] = settings
new_form_data["column_config"] = merged_column_config


def create_metric_object(col: ColumnRef) -> Dict[str, Any] | str:
"""Create a metric object for a column with enhanced validation.

Expand Down
43 changes: 43 additions & 0 deletions superset/mcp_service/chart/schemas.py
Original file line number Diff line number Diff line change
Expand Up @@ -1540,6 +1540,37 @@ def validate_metric_aggregate(self) -> Self:
return self


class TableColumnConfig(UnknownFieldCheckMixin):
"""Display formatting supported by the MCP table-chart schema."""

model_config = ConfigDict(
extra="ignore",
populate_by_name=True,
json_schema_extra={"additionalProperties": False},
)

column_width: int | None = Field(
None,
alias="columnWidth",
description="Minimum column width in pixels.",
ge=0,
)
d3_number_format: str | None = Field(
None,
alias="d3NumberFormat",
description="D3 number format, for example ',.2f', '$,.2f', or '.1%'.",
min_length=1,
max_length=100,
)
d3_time_format: str | None = Field(
None,
alias="d3TimeFormat",
description="D3 time format, for example '%Y-%m-%d' or '%b %d, %Y'.",
min_length=1,
max_length=100,
)


class TableChartConfig(UnknownFieldCheckMixin):
model_config = ConfigDict(extra="ignore", populate_by_name=True)

Expand Down Expand Up @@ -1586,6 +1617,18 @@ class TableChartConfig(UnknownFieldCheckMixin):
),
max_length=100,
)
column_config: dict[str, "TableColumnConfig"] | None = Field(
None,
description=(
"Per-column display settings, keyed by the result column label "
"(for a raw column this is usually its column name; for a metric, use "
"its label). Use columnWidth for minimum width in pixels, "
"d3NumberFormat for D3 number formats such as ',.2f' or '.1%', and "
"d3TimeFormat for D3 time formats such as '%Y-%m-%d'. Example: "
"{'Total Sales': {'columnWidth': 120, 'd3NumberFormat': '$,.2f'}, "
"'Order Date': {'d3TimeFormat': '%Y-%m-%d'}}."
),
)

@model_validator(mode="after")
def reject_sql_expression_in_raw_mode(self) -> "TableChartConfig":
Expand Down
27 changes: 18 additions & 9 deletions superset/mcp_service/chart/tool/update_chart.py
Original file line number Diff line number Diff line change
Expand Up @@ -39,6 +39,7 @@
analyze_chart_semantics,
generate_chart_name,
map_config_to_form_data,
merge_table_column_config,
)
from superset.mcp_service.chart.compile import validate_and_compile
from superset.mcp_service.chart.schemas import (
Expand All @@ -59,6 +60,20 @@
logger = logging.getLogger(__name__)


def _get_existing_form_data(chart: Any) -> dict[str, Any]:
"""Return a chart's saved form data, treating malformed params as empty."""
if not getattr(chart, "params", None):
return {}
try:
parsed = json.loads(chart.params)
except (ValueError, TypeError):
parsed = None
if not isinstance(parsed, dict):
logger.warning("Failed to parse existing chart.params for chart %s", chart.id)
return {}
return parsed


def _validation_error_response(message: str, details: str) -> GenerateChartResponse:
return GenerateChartResponse.model_validate(
{
Expand Down Expand Up @@ -117,6 +132,7 @@ def _build_update_payload(
parsed_config, dataset_id=effective_dataset_id
)
new_form_data.pop("_mcp_warnings", None)
merge_table_column_config(_get_existing_form_data(chart), new_form_data)

chart_name = (
request.chart_name
Expand Down Expand Up @@ -164,15 +180,7 @@ def _build_preview_form_data(
GenerateChartResponse error when neither config nor chart_name is given.
``parsed_config`` is the pre-parsed chart config from the caller.
"""
existing_form_data: dict[str, Any] = {}
if getattr(chart, "params", None):
try:
existing_form_data = json.loads(chart.params) or {}
except (ValueError, TypeError):
logger.warning(
"Failed to parse existing chart.params for chart %s", chart.id
)
existing_form_data = {}
existing_form_data = _get_existing_form_data(chart)

effective_dataset_id = (
request.dataset_id
Expand All @@ -185,6 +193,7 @@ def _build_preview_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)
merged = {**existing_form_data, **new_form_data}
else:
if not request.chart_name and request.dataset_id is None:
Expand Down
3 changes: 3 additions & 0 deletions superset/mcp_service/chart/tool/update_chart_preview.py
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@
generate_chart_name,
generate_explore_link,
map_config_to_form_data,
merge_table_column_config,
)
from superset.mcp_service.chart.compile import validate_and_compile
from superset.mcp_service.chart.preview_utils import (
Expand Down Expand Up @@ -186,6 +187,8 @@ def update_chart_preview( # noqa: C901
old_adhoc_filters = previous_form_data.get("adhoc_filters")
if old_adhoc_filters:
new_form_data["adhoc_filters"] = old_adhoc_filters
if previous_form_data:
merge_table_column_config(previous_form_data, new_form_data)

# Tier-1 schema validation against the dataset (no DB roundtrip).
# Runs AFTER the filter merge so filter columns are also validated.
Expand Down
38 changes: 38 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 @@ -38,6 +38,7 @@
map_filter_operator,
map_table_config,
map_xy_config,
merge_table_column_config,
validate_chart_dataset,
)
from superset.mcp_service.chart.schemas import (
Expand Down Expand Up @@ -144,6 +145,43 @@ def test_map_filter_operator_unknown(self) -> None:
assert map_filter_operator("UNKNOWN") == "UNKNOWN"


class TestMergeTableColumnConfig:
def test_partial_update_preserves_other_labels_and_settings(self) -> None:
existing = {
"column_config": {
"Revenue": {"columnWidth": 80, "visible": False},
"Region": {"customColumnName": "Sales region"},
}
}
updated = {
"column_config": {
"Revenue": {"columnWidth": 120, "d3NumberFormat": "$,.2f"}
}
}

merge_table_column_config(existing, updated)

assert updated["column_config"] == {
"Revenue": {
"columnWidth": 120,
"d3NumberFormat": "$,.2f",
"visible": False,
},
"Region": {"customColumnName": "Sales region"},
}

def test_omitted_preserves_and_explicit_empty_clears(self) -> None:
existing = {"column_config": {"Revenue": {"columnWidth": 80}}}
omitted: dict[str, Any] = {}
explicit_empty: dict[str, Any] = {"column_config": {}}

merge_table_column_config(existing, omitted)
merge_table_column_config(existing, explicit_empty)

assert omitted["column_config"] == existing["column_config"]
assert explicit_empty["column_config"] == {}


class TestMapTableConfig:
"""Test map_table_config function"""

Expand Down
35 changes: 35 additions & 0 deletions tests/unit_tests/mcp_service/chart/test_new_chart_types.py
Original file line number Diff line number Diff line change
Expand Up @@ -1240,6 +1240,41 @@ def test_color_scheme_omitted_when_unset(self) -> None:

assert "color_scheme" not in result

def test_column_config_in_form_data(self) -> None:
config = TableChartConfig.model_validate(
{
"chart_type": "table",
"columns": [
{"name": "order_date"},
{"name": "revenue", "aggregate": "SUM", "label": "Revenue"},
],
"column_config": {
"order_date": {
"columnWidth": 120,
"d3TimeFormat": "%Y-%m-%d",
},
"Revenue": {"d3NumberFormat": "$,.2f"},
},
}
)

result = map_table_config(config)

assert result["column_config"] == {
"order_date": {"columnWidth": 120, "d3TimeFormat": "%Y-%m-%d"},
"Revenue": {"d3NumberFormat": "$,.2f"},
}

def test_column_config_rejects_unknown_setting(self) -> None:
with pytest.raises(ValidationError, match="Unknown field 'width'"):
TableChartConfig.model_validate(
{
"chart_type": "table",
"columns": [{"name": "product"}],
"column_config": {"product": {"width": 120}},
}
)


class TestCurrencyFormatModel:
"""CurrencyFormat schema validation."""
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -44,10 +44,17 @@ def test_xy_schema_has_expected_fields(self) -> None:
assert "y" in props
assert "kind" in props

def test_table_schema_has_columns(self) -> None:
def test_table_schema_has_columns_and_column_config(self) -> None:
result = _call_schema("table")
props = result["schema"]["properties"]
assert "columns" in props
assert "column_config" in props
description = props["column_config"]["description"]
assert "columnWidth" in description
assert "d3NumberFormat" in description
assert "d3TimeFormat" in description
column_config_schema = result["schema"]["$defs"]["TableColumnConfig"]
assert column_config_schema["additionalProperties"] is False

def test_pie_schema_has_dimension_metric(self) -> None:
result = _call_schema("pie")
Expand Down
Loading
Loading