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
3 changes: 3 additions & 0 deletions superset/mcp_service/app.py
Original file line number Diff line number Diff line change
Expand Up @@ -118,6 +118,9 @@ def get_default_instructions(branding: str = "Apache Superset") -> str:
- chart_type="xy", kind="scatter": Scatter plot for correlation analysis
- chart_type="table": Data table for detailed views
- chart_type="table", viz_type="ag-grid-table": Interactive AG Grid table
- chart_type="pie": Pie chart for proportional data (set donut=True for donut)
- chart_type="pivot_table": Interactive pivot table for cross-tabulation
- chart_type="mixed_timeseries": Dual-series chart combining two chart types

Time grain for temporal x-axis (time_grain parameter):
- PT1H (hourly), P1D (daily), P1W (weekly), P1M (monthly), P1Y (yearly)
Expand Down
289 changes: 255 additions & 34 deletions superset/mcp_service/chart/chart_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -30,6 +30,9 @@
ChartCapabilities,
ChartSemantics,
ColumnRef,
MixedTimeseriesChartConfig,
PieChartConfig,
PivotTableChartConfig,
TableChartConfig,
XYChartConfig,
)
Expand Down Expand Up @@ -301,14 +304,24 @@ def is_column_truly_temporal(column_name: str, dataset_id: int | str | None) ->


def map_config_to_form_data(
config: TableChartConfig | XYChartConfig,
config: TableChartConfig
| XYChartConfig
| PieChartConfig
| PivotTableChartConfig
| MixedTimeseriesChartConfig,
dataset_id: int | str | None = None,
) -> Dict[str, Any]:
"""Map chart config to Superset form_data."""
if isinstance(config, TableChartConfig):
return map_table_config(config)
elif isinstance(config, XYChartConfig):
return map_xy_config(config, dataset_id=dataset_id)
elif isinstance(config, PieChartConfig):
return map_pie_config(config)
elif isinstance(config, PivotTableChartConfig):
return map_pivot_table_config(config)
elif isinstance(config, MixedTimeseriesChartConfig):
return map_mixed_timeseries_config(config, dataset_id=dataset_id)
else:
raise ValueError(f"Unsupported config type: {type(config)}")

