mirror of
https://github.com/apache/superset.git
synced 2026-07-20 05:36:00 +00:00
feat(mcp): introduce chart type plugin registry for extensible chart generation
Replaces four scattered dispatch locations (schema_validator, dataset_validator, chart_utils, runtime validator) with a central ChartTypePlugin registry. Each of the 7 supported chart types (xy, table, pie, pivot_table, mixed_timeseries, handlebars, big_number) now owns its pre-validation, column extraction, form_data mapping, post-map validation, column normalization, and runtime warnings in a single plugin class. Key changes: - Add ChartTypePlugin protocol and BaseChartPlugin base class (plugin.py) - Add ChartTypeRegistry with register/get/all_types helpers (registry.py) - Add 7 chart type plugins under chart/plugins/ with full coverage - Fix 5-type column validation gap: pie, pivot_table, mixed_timeseries, handlebars, and big_number now participate in dataset column validation (previously silently skipped) - Move BigNumber trendline temporal check to BigNumberChartPlugin.post_map_validate() - Add get_runtime_warnings() to plugin protocol; XYChartPlugin implements format/cardinality checks, removing isinstance(config, XYChartConfig) from RuntimeValidator - Fix stale generate_chart.py docstring listing only 'xy' and 'table' chart types - Add missing pie, pivot_table, mixed_timeseries handlers to _enhance_validation_error; refactor into a data-driven lookup table to stay within complexity limits - Fix empty details fallback in Pydantic error handler
This commit is contained in:
@@ -662,6 +662,7 @@ from superset.mcp_service.annotation_layer.tool import ( # noqa: F401, E402
|
||||
list_annotation_layers,
|
||||
list_layer_annotations,
|
||||
)
|
||||
import superset.mcp_service.chart.plugins # noqa: F401, E402 — registers all chart type plugins
|
||||
from superset.mcp_service.chart import ( # noqa: F401, E402
|
||||
prompts as chart_prompts,
|
||||
resources as chart_resources,
|
||||
|
||||
@@ -321,29 +321,32 @@ def map_config_to_form_data(
|
||||
| BigNumberChartConfig,
|
||||
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)
|
||||
elif isinstance(config, HandlebarsChartConfig):
|
||||
return map_handlebars_config(config)
|
||||
elif isinstance(config, BigNumberChartConfig):
|
||||
if config.show_trendline and config.temporal_column:
|
||||
if not is_column_truly_temporal(config.temporal_column, dataset_id):
|
||||
raise ValueError(
|
||||
f"Big Number trendline requires a temporal SQL column; "
|
||||
f"'{config.temporal_column}' is not temporal."
|
||||
)
|
||||
return map_big_number_config(config)
|
||||
else:
|
||||
raise ValueError(f"Unsupported config type: {type(config)}")
|
||||
"""Map chart config to Superset form_data via the plugin registry.
|
||||
|
||||
The previous if/elif chain across all 7 chart types has been replaced by a
|
||||
single registry lookup. Cross-field constraints (e.g. BigNumber trendline
|
||||
temporal check) are now owned by each plugin's post_map_validate() method
|
||||
rather than being baked into this dispatcher.
|
||||
"""
|
||||
from superset.mcp_service.chart.registry import get_registry
|
||||
|
||||
chart_type = getattr(config, "chart_type", None)
|
||||
plugin = get_registry().get(chart_type) if chart_type else None
|
||||
|
||||
if plugin is None:
|
||||
raise ValueError(
|
||||
f"Unsupported config type: {type(config)} (chart_type={chart_type!r})"
|
||||
)
|
||||
|
||||
form_data = plugin.to_form_data(config, dataset_id=dataset_id)
|
||||
|
||||
# Run post-map validation (e.g. BigNumber trendline temporal type check).
|
||||
# Raise ValueError to preserve backward-compatible error handling in callers.
|
||||
error = plugin.post_map_validate(config, form_data, dataset_id=dataset_id)
|
||||
if error is not None:
|
||||
raise ValueError(error.message)
|
||||
|
||||
return form_data
|
||||
|
||||
|
||||
def _add_adhoc_filters(
|
||||
|
||||
193
superset/mcp_service/chart/plugin.py
Normal file
193
superset/mcp_service/chart/plugin.py
Normal file
@@ -0,0 +1,193 @@
|
||||
# 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.
|
||||
|
||||
"""
|
||||
ChartTypePlugin protocol and BaseChartPlugin base class.
|
||||
|
||||
Each chart type owns its pre-validation, column extraction, form_data mapping,
|
||||
and post-map validation in a single plugin class. This eliminates the previous
|
||||
pattern of 4 separate dispatch points (schema_validator.py, dataset_validator.py,
|
||||
chart_utils.py, pipeline.py) that had to be updated in sync whenever a new chart
|
||||
type was added.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Protocol, runtime_checkable
|
||||
|
||||
from superset.mcp_service.chart.schemas import ColumnRef
|
||||
from superset.mcp_service.common.error_schemas import ChartGenerationError
|
||||
|
||||
|
||||
@runtime_checkable
|
||||
class ChartTypePlugin(Protocol):
|
||||
"""
|
||||
Protocol that every chart-type plugin must satisfy.
|
||||
|
||||
Implementing all five methods in a single class guarantees that adding a
|
||||
new chart type requires only one new file — the plugin — rather than edits
|
||||
across four separate files.
|
||||
"""
|
||||
|
||||
#: Discriminator value matching ChartConfig's chart_type field.
|
||||
chart_type: str
|
||||
|
||||
def pre_validate(
|
||||
self,
|
||||
config: dict[str, Any],
|
||||
) -> ChartGenerationError | None:
|
||||
"""
|
||||
Early validation of the raw config dict before Pydantic parsing.
|
||||
|
||||
Called by SchemaValidator before attempting to parse the request.
|
||||
Should check that required top-level keys are present and well-typed.
|
||||
|
||||
Returns None if valid, ChartGenerationError if invalid.
|
||||
"""
|
||||
...
|
||||
|
||||
def extract_column_refs(
|
||||
self,
|
||||
config: Any,
|
||||
) -> list[ColumnRef]:
|
||||
"""
|
||||
Extract all column references from a parsed chart config.
|
||||
|
||||
Called by DatasetValidator to validate that all referenced columns exist
|
||||
in the dataset. Must cover every field that holds a column name,
|
||||
including filters.
|
||||
|
||||
Returns a list of ColumnRef objects (may be empty).
|
||||
"""
|
||||
...
|
||||
|
||||
def to_form_data(
|
||||
self,
|
||||
config: Any,
|
||||
dataset_id: int | str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
"""
|
||||
Map a parsed chart config to Superset's internal form_data dict.
|
||||
|
||||
Replaces the if/elif chain in chart_utils.map_config_to_form_data().
|
||||
|
||||
Returns a Superset form_data dict ready for caching and rendering.
|
||||
"""
|
||||
...
|
||||
|
||||
def post_map_validate(
|
||||
self,
|
||||
config: Any,
|
||||
form_data: dict[str, Any],
|
||||
dataset_id: int | str | None = None,
|
||||
) -> ChartGenerationError | None:
|
||||
"""
|
||||
Validate the mapped form_data after to_form_data() runs.
|
||||
|
||||
Use this for cross-field constraints that can only be checked once
|
||||
form_data is assembled (e.g. BigNumber trendline requires a temporal
|
||||
column whose type must be verified against the dataset).
|
||||
|
||||
Returns None if valid, ChartGenerationError if invalid.
|
||||
"""
|
||||
...
|
||||
|
||||
def normalize_column_refs(
|
||||
self,
|
||||
config: Any,
|
||||
dataset_context: Any,
|
||||
) -> Any:
|
||||
"""
|
||||
Return a new config with column names normalized to canonical dataset casing.
|
||||
|
||||
Called by DatasetValidator.normalize_column_names(). The default
|
||||
implementation (in BaseChartPlugin) returns the config unchanged; plugins
|
||||
with column fields override this to fix case sensitivity mismatches.
|
||||
|
||||
Returns a new config object (or the original if no normalization needed).
|
||||
"""
|
||||
...
|
||||
|
||||
def get_runtime_warnings(
|
||||
self,
|
||||
config: Any,
|
||||
dataset_id: int | str,
|
||||
) -> list[str]:
|
||||
"""
|
||||
Return chart-type-specific runtime warnings (performance, compatibility).
|
||||
|
||||
Called by RuntimeValidator to collect per-type warnings. Warnings are
|
||||
informational only — they never block chart generation. The default
|
||||
implementation returns an empty list; plugins override this to emit
|
||||
chart-type-specific warnings (e.g. XY cardinality checks).
|
||||
|
||||
Returns a list of warning message strings (may be empty).
|
||||
"""
|
||||
...
|
||||
|
||||
|
||||
class BaseChartPlugin:
|
||||
"""
|
||||
Base class providing sensible defaults for all ChartTypePlugin methods.
|
||||
|
||||
Concrete plugins extend this and override only what they need.
|
||||
"""
|
||||
|
||||
chart_type: str = ""
|
||||
|
||||
def pre_validate(
|
||||
self,
|
||||
config: dict[str, Any],
|
||||
) -> ChartGenerationError | None:
|
||||
return None
|
||||
|
||||
def extract_column_refs(
|
||||
self,
|
||||
config: Any,
|
||||
) -> list[ColumnRef]:
|
||||
return []
|
||||
|
||||
def to_form_data(
|
||||
self,
|
||||
config: Any,
|
||||
dataset_id: int | str | None = None,
|
||||
) -> dict[str, Any]:
|
||||
raise NotImplementedError(
|
||||
f"{self.__class__.__name__}.to_form_data() is not implemented"
|
||||
)
|
||||
|
||||
def post_map_validate(
|
||||
self,
|
||||
config: Any,
|
||||
form_data: dict[str, Any],
|
||||
dataset_id: int | str | None = None,
|
||||
) -> ChartGenerationError | None:
|
||||
return None
|
||||
|
||||
def normalize_column_refs(
|
||||
self,
|
||||
config: Any,
|
||||
dataset_context: Any,
|
||||
) -> Any:
|
||||
return config
|
||||
|
||||
def get_runtime_warnings(
|
||||
self,
|
||||
config: Any,
|
||||
dataset_id: int | str,
|
||||
) -> list[str]:
|
||||
return []
|
||||
58
superset/mcp_service/chart/plugins/__init__.py
Normal file
58
superset/mcp_service/chart/plugins/__init__.py
Normal file
@@ -0,0 +1,58 @@
|
||||
# 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.
|
||||
|
||||
"""
|
||||
Chart type plugins package.
|
||||
|
||||
Importing this module registers all built-in chart type plugins in the global
|
||||
registry. This module is imported by app.py at startup.
|
||||
|
||||
To add a new chart type:
|
||||
1. Create ``superset/mcp_service/chart/plugins/{chart_type}.py``
|
||||
2. Implement a class extending ``BaseChartPlugin``
|
||||
3. Import and register it here
|
||||
"""
|
||||
|
||||
from superset.mcp_service.chart.plugins.big_number import BigNumberChartPlugin
|
||||
from superset.mcp_service.chart.plugins.handlebars import HandlebarsChartPlugin
|
||||
from superset.mcp_service.chart.plugins.mixed_timeseries import (
|
||||
MixedTimeseriesChartPlugin,
|
||||
)
|
||||
from superset.mcp_service.chart.plugins.pie import PieChartPlugin
|
||||
from superset.mcp_service.chart.plugins.pivot_table import PivotTableChartPlugin
|
||||
from superset.mcp_service.chart.plugins.table import TableChartPlugin
|
||||
from superset.mcp_service.chart.plugins.xy import XYChartPlugin
|
||||
from superset.mcp_service.chart.registry import register
|
||||
|
||||
# Register all built-in chart type plugins
|
||||
register(XYChartPlugin())
|
||||
register(TableChartPlugin())
|
||||
register(PieChartPlugin())
|
||||
register(PivotTableChartPlugin())
|
||||
register(MixedTimeseriesChartPlugin())
|
||||
register(HandlebarsChartPlugin())
|
||||
register(BigNumberChartPlugin())
|
||||
|
||||
__all__ = [
|
||||
"BigNumberChartPlugin",
|
||||
"HandlebarsChartPlugin",
|
||||
"MixedTimeseriesChartPlugin",
|
||||
"PieChartPlugin",
|
||||
"PivotTableChartPlugin",
|
||||
"TableChartPlugin",
|
||||
"XYChartPlugin",
|
||||
]
|
||||
193
superset/mcp_service/chart/plugins/big_number.py
Normal file
193
superset/mcp_service/chart/plugins/big_number.py
Normal file
@@ -0,0 +1,193 @@
|
||||
# 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.
|
||||
|
||||
"""Big number chart type plugin."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from superset.mcp_service.chart.plugin import BaseChartPlugin
|
||||
from superset.mcp_service.chart.schemas import ColumnRef
|
||||
from superset.mcp_service.common.error_schemas import ChartGenerationError
|
||||
|
||||
|
||||
class BigNumberChartPlugin(BaseChartPlugin):
|
||||
"""Plugin for big_number chart type."""
|
||||
|
||||
chart_type = "big_number"
|
||||
|
||||
def pre_validate(
|
||||
self,
|
||||
config: dict[str, Any],
|
||||
) -> ChartGenerationError | None:
|
||||
if "metric" not in config:
|
||||
return ChartGenerationError(
|
||||
error_type="missing_metric",
|
||||
message="Big Number chart missing required field: metric",
|
||||
details=(
|
||||
"Big Number charts require a 'metric' field "
|
||||
"specifying the value to display"
|
||||
),
|
||||
suggestions=[
|
||||
"Add 'metric' with name and aggregate: "
|
||||
"{'name': 'revenue', 'aggregate': 'SUM'}",
|
||||
"The aggregate function is required (SUM, COUNT, AVG, MIN, MAX)",
|
||||
"Example: {'chart_type': 'big_number', "
|
||||
"'metric': {'name': 'sales', 'aggregate': 'SUM'}}",
|
||||
],
|
||||
error_code="MISSING_BIG_NUMBER_METRIC",
|
||||
)
|
||||
|
||||
metric = config.get("metric", {})
|
||||
if not isinstance(metric, dict):
|
||||
return ChartGenerationError(
|
||||
error_type="invalid_metric_type",
|
||||
message="Big Number metric must be a dict with 'name' and 'aggregate'",
|
||||
details=(
|
||||
"The 'metric' field must be an object, got "
|
||||
f"{type(metric).__name__}"
|
||||
),
|
||||
suggestions=[
|
||||
"Use a dict: {'name': 'col', 'aggregate': 'SUM'}",
|
||||
"Valid aggregates: SUM, COUNT, AVG, MIN, MAX",
|
||||
],
|
||||
error_code="INVALID_BIG_NUMBER_METRIC_TYPE",
|
||||
)
|
||||
if not metric.get("aggregate") and not metric.get("saved_metric"):
|
||||
return ChartGenerationError(
|
||||
error_type="missing_metric_aggregate",
|
||||
message=(
|
||||
"Big Number metric must include an aggregate function "
|
||||
"or reference a saved metric"
|
||||
),
|
||||
details=(
|
||||
"The metric must have an 'aggregate' field or 'saved_metric': true"
|
||||
),
|
||||
suggestions=[
|
||||
"Add 'aggregate': {'name': 'col', 'aggregate': 'SUM'}",
|
||||
"Or use a saved metric: {'name': 'metric', 'saved_metric': true}",
|
||||
"Valid aggregates: SUM, COUNT, AVG, MIN, MAX",
|
||||
],
|
||||
error_code="MISSING_BIG_NUMBER_AGGREGATE",
|
||||
)
|
||||
|
||||
show_trendline = config.get("show_trendline", False)
|
||||
temporal_column = config.get("temporal_column")
|
||||
if show_trendline and not temporal_column:
|
||||
return ChartGenerationError(
|
||||
error_type="missing_temporal_column",
|
||||
message="Trendline requires a temporal column",
|
||||
details=(
|
||||
"When 'show_trendline' is True, "
|
||||
"a 'temporal_column' must be specified"
|
||||
),
|
||||
suggestions=[
|
||||
"Add 'temporal_column': 'date_column_name'",
|
||||
"Or set 'show_trendline': false for number only",
|
||||
"Use get_dataset_info to find temporal columns",
|
||||
],
|
||||
error_code="MISSING_TEMPORAL_COLUMN",
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
def extract_column_refs(self, config: Any) -> list[ColumnRef]:
|
||||
from superset.mcp_service.chart.schemas import BigNumberChartConfig
|
||||
|
||||
if not isinstance(config, BigNumberChartConfig):
|
||||
return []
|
||||
refs: list[ColumnRef] = [config.metric]
|
||||
# temporal_column is a str field, not a ColumnRef — validate it exists
|
||||
if config.temporal_column:
|
||||
refs.append(ColumnRef(name=config.temporal_column))
|
||||
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]:
|
||||
from superset.mcp_service.chart.chart_utils import map_big_number_config
|
||||
|
||||
return map_big_number_config(config)
|
||||
|
||||
def post_map_validate(
|
||||
self,
|
||||
config: Any,
|
||||
form_data: dict[str, Any],
|
||||
dataset_id: int | str | None = None,
|
||||
) -> ChartGenerationError | None:
|
||||
"""Verify the trendline temporal column is a real temporal SQL type.
|
||||
|
||||
This check was previously baked into map_config_to_form_data() in
|
||||
chart_utils.py as a special case. Moving it here keeps the dispatcher
|
||||
clean and makes the constraint explicit and discoverable.
|
||||
"""
|
||||
from superset.mcp_service.chart.schemas import BigNumberChartConfig
|
||||
|
||||
if not isinstance(config, BigNumberChartConfig):
|
||||
return None
|
||||
if not (config.show_trendline and config.temporal_column):
|
||||
return None
|
||||
|
||||
from superset.mcp_service.chart.chart_utils import is_column_truly_temporal
|
||||
|
||||
if not is_column_truly_temporal(config.temporal_column, dataset_id):
|
||||
return ChartGenerationError(
|
||||
error_type="non_temporal_trendline_column",
|
||||
message=(
|
||||
f"Big Number trendline requires a temporal SQL column; "
|
||||
f"'{config.temporal_column}' is not temporal."
|
||||
),
|
||||
details=(
|
||||
f"Column '{config.temporal_column}' does not have a temporal "
|
||||
f"SQL type (DATE, DATETIME, TIMESTAMP). The trendline requires "
|
||||
f"a true temporal column for DATE_TRUNC to work."
|
||||
),
|
||||
suggestions=[
|
||||
"Use get_dataset_info to find columns with temporal SQL types",
|
||||
"Set 'show_trendline': false to use any column as the metric",
|
||||
"If the column contains dates stored as integers, "
|
||||
"consider casting it in a virtual dataset",
|
||||
],
|
||||
error_code="NON_TEMPORAL_TRENDLINE_COLUMN",
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any:
|
||||
from superset.mcp_service.chart.schemas import BigNumberChartConfig
|
||||
from superset.mcp_service.chart.validation.dataset_validator import (
|
||||
DatasetValidator,
|
||||
)
|
||||
|
||||
config_dict = config.model_dump()
|
||||
|
||||
if config_dict.get("metric") and not config_dict["metric"].get("saved_metric"):
|
||||
config_dict["metric"]["name"] = DatasetValidator._get_canonical_column_name(
|
||||
config_dict["metric"]["name"], dataset_context
|
||||
)
|
||||
if config_dict.get("temporal_column"):
|
||||
config_dict["temporal_column"] = (
|
||||
DatasetValidator._get_canonical_column_name(
|
||||
config_dict["temporal_column"], dataset_context
|
||||
)
|
||||
)
|
||||
DatasetValidator._normalize_filters(config_dict, dataset_context)
|
||||
return BigNumberChartConfig.model_validate(config_dict)
|
||||
161
superset/mcp_service/chart/plugins/handlebars.py
Normal file
161
superset/mcp_service/chart/plugins/handlebars.py
Normal file
@@ -0,0 +1,161 @@
|
||||
# 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.
|
||||
|
||||
"""Handlebars chart type plugin."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from superset.mcp_service.chart.plugin import BaseChartPlugin
|
||||
from superset.mcp_service.chart.schemas import ColumnRef
|
||||
from superset.mcp_service.common.error_schemas import ChartGenerationError
|
||||
|
||||
|
||||
class HandlebarsChartPlugin(BaseChartPlugin):
|
||||
"""Plugin for handlebars chart type (custom HTML template charts)."""
|
||||
|
||||
chart_type = "handlebars"
|
||||
|
||||
def pre_validate(
|
||||
self,
|
||||
config: dict[str, Any],
|
||||
) -> ChartGenerationError | None:
|
||||
if "handlebars_template" not in config:
|
||||
return ChartGenerationError(
|
||||
error_type="missing_handlebars_template",
|
||||
message="Handlebars chart missing required field: handlebars_template",
|
||||
details=(
|
||||
"Handlebars charts require a 'handlebars_template' string "
|
||||
"containing Handlebars HTML template markup"
|
||||
),
|
||||
suggestions=[
|
||||
"Add 'handlebars_template' with a Handlebars HTML template",
|
||||
"Data is available as {{data}} array in the template",
|
||||
"Example: '<ul>{{#each data}}<li>{{this.name}}: "
|
||||
"{{this.value}}</li>{{/each}}</ul>'",
|
||||
],
|
||||
error_code="MISSING_HANDLEBARS_TEMPLATE",
|
||||
)
|
||||
|
||||
template = config.get("handlebars_template")
|
||||
if not isinstance(template, str) or not template.strip():
|
||||
return ChartGenerationError(
|
||||
error_type="invalid_handlebars_template",
|
||||
message="Handlebars template must be a non-empty string",
|
||||
details=(
|
||||
"The 'handlebars_template' field must be a non-empty string "
|
||||
"containing valid Handlebars HTML template markup"
|
||||
),
|
||||
suggestions=[
|
||||
"Ensure handlebars_template is a non-empty string",
|
||||
"Example: '<ul>{{#each data}}<li>{{this.name}}</li>"
|
||||
"{{/each}}</ul>'",
|
||||
],
|
||||
error_code="INVALID_HANDLEBARS_TEMPLATE",
|
||||
)
|
||||
|
||||
query_mode = config.get("query_mode", "aggregate")
|
||||
if query_mode not in ("aggregate", "raw"):
|
||||
return ChartGenerationError(
|
||||
error_type="invalid_query_mode",
|
||||
message="Invalid query_mode for handlebars chart",
|
||||
details="query_mode must be either 'aggregate' or 'raw'",
|
||||
suggestions=[
|
||||
"Use 'aggregate' for aggregated data (default)",
|
||||
"Use 'raw' for individual rows",
|
||||
],
|
||||
error_code="INVALID_QUERY_MODE",
|
||||
)
|
||||
|
||||
if query_mode == "raw" and not config.get("columns"):
|
||||
return ChartGenerationError(
|
||||
error_type="missing_raw_columns",
|
||||
message="Handlebars chart in 'raw' mode requires 'columns'",
|
||||
details=(
|
||||
"When query_mode is 'raw', you must specify which columns "
|
||||
"to include in the query results"
|
||||
),
|
||||
suggestions=[
|
||||
"Add 'columns': [{'name': 'column_name'}] for raw mode",
|
||||
"Or use query_mode='aggregate' with 'metrics' and optional 'groupby'", # noqa: E501
|
||||
],
|
||||
error_code="MISSING_RAW_COLUMNS",
|
||||
)
|
||||
|
||||
if query_mode == "aggregate" and not config.get("metrics"):
|
||||
return ChartGenerationError(
|
||||
error_type="missing_aggregate_metrics",
|
||||
message="Handlebars chart in 'aggregate' mode requires 'metrics'",
|
||||
details=(
|
||||
"When query_mode is 'aggregate' (default), you must specify "
|
||||
"at least one metric with an aggregate function"
|
||||
),
|
||||
suggestions=[
|
||||
"Add 'metrics': [{'name': 'column', 'aggregate': 'SUM'}]",
|
||||
"Or use query_mode='raw' with 'columns' for individual rows",
|
||||
],
|
||||
error_code="MISSING_AGGREGATE_METRICS",
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
def extract_column_refs(self, config: Any) -> list[ColumnRef]:
|
||||
from superset.mcp_service.chart.schemas import HandlebarsChartConfig
|
||||
|
||||
if not isinstance(config, HandlebarsChartConfig):
|
||||
return []
|
||||
refs: list[ColumnRef] = []
|
||||
if config.columns:
|
||||
refs.extend(config.columns)
|
||||
if config.metrics:
|
||||
refs.extend(config.metrics)
|
||||
if config.groupby:
|
||||
refs.extend(config.groupby)
|
||||
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]:
|
||||
from superset.mcp_service.chart.chart_utils import map_handlebars_config
|
||||
|
||||
return map_handlebars_config(config)
|
||||
|
||||
def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any:
|
||||
from superset.mcp_service.chart.schemas import HandlebarsChartConfig
|
||||
from superset.mcp_service.chart.validation.dataset_validator import (
|
||||
DatasetValidator,
|
||||
)
|
||||
|
||||
config_dict = config.model_dump()
|
||||
|
||||
def _norm_list(key: str) -> None:
|
||||
if config_dict.get(key):
|
||||
for col in config_dict[key]:
|
||||
if not col.get("saved_metric"):
|
||||
col["name"] = DatasetValidator._get_canonical_column_name(
|
||||
col["name"], dataset_context
|
||||
)
|
||||
|
||||
_norm_list("columns")
|
||||
_norm_list("metrics")
|
||||
_norm_list("groupby")
|
||||
DatasetValidator._normalize_filters(config_dict, dataset_context)
|
||||
return HandlebarsChartConfig.model_validate(config_dict)
|
||||
136
superset/mcp_service/chart/plugins/mixed_timeseries.py
Normal file
136
superset/mcp_service/chart/plugins/mixed_timeseries.py
Normal file
@@ -0,0 +1,136 @@
|
||||
# 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.
|
||||
|
||||
"""Mixed timeseries chart type plugin."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from superset.mcp_service.chart.plugin import BaseChartPlugin
|
||||
from superset.mcp_service.chart.schemas import ColumnRef
|
||||
from superset.mcp_service.common.error_schemas import ChartGenerationError
|
||||
|
||||
|
||||
class MixedTimeseriesChartPlugin(BaseChartPlugin):
|
||||
"""Plugin for mixed_timeseries chart type."""
|
||||
|
||||
chart_type = "mixed_timeseries"
|
||||
|
||||
def pre_validate(
|
||||
self,
|
||||
config: dict[str, Any],
|
||||
) -> ChartGenerationError | None:
|
||||
missing_fields = []
|
||||
|
||||
if "x" not in config:
|
||||
missing_fields.append("'x' (X-axis temporal column)")
|
||||
if "y" not in config:
|
||||
missing_fields.append("'y' (primary Y-axis metrics)")
|
||||
if "y_secondary" not in config:
|
||||
missing_fields.append("'y_secondary' (secondary Y-axis metrics)")
|
||||
|
||||
if missing_fields:
|
||||
return ChartGenerationError(
|
||||
error_type="missing_mixed_timeseries_fields",
|
||||
message=(
|
||||
f"Mixed timeseries chart missing required fields: "
|
||||
f"{', '.join(missing_fields)}"
|
||||
),
|
||||
details=(
|
||||
"Mixed timeseries charts require an x-axis, primary metrics, "
|
||||
"and secondary metrics"
|
||||
),
|
||||
suggestions=[
|
||||
"Add 'x' field: {'name': 'date_column'}",
|
||||
"Add 'y' field: [{'name': 'revenue', 'aggregate': 'SUM'}]",
|
||||
"Add 'y_secondary': [{'name': 'orders', 'aggregate': 'COUNT'}]",
|
||||
"Optional: 'primary_kind' and 'secondary_kind' for chart types",
|
||||
],
|
||||
error_code="MISSING_MIXED_TIMESERIES_FIELDS",
|
||||
)
|
||||
|
||||
for field_name in ["y", "y_secondary"]:
|
||||
if not isinstance(config.get(field_name, []), list):
|
||||
return ChartGenerationError(
|
||||
error_type=f"invalid_{field_name}_format",
|
||||
message=f"'{field_name}' must be a list of metrics",
|
||||
details=(
|
||||
f"The '{field_name}' field must be an array of metric "
|
||||
"specifications"
|
||||
),
|
||||
suggestions=[
|
||||
f"Wrap in array: '{field_name}': "
|
||||
"[{'name': 'col', 'aggregate': 'SUM'}]",
|
||||
],
|
||||
error_code=f"INVALID_{field_name.upper()}_FORMAT",
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
def extract_column_refs(self, config: Any) -> list[ColumnRef]:
|
||||
from superset.mcp_service.chart.schemas import MixedTimeseriesChartConfig
|
||||
|
||||
if not isinstance(config, MixedTimeseriesChartConfig):
|
||||
return []
|
||||
refs: list[ColumnRef] = [config.x]
|
||||
refs.extend(config.y)
|
||||
refs.extend(config.y_secondary)
|
||||
if config.group_by:
|
||||
refs.extend(config.group_by)
|
||||
if config.group_by_secondary:
|
||||
refs.extend(config.group_by_secondary)
|
||||
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]:
|
||||
from superset.mcp_service.chart.chart_utils import map_mixed_timeseries_config
|
||||
|
||||
return map_mixed_timeseries_config(config, dataset_id=dataset_id)
|
||||
|
||||
def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any:
|
||||
from superset.mcp_service.chart.schemas import MixedTimeseriesChartConfig
|
||||
from superset.mcp_service.chart.validation.dataset_validator import (
|
||||
DatasetValidator,
|
||||
)
|
||||
|
||||
config_dict = config.model_dump()
|
||||
|
||||
def _norm_single(key: str) -> None:
|
||||
if config_dict.get(key):
|
||||
config_dict[key]["name"] = DatasetValidator._get_canonical_column_name(
|
||||
config_dict[key]["name"], dataset_context
|
||||
)
|
||||
|
||||
def _norm_list(key: str) -> None:
|
||||
if config_dict.get(key):
|
||||
for col in config_dict[key]:
|
||||
col["name"] = DatasetValidator._get_canonical_column_name(
|
||||
col["name"], dataset_context
|
||||
)
|
||||
|
||||
_norm_single("x")
|
||||
_norm_list("y")
|
||||
_norm_list("y_secondary")
|
||||
_norm_list("group_by")
|
||||
_norm_list("group_by_secondary")
|
||||
DatasetValidator._normalize_filters(config_dict, dataset_context)
|
||||
return MixedTimeseriesChartConfig.model_validate(config_dict)
|
||||
102
superset/mcp_service/chart/plugins/pie.py
Normal file
102
superset/mcp_service/chart/plugins/pie.py
Normal file
@@ -0,0 +1,102 @@
|
||||
# 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.
|
||||
|
||||
"""Pie chart type plugin."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from superset.mcp_service.chart.plugin import BaseChartPlugin
|
||||
from superset.mcp_service.chart.schemas import ColumnRef
|
||||
from superset.mcp_service.common.error_schemas import ChartGenerationError
|
||||
|
||||
|
||||
class PieChartPlugin(BaseChartPlugin):
|
||||
"""Plugin for pie chart type."""
|
||||
|
||||
chart_type = "pie"
|
||||
|
||||
def pre_validate(
|
||||
self,
|
||||
config: dict[str, Any],
|
||||
) -> ChartGenerationError | None:
|
||||
missing_fields = []
|
||||
|
||||
if "dimension" not in config:
|
||||
missing_fields.append("'dimension' (category column for slices)")
|
||||
if "metric" not in config:
|
||||
missing_fields.append("'metric' (value metric for slice sizes)")
|
||||
|
||||
if missing_fields:
|
||||
return ChartGenerationError(
|
||||
error_type="missing_pie_fields",
|
||||
message=(
|
||||
f"Pie chart missing required fields: {', '.join(missing_fields)}"
|
||||
),
|
||||
details=(
|
||||
"Pie charts require a dimension (categories) and a metric (values)"
|
||||
),
|
||||
suggestions=[
|
||||
"Add 'dimension' field: {'name': 'category_column'}",
|
||||
"Add 'metric' field: {'name': 'value_column', 'aggregate': 'SUM'}",
|
||||
"Example: {'chart_type': 'pie', 'dimension': {'name': 'product'}, "
|
||||
"'metric': {'name': 'revenue', 'aggregate': 'SUM'}}",
|
||||
],
|
||||
error_code="MISSING_PIE_FIELDS",
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
def extract_column_refs(self, config: Any) -> list[ColumnRef]:
|
||||
from superset.mcp_service.chart.schemas import PieChartConfig
|
||||
|
||||
if not isinstance(config, PieChartConfig):
|
||||
return []
|
||||
refs: list[ColumnRef] = [config.dimension, 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]:
|
||||
from superset.mcp_service.chart.chart_utils import map_pie_config
|
||||
|
||||
return map_pie_config(config)
|
||||
|
||||
def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any:
|
||||
from superset.mcp_service.chart.schemas import PieChartConfig
|
||||
from superset.mcp_service.chart.validation.dataset_validator import (
|
||||
DatasetValidator,
|
||||
)
|
||||
|
||||
config_dict = config.model_dump()
|
||||
|
||||
if config_dict.get("dimension"):
|
||||
config_dict["dimension"]["name"] = (
|
||||
DatasetValidator._get_canonical_column_name(
|
||||
config_dict["dimension"]["name"], dataset_context
|
||||
)
|
||||
)
|
||||
if config_dict.get("metric") and not config_dict["metric"].get("saved_metric"):
|
||||
config_dict["metric"]["name"] = DatasetValidator._get_canonical_column_name(
|
||||
config_dict["metric"]["name"], dataset_context
|
||||
)
|
||||
DatasetValidator._normalize_filters(config_dict, dataset_context)
|
||||
return PieChartConfig.model_validate(config_dict)
|
||||
125
superset/mcp_service/chart/plugins/pivot_table.py
Normal file
125
superset/mcp_service/chart/plugins/pivot_table.py
Normal file
@@ -0,0 +1,125 @@
|
||||
# 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.
|
||||
|
||||
"""Pivot table chart type plugin."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from superset.mcp_service.chart.plugin import BaseChartPlugin
|
||||
from superset.mcp_service.chart.schemas import ColumnRef
|
||||
from superset.mcp_service.common.error_schemas import ChartGenerationError
|
||||
|
||||
|
||||
class PivotTableChartPlugin(BaseChartPlugin):
|
||||
"""Plugin for pivot_table chart type."""
|
||||
|
||||
chart_type = "pivot_table"
|
||||
|
||||
def pre_validate(
|
||||
self,
|
||||
config: dict[str, Any],
|
||||
) -> ChartGenerationError | None:
|
||||
missing_fields = []
|
||||
|
||||
if "rows" not in config:
|
||||
missing_fields.append("'rows' (row grouping columns)")
|
||||
if "metrics" not in config:
|
||||
missing_fields.append("'metrics' (aggregation metrics)")
|
||||
|
||||
if missing_fields:
|
||||
return ChartGenerationError(
|
||||
error_type="missing_pivot_fields",
|
||||
message=(
|
||||
f"Pivot table missing required fields: {', '.join(missing_fields)}"
|
||||
),
|
||||
details="Pivot tables require row groupings and metrics",
|
||||
suggestions=[
|
||||
"Add 'rows' field: [{'name': 'category'}]",
|
||||
"Add 'metrics' field: [{'name': 'sales', 'aggregate': 'SUM'}]",
|
||||
"Optional 'columns' for cross-tabulation: [{'name': 'region'}]",
|
||||
],
|
||||
error_code="MISSING_PIVOT_FIELDS",
|
||||
)
|
||||
|
||||
if not isinstance(config.get("rows", []), list):
|
||||
return ChartGenerationError(
|
||||
error_type="invalid_rows_format",
|
||||
message="Rows must be a list of columns",
|
||||
details="The 'rows' field must be an array of column specifications",
|
||||
suggestions=[
|
||||
"Wrap row columns in array: 'rows': [{'name': 'category'}]",
|
||||
],
|
||||
error_code="INVALID_ROWS_FORMAT",
|
||||
)
|
||||
|
||||
if not isinstance(config.get("metrics", []), list):
|
||||
return ChartGenerationError(
|
||||
error_type="invalid_metrics_format",
|
||||
message="Metrics must be a list",
|
||||
details="The 'metrics' field must be an array of metric specifications",
|
||||
suggestions=[
|
||||
"Wrap metrics in array: 'metrics': [{'name': 'sales', "
|
||||
"'aggregate': 'SUM'}]",
|
||||
],
|
||||
error_code="INVALID_METRICS_FORMAT",
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
def extract_column_refs(self, config: Any) -> list[ColumnRef]:
|
||||
from superset.mcp_service.chart.schemas import PivotTableChartConfig
|
||||
|
||||
if not isinstance(config, PivotTableChartConfig):
|
||||
return []
|
||||
refs: list[ColumnRef] = list(config.rows)
|
||||
refs.extend(config.metrics)
|
||||
if config.columns:
|
||||
refs.extend(config.columns)
|
||||
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]:
|
||||
from superset.mcp_service.chart.chart_utils import map_pivot_table_config
|
||||
|
||||
return map_pivot_table_config(config)
|
||||
|
||||
def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any:
|
||||
from superset.mcp_service.chart.schemas import PivotTableChartConfig
|
||||
from superset.mcp_service.chart.validation.dataset_validator import (
|
||||
DatasetValidator,
|
||||
)
|
||||
|
||||
config_dict = config.model_dump()
|
||||
|
||||
def _norm_col_list(key: str) -> None:
|
||||
if config_dict.get(key):
|
||||
for col in config_dict[key]:
|
||||
col["name"] = DatasetValidator._get_canonical_column_name(
|
||||
col["name"], dataset_context
|
||||
)
|
||||
|
||||
_norm_col_list("rows")
|
||||
_norm_col_list("metrics")
|
||||
_norm_col_list("columns")
|
||||
DatasetValidator._normalize_filters(config_dict, dataset_context)
|
||||
return PivotTableChartConfig.model_validate(config_dict)
|
||||
96
superset/mcp_service/chart/plugins/table.py
Normal file
96
superset/mcp_service/chart/plugins/table.py
Normal file
@@ -0,0 +1,96 @@
|
||||
# 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.
|
||||
|
||||
"""Table chart type plugin."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from superset.mcp_service.chart.plugin import BaseChartPlugin
|
||||
from superset.mcp_service.chart.schemas import ColumnRef
|
||||
from superset.mcp_service.common.error_schemas import ChartGenerationError
|
||||
|
||||
|
||||
class TableChartPlugin(BaseChartPlugin):
|
||||
"""Plugin for table chart type."""
|
||||
|
||||
chart_type = "table"
|
||||
|
||||
def pre_validate(
|
||||
self,
|
||||
config: dict[str, Any],
|
||||
) -> ChartGenerationError | None:
|
||||
if "columns" not in config:
|
||||
return ChartGenerationError(
|
||||
error_type="missing_columns",
|
||||
message="Table chart missing required field: columns",
|
||||
details=(
|
||||
"Table charts require a 'columns' array to specify which "
|
||||
"columns to display"
|
||||
),
|
||||
suggestions=[
|
||||
"Add 'columns' field with array of column specifications",
|
||||
"Example: 'columns': [{'name': 'product'}, {'name': 'sales', "
|
||||
"'aggregate': 'SUM'}]",
|
||||
"Each column can have optional 'aggregate' for metrics",
|
||||
],
|
||||
error_code="MISSING_COLUMNS",
|
||||
)
|
||||
|
||||
if not isinstance(config.get("columns", []), list):
|
||||
return ChartGenerationError(
|
||||
error_type="invalid_columns_format",
|
||||
message="Columns must be a list",
|
||||
details="The 'columns' field must be an array of column specifications",
|
||||
suggestions=[
|
||||
"Ensure columns is an array: 'columns': [...]",
|
||||
"Each column should be an object with 'name' field",
|
||||
],
|
||||
error_code="INVALID_COLUMNS_FORMAT",
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
def extract_column_refs(self, config: Any) -> list[ColumnRef]:
|
||||
from superset.mcp_service.chart.schemas import TableChartConfig
|
||||
|
||||
if not isinstance(config, TableChartConfig):
|
||||
return []
|
||||
refs: list[ColumnRef] = list(config.columns)
|
||||
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]:
|
||||
from superset.mcp_service.chart.chart_utils import map_table_config
|
||||
|
||||
return map_table_config(config)
|
||||
|
||||
def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any:
|
||||
from superset.mcp_service.chart.schemas import TableChartConfig
|
||||
from superset.mcp_service.chart.validation.dataset_validator import (
|
||||
DatasetValidator,
|
||||
)
|
||||
|
||||
config_dict = config.model_dump()
|
||||
DatasetValidator._normalize_table_config(config_dict, dataset_context)
|
||||
DatasetValidator._normalize_filters(config_dict, dataset_context)
|
||||
return TableChartConfig.model_validate(config_dict)
|
||||
149
superset/mcp_service/chart/plugins/xy.py
Normal file
149
superset/mcp_service/chart/plugins/xy.py
Normal file
@@ -0,0 +1,149 @@
|
||||
# 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.
|
||||
|
||||
"""XY chart type plugin (line, bar, area, scatter)."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any
|
||||
|
||||
from superset.mcp_service.chart.plugin import BaseChartPlugin
|
||||
from superset.mcp_service.chart.schemas import ColumnRef
|
||||
from superset.mcp_service.common.error_schemas import ChartGenerationError
|
||||
|
||||
|
||||
class XYChartPlugin(BaseChartPlugin):
|
||||
"""Plugin for xy chart type (line, bar, area, scatter)."""
|
||||
|
||||
chart_type = "xy"
|
||||
|
||||
def pre_validate(
|
||||
self,
|
||||
config: dict[str, Any],
|
||||
) -> ChartGenerationError | None:
|
||||
# x is optional — defaults to dataset's main_dttm_col in map_xy_config
|
||||
if "y" not in config:
|
||||
return ChartGenerationError(
|
||||
error_type="missing_xy_fields",
|
||||
message="XY chart missing required field: 'y' (Y-axis metrics)",
|
||||
details=(
|
||||
"XY charts require Y-axis (metrics) specifications. "
|
||||
"X-axis is optional and defaults to the dataset's primary "
|
||||
"datetime column when omitted."
|
||||
),
|
||||
suggestions=[
|
||||
"Add 'y' field: [{'name': 'metric_column', 'aggregate': 'SUM'}]",
|
||||
"Example: {'chart_type': 'xy', 'x': {'name': 'date'}, "
|
||||
"'y': [{'name': 'sales', 'aggregate': 'SUM'}]}",
|
||||
],
|
||||
error_code="MISSING_XY_FIELDS",
|
||||
)
|
||||
|
||||
if not isinstance(config.get("y", []), list):
|
||||
return ChartGenerationError(
|
||||
error_type="invalid_y_format",
|
||||
message="Y-axis must be a list of metrics",
|
||||
details="The 'y' field must be an array of metric specifications",
|
||||
suggestions=[
|
||||
"Wrap Y-axis metric in array: 'y': [{'name': 'column', "
|
||||
"'aggregate': 'SUM'}]",
|
||||
"Multiple metrics supported: 'y': [metric1, metric2, ...]",
|
||||
],
|
||||
error_code="INVALID_Y_FORMAT",
|
||||
)
|
||||
|
||||
return None
|
||||
|
||||
def extract_column_refs(self, config: Any) -> list[ColumnRef]:
|
||||
from superset.mcp_service.chart.schemas import XYChartConfig
|
||||
|
||||
if not isinstance(config, XYChartConfig):
|
||||
return []
|
||||
refs: list[ColumnRef] = []
|
||||
if config.x is not None:
|
||||
refs.append(config.x)
|
||||
refs.extend(config.y)
|
||||
if config.group_by:
|
||||
refs.extend(config.group_by)
|
||||
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]:
|
||||
from superset.mcp_service.chart.chart_utils import map_xy_config
|
||||
|
||||
return map_xy_config(config, dataset_id=dataset_id)
|
||||
|
||||
def normalize_column_refs(self, config: Any, dataset_context: Any) -> Any:
|
||||
from superset.mcp_service.chart.schemas import XYChartConfig
|
||||
from superset.mcp_service.chart.validation.dataset_validator import (
|
||||
DatasetValidator,
|
||||
)
|
||||
|
||||
config_dict = config.model_dump()
|
||||
DatasetValidator._normalize_xy_config(config_dict, dataset_context)
|
||||
DatasetValidator._normalize_filters(config_dict, dataset_context)
|
||||
return XYChartConfig.model_validate(config_dict)
|
||||
|
||||
def get_runtime_warnings(self, config: Any, dataset_id: int | str) -> list[str]:
|
||||
"""Return format-compatibility and cardinality warnings for XY charts."""
|
||||
import logging
|
||||
|
||||
from superset.mcp_service.chart.schemas import XYChartConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
if not isinstance(config, XYChartConfig):
|
||||
return []
|
||||
|
||||
warnings: list[str] = []
|
||||
|
||||
try:
|
||||
from superset.mcp_service.chart.validation.runtime.format_validator import (
|
||||
FormatTypeValidator,
|
||||
)
|
||||
|
||||
_valid, format_warnings = FormatTypeValidator.validate_format_compatibility(
|
||||
config
|
||||
)
|
||||
if format_warnings:
|
||||
warnings.extend(format_warnings)
|
||||
except Exception as exc:
|
||||
logger.warning("XY format validation failed: %s", exc)
|
||||
|
||||
try:
|
||||
from superset.mcp_service.chart.validation.runtime.cardinality_validator import ( # noqa: E501
|
||||
CardinalityValidator,
|
||||
)
|
||||
|
||||
chart_kind = config.kind if hasattr(config, "kind") else "default"
|
||||
group_by_col = config.group_by[0].name if config.group_by else None
|
||||
if config.x is not None:
|
||||
_ok, card_info = CardinalityValidator.check_cardinality(
|
||||
dataset_id=dataset_id,
|
||||
x_column=config.x.name,
|
||||
chart_type=chart_kind,
|
||||
group_by_column=group_by_col,
|
||||
)
|
||||
if not _ok and card_info:
|
||||
warnings.extend(card_info.get("warnings", []))
|
||||
except Exception as exc:
|
||||
logger.warning("XY cardinality validation failed: %s", exc)
|
||||
|
||||
return warnings
|
||||
90
superset/mcp_service/chart/registry.py
Normal file
90
superset/mcp_service/chart/registry.py
Normal file
@@ -0,0 +1,90 @@
|
||||
# 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.
|
||||
|
||||
"""
|
||||
ChartTypeRegistry — central registry mapping chart_type strings to plugins.
|
||||
|
||||
Replaces the four previously-scattered dispatch locations:
|
||||
- schema_validator.py: chart_type_validators dict
|
||||
- dataset_validator.py: isinstance branches in _extract_column_references()
|
||||
- chart_utils.py: if/elif chain in map_config_to_form_data()
|
||||
- dataset_validator.py: isinstance branches in normalize_column_names()
|
||||
|
||||
Usage::
|
||||
|
||||
from superset.mcp_service.chart.registry import get_registry
|
||||
|
||||
plugin = get_registry().get("xy")
|
||||
if plugin is None:
|
||||
raise ValueError("Unknown chart type: xy")
|
||||
form_data = plugin.to_form_data(config, dataset_id)
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
if TYPE_CHECKING:
|
||||
from superset.mcp_service.chart.plugin import ChartTypePlugin
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_REGISTRY: dict[str, "ChartTypePlugin"] = {}
|
||||
|
||||
|
||||
def register(plugin: "ChartTypePlugin") -> None:
|
||||
"""Register a chart type plugin in the global registry."""
|
||||
if plugin.chart_type in _REGISTRY:
|
||||
logger.warning(
|
||||
"Overwriting existing plugin for chart_type=%r", plugin.chart_type
|
||||
)
|
||||
_REGISTRY[plugin.chart_type] = plugin
|
||||
logger.debug("Registered chart plugin: %r", plugin.chart_type)
|
||||
|
||||
|
||||
def get(chart_type: str) -> "ChartTypePlugin | None":
|
||||
"""Return the plugin for a given chart_type, or None if not registered."""
|
||||
return _REGISTRY.get(chart_type)
|
||||
|
||||
|
||||
def all_types() -> list[str]:
|
||||
"""Return all registered chart type strings in insertion order."""
|
||||
return list(_REGISTRY.keys())
|
||||
|
||||
|
||||
def is_registered(chart_type: str) -> bool:
|
||||
"""Return True if chart_type has a registered plugin."""
|
||||
return chart_type in _REGISTRY
|
||||
|
||||
|
||||
def get_registry() -> "_RegistryProxy":
|
||||
"""Return a proxy object for registry access (convenience wrapper)."""
|
||||
return _RegistryProxy()
|
||||
|
||||
|
||||
class _RegistryProxy:
|
||||
"""Thin proxy exposing registry functions as instance methods."""
|
||||
|
||||
def get(self, chart_type: str) -> "ChartTypePlugin | None":
|
||||
return _REGISTRY.get(chart_type)
|
||||
|
||||
def all_types(self) -> list[str]:
|
||||
return list(_REGISTRY.keys())
|
||||
|
||||
def is_registered(self, chart_type: str) -> bool:
|
||||
return chart_type in _REGISTRY
|
||||
@@ -105,7 +105,8 @@ async def generate_chart( # noqa: C901
|
||||
- Set save_chart=True to permanently save the chart
|
||||
- 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 (either 'xy' or 'table')
|
||||
- MUST include chart_type in config (one of: 'xy', 'table', 'pie',
|
||||
'pivot_table', 'mixed_timeseries', 'handlebars', 'big_number')
|
||||
|
||||
IMPORTANT: The 'chart_type' field in the config is a DISCRIMINATOR that determines
|
||||
which chart configuration schema to use. It MUST be included and MUST match the
|
||||
@@ -117,6 +118,21 @@ async def generate_chart( # noqa: C901
|
||||
- Use chart_type='table' for tabular visualizations
|
||||
Required fields: columns
|
||||
|
||||
- Use chart_type='pie' for pie/donut charts
|
||||
Required fields: dimension, metric
|
||||
|
||||
- Use chart_type='pivot_table' for pivot table visualizations
|
||||
Required fields: rows, metrics
|
||||
|
||||
- Use chart_type='mixed_timeseries' for dual-axis time-series charts
|
||||
Required fields: x, y, y_secondary
|
||||
|
||||
- Use chart_type='handlebars' for custom template-based visualizations
|
||||
Required fields: handlebars_template
|
||||
|
||||
- Use chart_type='big_number' for single KPI metric displays
|
||||
Required fields: metric
|
||||
|
||||
Example usage for XY chart:
|
||||
```json
|
||||
{
|
||||
|
||||
@@ -25,14 +25,8 @@ import logging
|
||||
from typing import Any, Dict, List, Tuple
|
||||
|
||||
from superset.mcp_service.chart.schemas import (
|
||||
BigNumberChartConfig,
|
||||
ChartConfig,
|
||||
ColumnRef,
|
||||
HandlebarsChartConfig,
|
||||
MixedTimeseriesChartConfig,
|
||||
PieChartConfig,
|
||||
PivotTableChartConfig,
|
||||
TableChartConfig,
|
||||
XYChartConfig,
|
||||
)
|
||||
from superset.mcp_service.common.error_schemas import (
|
||||
ChartGenerationError,
|
||||
@@ -58,7 +52,7 @@ class DatasetValidator:
|
||||
|
||||
@staticmethod
|
||||
def validate_against_dataset(
|
||||
config: Any,
|
||||
config: ChartConfig,
|
||||
dataset_id: int | str,
|
||||
dataset_context: DatasetContext | None = None,
|
||||
) -> Tuple[bool, ChartGenerationError | None]:
|
||||
@@ -269,59 +263,30 @@ class DatasetValidator:
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
def _extract_column_references(config: Any) -> List[ColumnRef]: # noqa: C901
|
||||
"""Extract all column references from a chart configuration.
|
||||
def _extract_column_references(
|
||||
config: ChartConfig,
|
||||
) -> List[ColumnRef]:
|
||||
"""Extract all column references from configuration via the plugin registry.
|
||||
|
||||
Covers every supported ``ChartConfig`` variant so fast-path tools
|
||||
(``generate_explore_link``, ``update_chart_preview``) that only run
|
||||
Tier-1 validation still catch bad column refs in pie / pivot table /
|
||||
mixed timeseries / handlebars / big number charts — not just XY and
|
||||
table.
|
||||
Previously only handled TableChartConfig and XYChartConfig, causing
|
||||
5 of 7 chart types to silently skip column validation. Now delegates
|
||||
to the plugin for each chart type so all types are covered.
|
||||
"""
|
||||
refs: List[ColumnRef] = []
|
||||
# Local import: plugins call DatasetValidator helpers from normalize_column_refs().
|
||||
# A top-level import of registry in dataset_validator would make loading this
|
||||
# module implicitly trigger plugin registration, creating a circular dependency.
|
||||
from superset.mcp_service.chart.registry import get_registry
|
||||
|
||||
if isinstance(config, TableChartConfig):
|
||||
refs.extend(config.columns)
|
||||
elif isinstance(config, XYChartConfig):
|
||||
if config.x is not None:
|
||||
refs.append(config.x)
|
||||
refs.extend(config.y)
|
||||
if config.group_by:
|
||||
refs.extend(config.group_by)
|
||||
elif isinstance(config, PieChartConfig):
|
||||
refs.append(config.dimension)
|
||||
refs.append(config.metric)
|
||||
elif isinstance(config, PivotTableChartConfig):
|
||||
refs.extend(config.rows)
|
||||
if config.columns:
|
||||
refs.extend(config.columns)
|
||||
refs.extend(config.metrics)
|
||||
elif isinstance(config, MixedTimeseriesChartConfig):
|
||||
refs.append(config.x)
|
||||
refs.extend(config.y)
|
||||
if config.group_by:
|
||||
refs.extend(config.group_by)
|
||||
refs.extend(config.y_secondary)
|
||||
if config.group_by_secondary:
|
||||
refs.extend(config.group_by_secondary)
|
||||
elif isinstance(config, HandlebarsChartConfig):
|
||||
if config.columns:
|
||||
refs.extend(config.columns)
|
||||
if config.groupby:
|
||||
refs.extend(config.groupby)
|
||||
if config.metrics:
|
||||
refs.extend(config.metrics)
|
||||
elif isinstance(config, BigNumberChartConfig):
|
||||
refs.append(config.metric)
|
||||
if config.temporal_column:
|
||||
refs.append(ColumnRef(name=config.temporal_column))
|
||||
chart_type = getattr(config, "chart_type", None)
|
||||
if chart_type is None:
|
||||
return []
|
||||
|
||||
# Filter columns (shared by every config type that defines ``filters``).
|
||||
if filters := getattr(config, "filters", None):
|
||||
for filter_config in filters:
|
||||
refs.append(ColumnRef(name=filter_config.column))
|
||||
plugin = get_registry().get(chart_type)
|
||||
if plugin is None:
|
||||
logger.warning("No plugin registered for chart_type=%r", chart_type)
|
||||
return []
|
||||
|
||||
return refs
|
||||
return plugin.extract_column_refs(config)
|
||||
|
||||
@staticmethod
|
||||
def _column_exists(column_name: str, dataset_context: DatasetContext) -> bool:
|
||||
@@ -433,10 +398,10 @@ class DatasetValidator:
|
||||
|
||||
@staticmethod
|
||||
def normalize_column_names(
|
||||
config: TableChartConfig | XYChartConfig,
|
||||
config: ChartConfig,
|
||||
dataset_id: int | str,
|
||||
dataset_context: DatasetContext | None = None,
|
||||
) -> TableChartConfig | XYChartConfig:
|
||||
) -> ChartConfig:
|
||||
"""
|
||||
Normalize column names in config to match the canonical dataset column names.
|
||||
|
||||
@@ -445,6 +410,9 @@ class DatasetValidator:
|
||||
(e.g., 'OrderDate'). The frontend performs case-sensitive comparisons,
|
||||
so we need to ensure column names match exactly.
|
||||
|
||||
Previously only XYChartConfig and TableChartConfig were normalized; now
|
||||
all 7 chart types are handled via the plugin registry.
|
||||
|
||||
Args:
|
||||
config: Chart configuration with column references
|
||||
dataset_id: Dataset ID to get canonical column names from
|
||||
@@ -459,22 +427,20 @@ class DatasetValidator:
|
||||
if not dataset_context:
|
||||
return config
|
||||
|
||||
# Create a mutable copy of the config
|
||||
config_dict = config.model_dump()
|
||||
from superset.mcp_service.chart.registry import get_registry
|
||||
|
||||
# Normalize based on config type
|
||||
if isinstance(config, XYChartConfig):
|
||||
DatasetValidator._normalize_xy_config(config_dict, dataset_context)
|
||||
elif isinstance(config, TableChartConfig):
|
||||
DatasetValidator._normalize_table_config(config_dict, dataset_context)
|
||||
chart_type = getattr(config, "chart_type", None)
|
||||
if chart_type is None:
|
||||
return config
|
||||
|
||||
# Normalize filter columns (common to both config types)
|
||||
DatasetValidator._normalize_filters(config_dict, dataset_context)
|
||||
plugin = get_registry().get(chart_type)
|
||||
if plugin is None:
|
||||
logger.warning(
|
||||
"No plugin for chart_type=%r; skipping column normalization", chart_type
|
||||
)
|
||||
return config
|
||||
|
||||
# Reconstruct the config with normalized names
|
||||
if isinstance(config, XYChartConfig):
|
||||
return XYChartConfig.model_validate(config_dict)
|
||||
return TableChartConfig.model_validate(config_dict)
|
||||
return plugin.normalize_column_refs(config, dataset_context)
|
||||
|
||||
@staticmethod
|
||||
def _get_column_suggestions(
|
||||
|
||||
@@ -23,10 +23,7 @@ Validates performance, compatibility, and user experience issues.
|
||||
import logging
|
||||
from typing import Any, Dict, List, Tuple
|
||||
|
||||
from superset.mcp_service.chart.schemas import (
|
||||
ChartConfig,
|
||||
XYChartConfig,
|
||||
)
|
||||
from superset.mcp_service.chart.schemas import ChartConfig
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -56,20 +53,10 @@ class RuntimeValidator:
|
||||
warnings: List[str] = []
|
||||
suggestions: List[str] = []
|
||||
|
||||
# Only check XY charts for format and cardinality issues
|
||||
if isinstance(config, XYChartConfig):
|
||||
# Format-type compatibility validation
|
||||
format_warnings = RuntimeValidator._validate_format_compatibility(config)
|
||||
if format_warnings:
|
||||
warnings.extend(format_warnings)
|
||||
|
||||
# Cardinality validation
|
||||
cardinality_warnings, cardinality_suggestions = (
|
||||
RuntimeValidator._validate_cardinality(config, dataset_id)
|
||||
)
|
||||
if cardinality_warnings:
|
||||
warnings.extend(cardinality_warnings)
|
||||
suggestions.extend(cardinality_suggestions)
|
||||
# Per-plugin runtime warnings (format, cardinality, etc.)
|
||||
plugin_warnings = RuntimeValidator._validate_plugin_runtime(config, dataset_id)
|
||||
if plugin_warnings:
|
||||
warnings.extend(plugin_warnings)
|
||||
|
||||
# Chart type appropriateness validation (for all chart types)
|
||||
type_warnings, type_suggestions = RuntimeValidator._validate_chart_type(
|
||||
@@ -98,61 +85,28 @@ class RuntimeValidator:
|
||||
return True, None
|
||||
|
||||
@staticmethod
|
||||
def _validate_format_compatibility(config: XYChartConfig) -> List[str]:
|
||||
"""Validate format-type compatibility."""
|
||||
warnings: List[str] = []
|
||||
def _validate_plugin_runtime(
|
||||
config: ChartConfig, dataset_id: int | str
|
||||
) -> List[str]:
|
||||
"""Delegate per-chart-type runtime warnings to the plugin registry.
|
||||
|
||||
Each plugin's get_runtime_warnings() method returns chart-type-specific
|
||||
warnings (e.g. format/cardinality for XY). The registry dispatch removes
|
||||
the previous isinstance(config, XYChartConfig) hardcoding.
|
||||
"""
|
||||
try:
|
||||
# Import here to avoid circular imports
|
||||
from .format_validator import FormatTypeValidator
|
||||
from superset.mcp_service.chart.registry import get_registry
|
||||
|
||||
is_valid, format_warnings = (
|
||||
FormatTypeValidator.validate_format_compatibility(config)
|
||||
)
|
||||
if format_warnings:
|
||||
warnings.extend(format_warnings)
|
||||
except ImportError:
|
||||
logger.warning("Format validator not available")
|
||||
except Exception as e:
|
||||
logger.warning("Format validation failed: %s", e)
|
||||
|
||||
return warnings
|
||||
|
||||
@staticmethod
|
||||
def _validate_cardinality(
|
||||
config: XYChartConfig, dataset_id: int | str
|
||||
) -> Tuple[List[str], List[str]]:
|
||||
"""Validate cardinality issues."""
|
||||
warnings: List[str] = []
|
||||
suggestions: List[str] = []
|
||||
|
||||
try:
|
||||
# Import here to avoid circular imports
|
||||
from .cardinality_validator import CardinalityValidator
|
||||
|
||||
# Determine chart type for cardinality thresholds
|
||||
chart_type = config.kind if hasattr(config, "kind") else "default"
|
||||
|
||||
# Check X-axis cardinality
|
||||
if config.x is None or config.x.name is None:
|
||||
return warnings, suggestions
|
||||
is_ok, cardinality_info = CardinalityValidator.check_cardinality(
|
||||
dataset_id=dataset_id,
|
||||
x_column=config.x.name,
|
||||
chart_type=chart_type,
|
||||
group_by_column=(config.group_by[0].name if config.group_by else None),
|
||||
)
|
||||
|
||||
if not is_ok and cardinality_info:
|
||||
warnings.extend(cardinality_info.get("warnings", []))
|
||||
suggestions.extend(cardinality_info.get("suggestions", []))
|
||||
|
||||
except ImportError:
|
||||
logger.warning("Cardinality validator not available")
|
||||
except Exception as e:
|
||||
logger.warning("Cardinality validation failed: %s", e)
|
||||
|
||||
return warnings, suggestions
|
||||
chart_type = getattr(config, "chart_type", None)
|
||||
if chart_type is None:
|
||||
return []
|
||||
plugin = get_registry().get(chart_type)
|
||||
if plugin is None:
|
||||
return []
|
||||
return plugin.get_runtime_warnings(config, dataset_id)
|
||||
except Exception as exc:
|
||||
logger.warning("Plugin runtime validation failed: %s", exc)
|
||||
return []
|
||||
|
||||
@staticmethod
|
||||
def _validate_chart_type(
|
||||
|
||||
@@ -147,19 +147,13 @@ class SchemaValidator:
|
||||
chart_type: str,
|
||||
config: Dict[str, Any],
|
||||
) -> Tuple[bool, ChartGenerationError | None]:
|
||||
"""Validate chart type and dispatch to type-specific pre-validation."""
|
||||
chart_type_validators = {
|
||||
"xy": SchemaValidator._pre_validate_xy_config,
|
||||
"table": SchemaValidator._pre_validate_table_config,
|
||||
"pie": SchemaValidator._pre_validate_pie_config,
|
||||
"pivot_table": SchemaValidator._pre_validate_pivot_table_config,
|
||||
"mixed_timeseries": SchemaValidator._pre_validate_mixed_timeseries_config,
|
||||
"handlebars": SchemaValidator._pre_validate_handlebars_config,
|
||||
"big_number": SchemaValidator._pre_validate_big_number_config,
|
||||
}
|
||||
"""Validate chart type and dispatch to plugin pre-validation."""
|
||||
from superset.mcp_service.chart.registry import get_registry
|
||||
|
||||
if not isinstance(chart_type, str) or chart_type not in chart_type_validators:
|
||||
valid_types = ", ".join(chart_type_validators.keys())
|
||||
registry = get_registry()
|
||||
|
||||
if not isinstance(chart_type, str) or not registry.is_registered(chart_type):
|
||||
valid_types = ", ".join(registry.all_types())
|
||||
return False, ChartGenerationError(
|
||||
error_type="invalid_chart_type",
|
||||
message=f"Invalid chart_type: '{chart_type}'",
|
||||
@@ -178,7 +172,19 @@ class SchemaValidator:
|
||||
error_code="INVALID_CHART_TYPE",
|
||||
)
|
||||
|
||||
return chart_type_validators[chart_type](config)
|
||||
plugin = registry.get(chart_type)
|
||||
if plugin is None:
|
||||
return False, ChartGenerationError(
|
||||
error_type="invalid_chart_type",
|
||||
message=f"Chart type '{chart_type}' has no registered plugin",
|
||||
details="Internal error: chart type is listed but has no plugin",
|
||||
suggestions=["Use a supported chart_type"],
|
||||
error_code="INVALID_CHART_TYPE",
|
||||
)
|
||||
|
||||
if (error := plugin.pre_validate(config)) is not None:
|
||||
return False, error
|
||||
return True, None
|
||||
|
||||
@staticmethod
|
||||
def _pre_validate_xy_config(
|
||||
@@ -550,6 +556,110 @@ class SchemaValidator:
|
||||
|
||||
return True, None
|
||||
|
||||
# Per-chart-type error details used by _enhance_validation_error.
|
||||
# Keyed by chart_type discriminator value.
|
||||
_CHART_TYPE_ERROR_HINTS: Dict[str, Dict[str, Any]] = {
|
||||
"xy": {
|
||||
"error_type": "xy_validation_error",
|
||||
"message": "XY chart configuration validation failed",
|
||||
"details": "The XY chart configuration is missing required "
|
||||
"fields or has invalid structure",
|
||||
"suggestions": [
|
||||
"Ensure 'x' field exists with {'name': 'column_name'}",
|
||||
"Ensure 'y' is an array: [{'name': 'metric', 'aggregate': 'SUM'}]",
|
||||
"Check that all column names are strings",
|
||||
"Verify aggregate functions are valid: SUM, COUNT, AVG, MIN, MAX",
|
||||
],
|
||||
"error_code": "XY_VALIDATION_ERROR",
|
||||
},
|
||||
"table": {
|
||||
"error_type": "table_validation_error",
|
||||
"message": "Table chart configuration validation failed",
|
||||
"details": "The table chart configuration is missing required "
|
||||
"fields or has invalid structure",
|
||||
"suggestions": [
|
||||
"Ensure 'columns' field is an array of column specifications",
|
||||
"Each column needs {'name': 'column_name'}",
|
||||
"Optional: add 'aggregate' for metrics",
|
||||
"Example: 'columns': [{'name': 'product'}, "
|
||||
"{'name': 'sales', 'aggregate': 'SUM'}]",
|
||||
],
|
||||
"error_code": "TABLE_VALIDATION_ERROR",
|
||||
},
|
||||
"pie": {
|
||||
"error_type": "pie_validation_error",
|
||||
"message": "Pie chart configuration validation failed",
|
||||
"details": "The pie chart configuration is missing required "
|
||||
"fields or has invalid structure",
|
||||
"suggestions": [
|
||||
"Ensure 'dimension' field has 'name' for the slice label",
|
||||
"Ensure 'metric' field has 'name' and 'aggregate'",
|
||||
"Example: {'chart_type': 'pie', 'dimension': {'name': 'category'}, "
|
||||
"'metric': {'name': 'revenue', 'aggregate': 'SUM'}}",
|
||||
],
|
||||
"error_code": "PIE_VALIDATION_ERROR",
|
||||
},
|
||||
"pivot_table": {
|
||||
"error_type": "pivot_table_validation_error",
|
||||
"message": "Pivot table configuration validation failed",
|
||||
"details": "The pivot table configuration is missing required "
|
||||
"fields or has invalid structure",
|
||||
"suggestions": [
|
||||
"Ensure 'rows' field is an array of column specs",
|
||||
"Ensure 'metrics' field is an array with aggregate funcs",
|
||||
"Optional: add 'columns' for column grouping",
|
||||
"Example: {'chart_type': 'pivot_table', 'rows': [{'name': 'region'}], "
|
||||
"'metrics': [{'name': 'revenue', 'aggregate': 'SUM'}]}",
|
||||
],
|
||||
"error_code": "PIVOT_TABLE_VALIDATION_ERROR",
|
||||
},
|
||||
"mixed_timeseries": {
|
||||
"error_type": "mixed_timeseries_validation_error",
|
||||
"message": "Mixed timeseries chart configuration validation failed",
|
||||
"details": "The mixed timeseries configuration is missing "
|
||||
"required fields or has invalid structure",
|
||||
"suggestions": [
|
||||
"Ensure 'x' field has 'name' for the time axis column",
|
||||
"Ensure 'y' is an array of primary-axis metrics",
|
||||
"Ensure 'y_secondary' is an array of secondary-axis metrics",
|
||||
"Example: {'chart_type': 'mixed_timeseries', "
|
||||
"'x': {'name': 'order_date'}, "
|
||||
"'y': [{'name': 'revenue', 'aggregate': 'SUM'}], "
|
||||
"'y_secondary': [{'name': 'orders', 'aggregate': 'COUNT'}]}",
|
||||
],
|
||||
"error_code": "MIXED_TIMESERIES_VALIDATION_ERROR",
|
||||
},
|
||||
"handlebars": {
|
||||
"error_type": "handlebars_validation_error",
|
||||
"message": "Handlebars chart configuration validation failed",
|
||||
"details": "The handlebars chart configuration is missing "
|
||||
"required fields or has invalid structure",
|
||||
"suggestions": [
|
||||
"Ensure 'handlebars_template' is a non-empty string",
|
||||
"For aggregate mode: add 'metrics' with aggregate functions",
|
||||
"For raw mode: set 'query_mode': 'raw' and add 'columns'",
|
||||
"Example: {'chart_type': 'handlebars', "
|
||||
"'handlebars_template': '<ul>{{#each data}}<li>"
|
||||
"{{this.name}}</li>{{/each}}</ul>', "
|
||||
"'metrics': [{'name': 'sales', 'aggregate': 'SUM'}]}",
|
||||
],
|
||||
"error_code": "HANDLEBARS_VALIDATION_ERROR",
|
||||
},
|
||||
"big_number": {
|
||||
"error_type": "big_number_validation_error",
|
||||
"message": "Big Number chart configuration validation failed",
|
||||
"details": "The Big Number chart configuration is missing required "
|
||||
"fields or has invalid structure",
|
||||
"suggestions": [
|
||||
"Ensure 'metric' field has 'name' and 'aggregate'",
|
||||
"Example: 'metric': {'name': 'revenue', 'aggregate': 'SUM'}",
|
||||
"For trendline: add show_trendline=true and temporal_column='col'",
|
||||
"Without trendline: just provide the metric",
|
||||
],
|
||||
"error_code": "BIG_NUMBER_VALIDATION_ERROR",
|
||||
},
|
||||
}
|
||||
|
||||
@staticmethod
|
||||
def _enhance_validation_error(
|
||||
error: PydanticValidationError, request_data: Dict[str, Any]
|
||||
@@ -562,89 +672,22 @@ class SchemaValidator:
|
||||
if err.get("type") == "union_tag_invalid" or "discriminator" in str(
|
||||
err.get("ctx", {})
|
||||
):
|
||||
# This is the generic union error - provide better message
|
||||
config = request_data.get("config", {})
|
||||
chart_type = config.get("chart_type", "unknown")
|
||||
|
||||
if chart_type == "xy":
|
||||
return ChartGenerationError(
|
||||
error_type="xy_validation_error",
|
||||
message="XY chart configuration validation failed",
|
||||
details="The XY chart configuration is missing required "
|
||||
"fields or has invalid structure",
|
||||
suggestions=[
|
||||
"Ensure 'x' field exists with {'name': 'column_name'}",
|
||||
"Ensure 'y' field is an array: [{'name': 'metric', "
|
||||
"'aggregate': 'SUM'}]",
|
||||
"Check that all column names are strings",
|
||||
"Verify aggregate functions are valid: SUM, COUNT, AVG, "
|
||||
"MIN, MAX",
|
||||
],
|
||||
error_code="XY_VALIDATION_ERROR",
|
||||
)
|
||||
elif chart_type == "table":
|
||||
return ChartGenerationError(
|
||||
error_type="table_validation_error",
|
||||
message="Table chart configuration validation failed",
|
||||
details="The table chart configuration is missing required "
|
||||
"fields or has invalid structure",
|
||||
suggestions=[
|
||||
"Ensure 'columns' field is an array of column "
|
||||
"specifications",
|
||||
"Each column needs {'name': 'column_name'}",
|
||||
"Optional: add 'aggregate' for metrics",
|
||||
"Example: 'columns': [{'name': 'product'}, {'name': "
|
||||
"'sales', 'aggregate': 'SUM'}]",
|
||||
],
|
||||
error_code="TABLE_VALIDATION_ERROR",
|
||||
)
|
||||
elif chart_type == "handlebars":
|
||||
return ChartGenerationError(
|
||||
error_type="handlebars_validation_error",
|
||||
message="Handlebars chart configuration validation failed",
|
||||
details="The handlebars chart configuration is missing "
|
||||
"required fields or has invalid structure",
|
||||
suggestions=[
|
||||
"Ensure 'handlebars_template' is a non-empty string",
|
||||
"For aggregate mode: add 'metrics' with aggregate "
|
||||
"functions",
|
||||
"For raw mode: set 'query_mode': 'raw' and add 'columns'",
|
||||
"Example: {'chart_type': 'handlebars', "
|
||||
"'handlebars_template': '<ul>{{#each data}}<li>"
|
||||
"{{this.name}}</li>{{/each}}</ul>', "
|
||||
"'metrics': [{'name': 'sales', 'aggregate': 'SUM'}]}",
|
||||
],
|
||||
error_code="HANDLEBARS_VALIDATION_ERROR",
|
||||
)
|
||||
elif chart_type == "big_number":
|
||||
return ChartGenerationError(
|
||||
error_type="big_number_validation_error",
|
||||
message="Big Number chart configuration validation failed",
|
||||
details="The Big Number chart configuration is "
|
||||
"missing required fields or has invalid "
|
||||
"structure",
|
||||
suggestions=[
|
||||
"Ensure 'metric' field has 'name' and 'aggregate'",
|
||||
"Example: 'metric': {'name': 'revenue', "
|
||||
"'aggregate': 'SUM'}",
|
||||
"For trendline: add 'show_trendline': true "
|
||||
"and 'temporal_column': 'date_col'",
|
||||
"Without trendline: just provide the metric",
|
||||
],
|
||||
error_code="BIG_NUMBER_VALIDATION_ERROR",
|
||||
)
|
||||
chart_type = request_data.get("config", {}).get("chart_type", "")
|
||||
hint = SchemaValidator._CHART_TYPE_ERROR_HINTS.get(chart_type)
|
||||
if hint:
|
||||
return ChartGenerationError(**hint)
|
||||
|
||||
# Default enhanced error
|
||||
error_details = []
|
||||
for err in errors[:3]: # Show first 3 errors
|
||||
loc = " -> ".join(str(location) for location in err.get("loc", []))
|
||||
msg = err.get("msg", "Validation failed")
|
||||
error_details.append(f"{loc}: {msg}")
|
||||
error_details.append(f"{loc}: {msg}" if loc else msg)
|
||||
|
||||
return ChartGenerationError(
|
||||
error_type="validation_error",
|
||||
message="Chart configuration validation failed",
|
||||
details="; ".join(error_details),
|
||||
details="; ".join(error_details) or "Invalid chart configuration structure",
|
||||
suggestions=[
|
||||
"Check that all required fields are present",
|
||||
"Ensure field types match the schema",
|
||||
|
||||
Reference in New Issue
Block a user