diff --git a/superset/common/form_data_query_context.py b/superset/common/form_data_query_context.py index eaa18dac4825..553362e68c61 100644 --- a/superset/common/form_data_query_context.py +++ b/superset/common/form_data_query_context.py @@ -130,6 +130,23 @@ def freeform_where_having(form_data: dict[str, Any]) -> dict[str, str]: return extras +def _as_column_list(value: Any) -> list[Any]: + """ + Normalize a ``groupby``/``columns`` value into a list. + + Single-select controls (e.g. the heatmap ``groupby`` Y axis, which is + ``multi: false``, and heatmap charts migrated via ``MigrateHeatmapChart``) + store the dimension as a bare string. Wrap a scalar in a one-element list, + mirroring ``chart_helpers.resolve_groupby``, so downstream list operations + (``.copy()``, ``.insert()``) do not blow up on a ``str``. + """ + if value is None: + return [] + if isinstance(value, str): + return [value] + return list(value) + + def columns_from_form_data(form_data: dict[str, Any]) -> list[Any]: """ Derive the query's grouping/raw columns from form data. @@ -141,10 +158,10 @@ def columns_from_form_data(form_data: dict[str, Any]) -> list[Any]: if form_data.get("query_mode") == "raw" and ( form_data.get("all_columns") or form_data.get("columns") ): - return list(form_data.get("all_columns") or form_data.get("columns") or []) + return _as_column_list(form_data.get("all_columns") or form_data.get("columns")) - groupby_columns: list[Any] = form_data.get("groupby") or [] - raw_columns: list[Any] = form_data.get("columns") or [] + groupby_columns: list[Any] = _as_column_list(form_data.get("groupby")) + raw_columns: list[Any] = _as_column_list(form_data.get("columns")) # Prefer explicit raw columns only when they are actually present; a stale # empty ``columns: []`` key must not shadow the group-by dimensions (which # would silently drop the grouping and change the aggregation). diff --git a/superset/mcp_service/app.py b/superset/mcp_service/app.py index 771b4cf6a0ed..1df3e06d9ac2 100644 --- a/superset/mcp_service/app.py +++ b/superset/mcp_service/app.py @@ -448,11 +448,11 @@ def get_default_instructions( chart_type_display_name field with a human-readable name when available. This field is populated for chart types known to the MCP registry (xy, pie, table, pivot_table, big_number, mixed_timeseries, handlebars, -histogram, box_plot, waterfall, gantt, bubble_v2, and interactive_pivot). -Availability gates creation and schema discovery, not display names for -existing charts. -For all other viz_types (Funnel, Gauge, Heatmap, etc.) it will be null — -use the raw viz_type field instead when referring to those chart types. +histogram, box_plot, waterfall, gantt, bubble_v2, heatmap_v2, and +interactive_pivot). Availability gates creation and schema discovery, not +display names for existing charts. For all other viz_types it will be +null — use the raw viz_type field instead when referring to those chart +types. Query Examples: - List all tables: diff --git a/superset/mcp_service/chart/chart_helpers.py b/superset/mcp_service/chart/chart_helpers.py index dbf4c5615cb7..6835c7328078 100644 --- a/superset/mcp_service/chart/chart_helpers.py +++ b/superset/mcp_service/chart/chart_helpers.py @@ -882,7 +882,10 @@ def build_query_dicts_from_form_data( qd["filters"] = [*(qd.get("filters") or []), *null_filters] return [qd] - if viz_type.startswith("echarts_timeseries"): + # Heatmap puts its x_axis in the query columns too: the frontend folds it + # into groupby in buildQuery, and the MCP path builds the query dict + # directly, so without this the x axis never reaches GROUP BY. + if viz_type.startswith("echarts_timeseries") or viz_type == "heatmap_v2": groupby = with_x_axis_column(form_data, groupby) return [ diff --git a/superset/mcp_service/chart/chart_utils.py b/superset/mcp_service/chart/chart_utils.py index f9b1a634a0a7..5b244b83b46d 100644 --- a/superset/mcp_service/chart/chart_utils.py +++ b/superset/mcp_service/chart/chart_utils.py @@ -48,6 +48,7 @@ GanttChartConfig, GaugeChartConfig, HandlebarsChartConfig, + HeatmapChartConfig, HistogramChartConfig, MixedTimeseriesChartConfig, PieChartConfig, @@ -1694,6 +1695,27 @@ def map_bubble_config(config: BubbleChartConfig) -> Dict[str, Any]: return form_data +def map_heatmap_config(config: HeatmapChartConfig) -> Dict[str, Any]: + """Map heatmap config to Superset form_data (viz_type ``heatmap_v2``). + + Matches the frontend Heatmap buildQuery contract: an ``x_axis`` column and + a single ``groupby`` Y column form the two axes, one ``metric`` colours + the cells, and ``normalize_across`` selects the rank-normalization range. + The Y axis is a single-select ``groupby`` (not a list). + """ + form_data: Dict[str, Any] = { + "viz_type": "heatmap_v2", + "x_axis": config.x_axis.name, + "groupby": config.y_axis.name, + "metric": create_metric_object(config.metric), + "normalize_across": config.normalize_across, + "normalized": config.normalized, + "row_limit": config.row_limit, + } + _add_adhoc_filters(form_data, config.filters) + return form_data + + def map_histogram_config(config: "HistogramChartConfig") -> Dict[str, Any]: """Map histogram config to Superset form_data (viz_type histogram_v2). @@ -2315,6 +2337,14 @@ def _bubble_chart_what(config: BubbleChartConfig) -> str: return f"{config.entity.name}: {x_label} vs {y_label}" +def _heatmap_chart_what(config: HeatmapChartConfig) -> str: + """Build the 'what' portion for a heatmap chart name.""" + metric_label = ( + config.metric.label or config.metric.name or config.metric.sql_expression + ) + return f"{config.x_axis.name} vs {config.y_axis.name} by {metric_label}" + + def _pivot_table_what(config: PivotTableChartConfig) -> str: """Build the 'what' portion for a pivot table chart name.""" # Pivot rows reject sql_expression at validation, so name is set. diff --git a/superset/mcp_service/chart/plugins/__init__.py b/superset/mcp_service/chart/plugins/__init__.py index aeabb7c40fea..fc9685199703 100644 --- a/superset/mcp_service/chart/plugins/__init__.py +++ b/superset/mcp_service/chart/plugins/__init__.py @@ -33,6 +33,7 @@ from superset.mcp_service.chart.plugins.gantt import GanttChartPlugin from superset.mcp_service.chart.plugins.gauge import GaugeChartPlugin from superset.mcp_service.chart.plugins.handlebars import HandlebarsChartPlugin +from superset.mcp_service.chart.plugins.heatmap import HeatmapChartPlugin from superset.mcp_service.chart.plugins.histogram import HistogramChartPlugin from superset.mcp_service.chart.plugins.interactive_pivot import ( InteractivePivotChartPlugin, @@ -62,6 +63,7 @@ register(BigNumberChartPlugin()) register(HistogramChartPlugin()) register(BoxPlotChartPlugin()) +register(HeatmapChartPlugin()) register(WaterfallChartPlugin()) register(GanttChartPlugin()) @@ -72,6 +74,7 @@ "GanttChartPlugin", "GaugeChartPlugin", "HandlebarsChartPlugin", + "HeatmapChartPlugin", "HistogramChartPlugin", "InteractivePivotChartPlugin", "MixedTimeseriesChartPlugin", diff --git a/superset/mcp_service/chart/plugins/heatmap.py b/superset/mcp_service/chart/plugins/heatmap.py new file mode 100644 index 000000000000..d51c314448fa --- /dev/null +++ b/superset/mcp_service/chart/plugins/heatmap.py @@ -0,0 +1,150 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +"""Heatmap chart type plugin.""" + +from __future__ import annotations + +from collections.abc import Mapping +from typing import Any, ClassVar + +from superset.mcp_service.chart.chart_utils import ( + _heatmap_chart_what, + _summarize_filters, + map_heatmap_config, +) +from superset.mcp_service.chart.plugin import BaseChartPlugin +from superset.mcp_service.chart.schemas import ColumnRef, HeatmapChartConfig +from superset.mcp_service.chart.validation.dataset_validator import DatasetValidator +from superset.mcp_service.common.error_schemas import ChartGenerationError + + +class HeatmapChartPlugin(BaseChartPlugin): + """Plugin for heatmap chart type.""" + + chart_type = "heatmap_v2" + display_name = "Heatmap" + native_viz_types: ClassVar[Mapping[str, str]] = { + "heatmap_v2": "Heatmap", + } + + def pre_validate( + self, + config: dict[str, Any], + ) -> ChartGenerationError | None: + missing_fields = [] + + if "x_axis" not in config: + missing_fields.append("'x_axis' (column along the X axis)") + if "y_axis" not in config and "groupby" not in config: + missing_fields.append("'y_axis' (column along the Y axis)") + if "metric" not in config: + missing_fields.append("'metric' (value colouring each cell)") + + if missing_fields: + return ChartGenerationError( + error_type="missing_heatmap_fields", + message=( + f"Heatmap chart missing required fields: " + f"{', '.join(missing_fields)}" + ), + details=( + "Heatmaps plot a metric across two dimensions — one on the " + "x_axis and one on the y_axis — colouring each cell by the " + "metric value" + ), + suggestions=[ + "Add 'x_axis': {'name': 'day_of_week'}", + "Add 'y_axis': {'name': 'hour'}", + "Add 'metric': {'name': 'trips', 'aggregate': 'COUNT'}", + "Example: {'chart_type': 'heatmap_v2', " + "'x_axis': {'name': 'day_of_week'}, " + "'y_axis': {'name': 'hour'}, " + "'metric': {'name': 'trips', 'aggregate': 'COUNT'}}", + ], + error_code="MISSING_HEATMAP_FIELDS", + ) + + return None + + def extract_column_refs(self, config: Any) -> list[ColumnRef]: + if not isinstance(config, HeatmapChartConfig): + return [] + refs: list[ColumnRef] = [config.x_axis, config.y_axis, config.metric] + if config.filters: + for f in config.filters: + refs.append(ColumnRef(name=f.column)) + return refs + + def to_form_data( + self, config: Any, dataset_id: int | str | None = None + ) -> dict[str, Any]: + return map_heatmap_config(config) + + def generate_name(self, config: Any, dataset_name: str | None = None) -> str: + what = _heatmap_chart_what(config) + context = _summarize_filters(config.filters) + return self._with_context(what, context) + + def resolve_viz_type(self, config: Any) -> str: + return "heatmap_v2" + + def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any: + config_dict = config.model_dump(exclude_unset=True) + + for key in ("x_axis", "y_axis"): + col = config_dict.get(key) + if col and not col.get("sql_expression") and not col.get("saved_metric"): + col["name"] = DatasetValidator.get_canonical_column_name( + col["name"], dataset_context + ) + if config_dict.get("metric"): + if config_dict["metric"].get("sql_expression"): + pass + elif config_dict["metric"].get("saved_metric"): + config_dict["metric"]["name"] = ( + DatasetValidator.get_canonical_metric_name( + config_dict["metric"]["name"], dataset_context + ) + ) + else: + config_dict["metric"]["name"] = ( + DatasetValidator.get_canonical_column_name( + config_dict["metric"]["name"], dataset_context + ) + ) + DatasetValidator.normalize_filters(config_dict, dataset_context) + return HeatmapChartConfig.model_validate(config_dict) + + def schema_error_hint(self) -> ChartGenerationError | None: + return ChartGenerationError( + error_type="heatmap_validation_error", + message="Heatmap chart configuration validation failed", + details=( + "The heatmap chart configuration is missing required " + "fields or has invalid structure" + ), + suggestions=[ + "Ensure 'x_axis' and 'y_axis' each have a 'name'", + "Ensure 'metric' field has 'name' and 'aggregate'", + "Example: {'chart_type': 'heatmap_v2', " + "'x_axis': {'name': 'day_of_week'}, " + "'y_axis': {'name': 'hour'}, " + "'metric': {'name': 'trips', 'aggregate': 'COUNT'}}", + ], + error_code="HEATMAP_VALIDATION_ERROR", + ) diff --git a/superset/mcp_service/chart/schemas.py b/superset/mcp_service/chart/schemas.py index 9b14c9ab8c2c..6027ca5b8f54 100644 --- a/superset/mcp_service/chart/schemas.py +++ b/superset/mcp_service/chart/schemas.py @@ -1727,6 +1727,67 @@ def record_implicit_metric_aggregate(self) -> "BubbleChartConfig": return self +class HeatmapChartConfig(BaseChartConfig): + """Config for heatmap charts (viz_type ``heatmap_v2``). + + Matches the frontend Heatmap buildQuery contract: an ``x_axis`` column, a + single ``groupby`` column for the Y axis, and one ``metric`` colouring each + cell. ``normalize_across`` drives the server-side rank normalization + (whole heatmap, per-x, or per-y). + """ + + model_config = ConfigDict(extra="ignore", populate_by_name=True) + + chart_type: Literal["heatmap_v2"] = "heatmap_v2" + x_axis: ColumnRef = Field( + ..., + description="Column along the X axis", + ) + y_axis: ColumnRef = Field( + ..., + description="Column along the Y axis (form_data 'groupby'; single-select)", + validation_alias=AliasChoices("y_axis", "groupby"), + ) + metric: ColumnRef = Field( + ..., + description="Value metric colouring each cell (use aggregate e.g. SUM, " + "COUNT for ad-hoc, or set saved_metric=True for a saved dataset metric)", + ) + normalize_across: Literal["heatmap", "x", "y"] = Field( + "heatmap", + description="Range the cell colour is normalized against: the whole " + "'heatmap', each 'x' column, or each 'y' row (frontend default: " + "'heatmap'). Only takes effect when normalized=true.", + ) + normalized: bool = Field( + False, + description="Colour cells by rank within 'normalize_across' rather than " + "the raw metric value. When false (the default) 'normalize_across' has " + "no visual effect.", + ) + row_limit: int = Field( + 10000, description="Max rows queried (cells = X × Y)", ge=1, le=100000 + ) + filters: List[FilterConfig] | None = Field( + None, + description="Structured filters (column/op/value). " + "Do NOT use adhoc_filters or raw SQL expressions.", + ) + + @model_validator(mode="after") + def reject_metric_style_dimensions(self) -> "HeatmapChartConfig": + """x_axis and y_axis are dimensions, not metrics.""" + for col, name in ((self.x_axis, "x_axis"), (self.y_axis, "y_axis")): + _reject_sql_expression_on_dimension(col, name) + if col and col.is_metric: + raise ValueError( + f"{name} must be a plain column, not a metric; drop " + "'aggregate'/'saved_metric' (metrics belong in the 'metric' " + "field)" + ) + return self + + class PivotTableChartConfig(BaseChartConfig): model_config = ConfigDict(extra="ignore", populate_by_name=True) @@ -3745,6 +3806,7 @@ def validate_gantt_roles(self) -> "GanttChartConfig": | GaugeChartConfig | TreemapChartConfig | BubbleChartConfig + | HeatmapChartConfig | PivotTableChartConfig | InteractivePivotChartConfig | MixedTimeseriesChartConfig @@ -3758,7 +3820,8 @@ def validate_gantt_roles(self) -> "GanttChartConfig": discriminator=CHART_TYPE_DISCRIMINATOR, description=( "Chart configuration - specify chart_type as 'xy', 'table', " - "'pie', 'gauge', 'treemap_v2', 'bubble_v2', 'pivot_table', " + "'pie', 'gauge', 'treemap_v2', 'bubble_v2', 'heatmap_v2', " + "'pivot_table', " "'interactive_pivot', 'mixed_timeseries', 'handlebars', " "'big_number', 'histogram', 'box_plot', 'waterfall', or 'gantt'" ), diff --git a/superset/mcp_service/chart/tool/generate_chart.py b/superset/mcp_service/chart/tool/generate_chart.py index 2e9948d5a460..9bf0a4514fb2 100644 --- a/superset/mcp_service/chart/tool/generate_chart.py +++ b/superset/mcp_service/chart/tool/generate_chart.py @@ -86,7 +86,8 @@ async def generate_chart( # noqa: C901 - LLM clients MUST display returned chart URL to users - Use numeric dataset ID or UUID (NOT schema.table_name format) - MUST include chart_type in config (one of: 'xy', 'table', 'pie', - 'gauge', 'treemap_v2', 'bubble_v2', 'pivot_table', 'mixed_timeseries', + 'gauge', 'treemap_v2', 'bubble_v2', 'heatmap_v2', 'pivot_table', + 'mixed_timeseries', 'handlebars', 'big_number', 'histogram', 'box_plot', 'waterfall', 'gantt', plus host-gated types returned by get_chart_type_schema such as 'interactive_pivot') @@ -135,6 +136,9 @@ async def generate_chart( # noqa: C901 - chart_type='bubble_v2' for a scatter of bubbles sized by a metric. Required fields: entity, x, y, size (x/y/size are metrics) + - chart_type='heatmap_v2' for a two-dimensional density grid. + Required fields: x_axis, y_axis, metric + - chart_type='histogram' for value-distribution charts. Required fields: column (numeric); optional: bins, groupby, normalize, cumulative @@ -165,6 +169,7 @@ async def generate_chart( # noqa: C901 - "gauge" / "dial" / "speedometer" -> chart_type='gauge' - "treemap" / "hierarchy" -> chart_type='treemap_v2' - "bubble" / "bubble chart" -> chart_type='bubble_v2' + - "heatmap" / "density grid" -> chart_type='heatmap_v2' - "custom HTML template" -> chart_type='handlebars' - "histogram" / "distribution" -> chart_type='histogram' - "box plot" / "box and whisker" -> chart_type='box_plot' diff --git a/superset/mcp_service/chart/tool/get_chart_type_schema.py b/superset/mcp_service/chart/tool/get_chart_type_schema.py index 36b45d784c90..89fe45d99bc6 100644 --- a/superset/mcp_service/chart/tool/get_chart_type_schema.py +++ b/superset/mcp_service/chart/tool/get_chart_type_schema.py @@ -38,6 +38,7 @@ GanttChartConfig, GaugeChartConfig, HandlebarsChartConfig, + HeatmapChartConfig, HistogramChartConfig, InteractivePivotChartConfig, MixedTimeseriesChartConfig, @@ -77,6 +78,7 @@ class ChartTypeSchemaResponse(TypedDict, total=False): "histogram": TypeAdapter(HistogramChartConfig), "box_plot": TypeAdapter(BoxPlotChartConfig), "bubble_v2": TypeAdapter(BubbleChartConfig), + "heatmap_v2": TypeAdapter(HeatmapChartConfig), "waterfall": TypeAdapter(WaterfallChartConfig), "gantt": TypeAdapter(GanttChartConfig), } @@ -221,6 +223,14 @@ class ChartTypeSchemaResponse(TypedDict, total=False): "size": {"name": "population", "aggregate": "SUM"}, }, ], + "heatmap_v2": [ + { + "chart_type": "heatmap_v2", + "x_axis": {"name": "day_of_week"}, + "y_axis": {"name": "hour"}, + "metric": {"name": "trips", "aggregate": "COUNT"}, + }, + ], "waterfall": [ { "chart_type": "waterfall", @@ -365,8 +375,9 @@ def get_chart_type_schema( for a chart configuration before calling generate_chart or update_chart. Valid chart_type values depend on the host deployment. Core types are xy, - table, pie, gauge, treemap_v2, bubble_v2, pivot_table, mixed_timeseries, - handlebars, big_number, histogram, box_plot, waterfall, and gantt. + table, pie, gauge, treemap_v2, bubble_v2, heatmap_v2, pivot_table, + mixed_timeseries, handlebars, big_number, histogram, box_plot, + waterfall, and gantt. Deployments that enable an AG Grid pivot extension also expose interactive_pivot. diff --git a/tests/unit_tests/common/test_form_data_query_context.py b/tests/unit_tests/common/test_form_data_query_context.py index b1b767423079..e36c92ebbe3c 100644 --- a/tests/unit_tests/common/test_form_data_query_context.py +++ b/tests/unit_tests/common/test_form_data_query_context.py @@ -71,6 +71,21 @@ def test_columns_empty_columns_key_does_not_shadow_groupby() -> None: assert columns_from_form_data(form_data) == ["country"] +def test_columns_scalar_groupby_is_coerced_to_list() -> None: + # A single-select ``groupby`` control (e.g. heatmap_v2's Y axis, or a + # heatmap chart migrated via ``MigrateHeatmapChart``) stores the dimension + # as a bare string. It must be coerced to a one-element list, mirroring + # ``chart_helpers.resolve_groupby``, rather than crashing on ``str.copy()``. + form_data = {"x_axis": "day", "groupby": "hour"} + assert columns_from_form_data(form_data) == ["day", "hour"] + + +def test_columns_scalar_columns_is_coerced_to_list() -> None: + # The raw-columns branch must likewise tolerate a scalar ``columns`` value. + form_data = {"columns": "region"} + assert columns_from_form_data(form_data) == ["region"] + + def test_build_context_maps_groupby_metrics_and_filters() -> None: form_data = { "groupby": ["country"], 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 9c4721ffbf10..661837bacec3 100644 --- a/tests/unit_tests/mcp_service/chart/test_chart_utils.py +++ b/tests/unit_tests/mcp_service/chart/test_chart_utils.py @@ -160,6 +160,53 @@ def test_merge_bubble_preserves_omitted_defaults( assert {key: merged[key] for key in expected} == expected +@pytest.mark.parametrize( + "updates,expected", + [ + ({}, {"row_limit": 42}), + ({"row_limit": 200}, {"row_limit": 200}), + ], +) +def test_merge_heatmap_preserves_omitted_row_limit( + updates: dict[str, Any], expected: dict[str, Any] +) -> None: + """Normalizing a heatmap config must not mark unsent fields as set. + + ``merge_chart_form_data`` keeps an omitted row limit only when the field + is absent from ``model_fields_set``. Dumping the config without + ``exclude_unset`` hands back one where every field is set, so an update + that never mentions the row limit still resets it. (Heatmap has no + ``color_scheme`` control, so only the row limit applies here.) + """ + from superset.mcp_service.chart.chart_utils import map_heatmap_config + from superset.mcp_service.chart.schemas import HeatmapChartConfig + + config = HeatmapChartConfig( + chart_type="heatmap_v2", + x_axis=ColumnRef(name="product"), + y_axis=ColumnRef(name="revenue"), + 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=[], + ), + ) + new_form_data = map_heatmap_config(config) + existing = {"viz_type": new_form_data["viz_type"], "row_limit": 42} + + merged = merge_chart_form_data(existing, new_form_data, config) + + assert {key: merged[key] for key in expected} == expected + + @pytest.mark.parametrize("dataset_rebind", [False, True]) def test_merge_chart_defaults_on_viz_change_or_dataset_rebind( dataset_rebind: bool, diff --git a/tests/unit_tests/mcp_service/chart/test_heatmap_chart.py b/tests/unit_tests/mcp_service/chart/test_heatmap_chart.py new file mode 100644 index 000000000000..75d2d0cd003d --- /dev/null +++ b/tests/unit_tests/mcp_service/chart/test_heatmap_chart.py @@ -0,0 +1,288 @@ +# Licensed to the Apache Software Foundation (ASF) under one +# or more contributor license agreements. See the NOTICE file +# distributed with this work for additional information +# regarding copyright ownership. The ASF licenses this file +# to you under the Apache License, Version 2.0 (the +# "License"); you may not use this file except in compliance +# with the License. You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, +# software distributed under the License is distributed on an +# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY +# KIND, either express or implied. See the License for the +# specific language governing permissions and limitations +# under the License. + +"""Tests for the heatmap chart type plugin. + +Schema validation, form_data mapping (matching the frontend Heatmap +buildQuery contract for viz_type ``heatmap_v2`` — an ``x_axis`` column, a +single ``groupby`` Y column, and one ``metric``), native ``groupby`` +aliasing for the Y axis, and registry integration. +""" + +import pytest +from pydantic import TypeAdapter, ValidationError + +from superset.mcp_service.chart.chart_utils import map_heatmap_config +from superset.mcp_service.chart.schemas import ChartConfig, HeatmapChartConfig + + +class TestHeatmapChartConfigSchema: + """HeatmapChartConfig schema validation.""" + + def test_basic_heatmap_config(self) -> None: + config = HeatmapChartConfig( + chart_type="heatmap_v2", + x_axis={"name": "day_of_week"}, + y_axis={"name": "hour"}, + metric={"name": "trips", "aggregate": "COUNT"}, + ) + assert config.x_axis.name == "day_of_week" + assert config.y_axis.name == "hour" + assert config.normalize_across == "heatmap" # frontend default + + def test_heatmap_missing_x_axis(self) -> None: + with pytest.raises(ValidationError): + HeatmapChartConfig( + chart_type="heatmap_v2", + y_axis={"name": "hour"}, + metric={"name": "trips", "aggregate": "COUNT"}, + ) + + def test_heatmap_missing_y_axis(self) -> None: + with pytest.raises(ValidationError): + HeatmapChartConfig( + chart_type="heatmap_v2", + x_axis={"name": "day_of_week"}, + metric={"name": "trips", "aggregate": "COUNT"}, + ) + + def test_heatmap_missing_metric(self) -> None: + with pytest.raises(ValidationError): + HeatmapChartConfig( + chart_type="heatmap_v2", + x_axis={"name": "day_of_week"}, + y_axis={"name": "hour"}, + ) + + def test_heatmap_rejects_extra_fields(self) -> None: + with pytest.raises(ValidationError): + HeatmapChartConfig( + chart_type="heatmap_v2", + x_axis={"name": "day_of_week"}, + y_axis={"name": "hour"}, + metric={"name": "trips", "aggregate": "COUNT"}, + bogus=1, + ) + + def test_heatmap_axis_rejects_aggregate(self) -> None: + """An aggregate makes an axis metric-like; x_axis/y_axis are dims.""" + with pytest.raises(ValidationError): + HeatmapChartConfig( + chart_type="heatmap_v2", + x_axis={"name": "day_of_week", "aggregate": "COUNT"}, + y_axis={"name": "hour"}, + metric={"name": "trips", "aggregate": "COUNT"}, + ) + with pytest.raises(ValidationError): + HeatmapChartConfig( + chart_type="heatmap_v2", + x_axis={"name": "day_of_week"}, + y_axis={"name": "hour", "aggregate": "COUNT"}, + metric={"name": "trips", "aggregate": "COUNT"}, + ) + + def test_heatmap_y_axis_rejects_saved_metric(self) -> None: + with pytest.raises(ValidationError): + HeatmapChartConfig( + chart_type="heatmap_v2", + x_axis={"name": "day_of_week"}, + y_axis={"name": "count", "saved_metric": True}, + metric={"name": "trips", "aggregate": "COUNT"}, + ) + + def test_heatmap_invalid_normalize_across_rejected(self) -> None: + with pytest.raises(ValidationError): + HeatmapChartConfig( + chart_type="heatmap_v2", + x_axis={"name": "day_of_week"}, + y_axis={"name": "hour"}, + metric={"name": "trips", "aggregate": "COUNT"}, + normalize_across="diagonal", + ) + + def test_groupby_alias_for_y_axis(self) -> None: + """Superset-native 'groupby' is accepted for the Y-axis field.""" + config = HeatmapChartConfig.model_validate( + { + "chart_type": "heatmap_v2", + "x_axis": {"name": "day_of_week"}, + "groupby": {"name": "hour"}, + "metric": {"name": "trips", "aggregate": "COUNT"}, + } + ) + assert config.y_axis.name == "hour" + + def test_chart_config_union_dispatches_heatmap(self) -> None: + config = TypeAdapter(ChartConfig).validate_python( + { + "chart_type": "heatmap_v2", + "x_axis": {"name": "day_of_week"}, + "y_axis": {"name": "hour"}, + "metric": {"name": "trips", "aggregate": "COUNT"}, + } + ) + assert isinstance(config, HeatmapChartConfig) + + +class TestMapHeatmapConfig: + """form_data mapping must match the frontend Heatmap buildQuery.""" + + def test_basic_heatmap_form_data(self) -> None: + config = HeatmapChartConfig( + chart_type="heatmap_v2", + x_axis={"name": "day_of_week"}, + y_axis={"name": "hour"}, + metric={"name": "trips", "aggregate": "COUNT"}, + ) + form_data = map_heatmap_config(config) + assert form_data["viz_type"] == "heatmap_v2" + assert form_data["x_axis"] == "day_of_week" + # Y axis uses the groupby key as a single column (control is multi:false) + assert form_data["groupby"] == "hour" + assert form_data["metric"]["label"] == "COUNT(trips)" + assert form_data["normalize_across"] == "heatmap" + + def test_heatmap_form_data_with_normalize_and_filters(self) -> None: + config = HeatmapChartConfig( + chart_type="heatmap_v2", + x_axis={"name": "day_of_week"}, + y_axis={"name": "hour"}, + metric={"name": "trips", "aggregate": "COUNT"}, + normalize_across="x", + filters=[{"column": "year", "op": "=", "value": 2026}], + ) + form_data = map_heatmap_config(config) + assert form_data["normalize_across"] == "x" + assert form_data["adhoc_filters"], "filters must map to adhoc_filters" + + def test_normalized_defaults_false(self) -> None: + # normalize_across has no visual effect on the frontend unless the + # 'normalized' flag is also set, so it must be threaded through. + config = HeatmapChartConfig( + chart_type="heatmap_v2", + x_axis={"name": "day_of_week"}, + y_axis={"name": "hour"}, + metric={"name": "trips", "aggregate": "COUNT"}, + ) + assert map_heatmap_config(config)["normalized"] is False + + def test_normalized_true_maps_through(self) -> None: + config = HeatmapChartConfig( + chart_type="heatmap_v2", + x_axis={"name": "day_of_week"}, + y_axis={"name": "hour"}, + metric={"name": "trips", "aggregate": "COUNT"}, + normalize_across="x", + normalized=True, + ) + assert map_heatmap_config(config)["normalized"] is True + + def test_heatmap_saved_metric_maps_to_name_string(self) -> None: + config = HeatmapChartConfig( + chart_type="heatmap_v2", + x_axis={"name": "day_of_week"}, + y_axis={"name": "hour"}, + metric={"name": "avg_fare", "saved_metric": True}, + ) + assert map_heatmap_config(config)["metric"] == "avg_fare" + + +class TestHeatmapQueryContext: + """The built query must GROUP BY both axes, not just the Y (groupby) column. + + map_heatmap_config emits X under 'x_axis' and Y under 'groupby'; the query + builder folds x_axis into the columns only for time-series viz types, so + heatmap_v2 needs an explicit fold or its X dimension is dropped. + """ + + def test_x_axis_reaches_group_by(self, monkeypatch) -> None: + from superset.mcp_service.chart import chart_helpers + + monkeypatch.setattr( + chart_helpers, + "resolve_datasource_engine", + lambda datasource_id, datasource_type: "base", + ) + config = HeatmapChartConfig( + chart_type="heatmap_v2", + x_axis={"name": "day_of_week"}, + y_axis={"name": "hour"}, + metric={"name": "trips", "aggregate": "COUNT"}, + ) + form_data = map_heatmap_config(config) + queries = chart_helpers.build_query_dicts_from_form_data(form_data, 1, "table") + columns = queries[0]["columns"] + assert "day_of_week" in columns, "x_axis must reach GROUP BY" + assert "hour" in columns, "y_axis must reach GROUP BY" + + def test_x_axis_reaches_columns_from_form_data(self) -> None: + # The generate_chart compile check derives columns through a different + # builder (columns_from_form_data, used by the dashboard export path + # too), which must tolerate the scalar 'groupby' this mapper emits and + # carry both axes rather than crashing on str.copy(). + from superset.common.form_data_query_context import columns_from_form_data + + config = HeatmapChartConfig( + chart_type="heatmap_v2", + x_axis={"name": "day_of_week"}, + y_axis={"name": "hour"}, + metric={"name": "trips", "aggregate": "COUNT"}, + ) + columns = columns_from_form_data(map_heatmap_config(config)) + assert columns == ["day_of_week", "hour"] + + +class TestHeatmapPluginRegistry: + """Plugin registration and viz-type resolution.""" + + def test_heatmap_plugin_registered(self) -> None: + from superset.mcp_service.chart import registry + + plugin = registry.get("heatmap_v2") + assert plugin is not None + assert plugin.resolve_viz_type(None) == "heatmap_v2" + + def test_display_name_resolves(self) -> None: + from superset.mcp_service.chart.registry import display_name_for_viz_type + + assert display_name_for_viz_type("heatmap_v2") == "Heatmap" + + def test_pre_validate_missing_fields(self) -> None: + from superset.mcp_service.chart import registry + + plugin = registry.get("heatmap_v2") + assert plugin is not None + error = plugin.pre_validate({"chart_type": "heatmap_v2"}) + assert error is not None + assert "x_axis" in error.message + assert "metric" in error.message + + +class TestHeatmapRecommendationCategory: + """Heatmap is categorized for chart recommendations and schema discovery.""" + + def test_heatmap_in_recommendation_category_map(self) -> None: + from superset.mcp_service.chart.tool.get_chart_data import _VIZ_CATEGORY + + assert _VIZ_CATEGORY.get("heatmap_v2") == "heatmap" + + def test_get_chart_type_schema_includes_heatmap(self) -> None: + from superset.mcp_service.chart.tool.get_chart_type_schema import ( + _CHART_TYPE_ADAPTERS, + ) + + assert "heatmap_v2" in _CHART_TYPE_ADAPTERS