Expand Down Expand Up @@ -567,6 +580,197 @@ def map_xy_config(
return form_data


def map_pie_config(config: PieChartConfig) -> Dict[str, Any]:
"""Map pie chart config to Superset form_data."""
metric = create_metric_object(config.metric)

form_data: Dict[str, Any] = {
"viz_type": "pie",
"groupby": [config.dimension.name],
"metric": metric,
"color_scheme": "supersetColors",
"show_labels": config.show_labels,
"show_legend": config.show_legend,
"label_type": config.label_type,
"number_format": config.number_format,
"sort_by_metric": config.sort_by_metric,
"row_limit": config.row_limit,
"donut": config.donut,
"show_total": config.show_total,
"labels_outside": config.labels_outside,
"outerRadius": config.outer_radius,
"innerRadius": config.inner_radius,
"date_format": "smart_date",
}

if config.filters:
form_data["adhoc_filters"] = [
{
"clause": "WHERE",
"expressionType": "SIMPLE",
"subject": filter_config.column,
"operator": map_filter_operator(filter_config.op),
"comparator": filter_config.value,
}
for filter_config in config.filters
if filter_config is not None
]

return form_data


def map_pivot_table_config(config: PivotTableChartConfig) -> Dict[str, Any]:
"""Map pivot table config to Superset form_data."""
if not config.rows:
raise ValueError("Pivot table must have at least one row grouping column")
if not config.metrics:
raise ValueError("Pivot table must have at least one metric")

metrics = [create_metric_object(col) for col in config.metrics]

form_data: Dict[str, Any] = {
"viz_type": "pivot_table_v2",
"groupbyRows": [col.name for col in config.rows],
"groupbyColumns": [col.name for col in config.columns]
if config.columns
else [],
"metrics": metrics,
"aggregateFunction": config.aggregate_function,
"rowTotals": config.show_row_totals,
"colTotals": config.show_column_totals,
"transposePivot": config.transpose,
"combineMetric": config.combine_metric,
"valueFormat": config.value_format,
"metricsLayout": "COLUMNS",
"rowOrder": "key_a_to_z",
"colOrder": "key_a_to_z",
"row_limit": config.row_limit,
}

if config.filters:
form_data["adhoc_filters"] = [
{
"clause": "WHERE",
"expressionType": "SIMPLE",
"subject": filter_config.column,
"operator": map_filter_operator(filter_config.op),
"comparator": filter_config.value,
}
for filter_config in config.filters
if filter_config is not None
]

return form_data


_MIXED_SERIES_TYPE_MAP = {
"line": "line",
"bar": "bar",
"area": "line", # area uses line type with area=True
"scatter": "scatter",
}


def _apply_axis_to_form_data(
form_data: Dict[str, Any],
axis_config: Any,
title_key: str,
format_key: str,
log_key: str | None = None,
) -> None:
"""Apply a single axis configuration to form_data."""
if not axis_config:
return
if axis_config.title:
form_data[title_key] = axis_config.title
if axis_config.format:
form_data[format_key] = axis_config.format
if log_key and axis_config.scale == "log":
form_data[log_key] = True


def _add_mixed_axis_config(
form_data: Dict[str, Any],
config: MixedTimeseriesChartConfig,
) -> None:
"""Add axis configurations to mixed timeseries form_data."""
_apply_axis_to_form_data(
form_data, config.x_axis, "xAxisTitle", "x_axis_time_format"
)
_apply_axis_to_form_data(
form_data, config.y_axis, "yAxisTitle", "y_axis_format", "logAxis"
)
_apply_axis_to_form_data(
form_data,
config.y_axis_secondary,
"yAxisTitleSecondary",
"y_axis_format_secondary",
"logAxisSecondary",
)


def map_mixed_timeseries_config(
config: MixedTimeseriesChartConfig,
dataset_id: int | str | None = None,
) -> Dict[str, Any]:
"""Map mixed timeseries chart config to Superset form_data."""
if not config.y:
raise ValueError("Mixed timeseries must have at least one primary metric")
if not config.y_secondary:
raise ValueError("Mixed timeseries must have at least one secondary metric")

# Check if x-axis column is truly temporal
x_is_temporal = is_column_truly_temporal(config.x.name, dataset_id)

form_data: Dict[str, Any] = {
"viz_type": "mixed_timeseries",
"x_axis": config.x.name,
# Query A
"metrics": [create_metric_object(col) for col in config.y],
"seriesType": _MIXED_SERIES_TYPE_MAP.get(config.primary_kind, "line"),
"area": config.primary_kind == "area",
"yAxisIndex": 0,
# Query B
"metrics_b": [create_metric_object(col) for col in config.y_secondary],
"seriesTypeB": _MIXED_SERIES_TYPE_MAP.get(config.secondary_kind, "bar"),
"areaB": config.secondary_kind == "area",
"yAxisIndexB": 1,
# Display
"show_legend": config.show_legend,
"zoomable": True,
"rich_tooltip": True,
}

# Configure temporal handling
configure_temporal_handling(form_data, x_is_temporal, config.time_grain)

# Primary groupby (Query A)
if config.group_by and config.group_by.name != config.x.name:
form_data["groupby"] = [config.group_by.name]

# Secondary groupby (Query B)
if config.group_by_secondary and config.group_by_secondary.name != config.x.name:
form_data["groupby_b"] = [config.group_by_secondary.name]

_add_mixed_axis_config(form_data, config)

# Filters
if config.filters:
form_data["adhoc_filters"] = [
{
"clause": "WHERE",
"expressionType": "SIMPLE",
"subject": filter_config.column,
"operator": map_filter_operator(filter_config.op),
"comparator": filter_config.value,
}
for filter_config in config.filters
if filter_config is not None
]

return form_data


def map_filter_operator(op: str) -> str:
"""Map filter operator to Superset format."""
operator_map = {
Expand All @@ -585,7 +789,13 @@ def map_filter_operator(op: str) -> str:
return operator_map.get(op, op)


def generate_chart_name(config: TableChartConfig | XYChartConfig) -> str:
def generate_chart_name(
config: TableChartConfig
| XYChartConfig
| PieChartConfig
| PivotTableChartConfig
| MixedTimeseriesChartConfig,
) -> str:
"""Generate a chart name based on the configuration."""
if isinstance(config, TableChartConfig):
return f"Table Chart - {', '.join(col.name for col in config.columns)}"
Expand All @@ -594,6 +804,16 @@ def generate_chart_name(config: TableChartConfig | XYChartConfig) -> str:
x_col = config.x.name
y_cols = ", ".join(col.name for col in config.y)
return f"{chart_type} Chart - {x_col} vs {y_cols}"
elif isinstance(config, PieChartConfig):
metric_label = config.metric.label or config.metric.name
return f"Pie Chart - {config.dimension.name} by {metric_label}"
elif isinstance(config, PivotTableChartConfig):
rows = ", ".join(col.name for col in config.rows)
return f"Pivot Table - {rows}"
elif isinstance(config, MixedTimeseriesChartConfig):
primary = ", ".join(col.name for col in config.y)
secondary = ", ".join(col.name for col in config.y_secondary)
return f"Mixed Chart - {primary} + {secondary}"
else:
return "Chart"

Expand All @@ -603,22 +823,7 @@ def analyze_chart_capabilities(chart: Any | None, config: Any) -> ChartCapabilit
if chart:
viz_type = getattr(chart, "viz_type", "unknown")
else:
# Map config chart_type to viz_type
chart_type = getattr(config, "chart_type", "unknown")
if chart_type == "xy":
kind = getattr(config, "kind", "line")
viz_type_map = {
"line": "echarts_timeseries_line",
"bar": "echarts_timeseries_bar",
"area": "echarts_area",
"scatter": "echarts_timeseries_scatter",
}
viz_type = viz_type_map.get(kind, "echarts_timeseries_line")
elif chart_type == "table":
# Use the viz_type from config if available (table or ag-grid-table)
viz_type = getattr(config, "viz_type", "table")
else:
viz_type = "unknown"
viz_type = _resolve_viz_type(config)

# Determine interaction capabilities based on chart type
interactive_types = [
Expand Down Expand Up @@ -663,27 +868,35 @@ def analyze_chart_capabilities(chart: Any | None, config: Any) -> ChartCapabilit
)


def _resolve_viz_type(config: Any) -> str:
"""Resolve viz_type from a chart config object."""
chart_type = getattr(config, "chart_type", "unknown")
if chart_type == "xy":
kind = getattr(config, "kind", "line")
viz_type_map = {
"line": "echarts_timeseries_line",
"bar": "echarts_timeseries_bar",
"area": "echarts_area",
"scatter": "echarts_timeseries_scatter",
}
return viz_type_map.get(kind, "echarts_timeseries_line")
elif chart_type == "table":
return getattr(config, "viz_type", "table")
elif chart_type == "pie":
return "pie"
elif chart_type == "pivot_table":
return "pivot_table_v2"
elif chart_type == "mixed_timeseries":
return "mixed_timeseries"
return "unknown"


def analyze_chart_semantics(chart: Any | None, config: Any) -> ChartSemantics:
"""Generate semantic understanding of the chart."""
if chart:
viz_type = getattr(chart, "viz_type", "unknown")
else:
# Map config chart_type to viz_type
chart_type = getattr(config, "chart_type", "unknown")
if chart_type == "xy":
kind = getattr(config, "kind", "line")
viz_type_map = {
"line": "echarts_timeseries_line",
"bar": "echarts_timeseries_bar",
"area": "echarts_area",
"scatter": "echarts_timeseries_scatter",
}
viz_type = viz_type_map.get(kind, "echarts_timeseries_line")
elif chart_type == "table":
# Use the viz_type from config if available (table or ag-grid-table)
viz_type = getattr(config, "viz_type", "table")
else:
viz_type = "unknown"
viz_type = _resolve_viz_type(config)

# Generate primary insight based on chart type
insights_map = {
Expand All @@ -696,6 +909,14 @@ def analyze_chart_semantics(chart: Any | None, config: Any) -> ChartSemantics:
),
"pie": "Shows proportional relationships within a dataset",
"echarts_area": "Emphasizes cumulative totals and part-to-whole relationships",
"pivot_table_v2": (
"Cross-tabulates data with rows, columns, and aggregated metrics "
"for multi-dimensional analysis"
),
"mixed_timeseries": (
"Combines two different chart types on the same time axis "
"for comparing related metrics with different scales"
),
}

primary_insight = insights_map.get(
Expand Down
Loading
Loading