aminghadersohi commented on code in PR #44148:
URL: https://github.com/apache/superset/pull/44148#discussion_r4190890506
##########
superset/mcp_service/chart/schemas.py:
##########
@@ -3738,6 +3738,187 @@ def validate_gantt_roles(self) -> "GanttChartConfig":
# Discriminated union for runtime validation (not exposed in JSON Schema)
+def _omit_inherited_descriptions(schema: dict[str, Any]) -> None:
+ """Drop field descriptions the published parent schema already carries.
+
+ The geographic references and filters only tighten bounds on the shared
+ ``ColumnRef``/``FilterConfig`` fields, whose descriptions appear in the
+ same tool schema.
+ """
+ for prop in schema.get("properties", {}).values():
+ prop.pop("description", None)
+
+
+class GeographicColumnRef(ColumnRef):
+ """Closed ColumnRef for geographic roles."""
+
+ model_config = ConfigDict(
+ extra="forbid",
+ populate_by_name=True,
+ strict=True,
+ json_schema_extra=_omit_inherited_descriptions,
+ )
+ dtype: str | None = Field(None, max_length=128)
+
+
+GeographicFilterValue = Annotated[str, Field(max_length=1000)] | int | float |
bool
+
+
+class GeographicFilterConfig(FilterConfig):
+ """Closed FilterConfig with bounded values."""
+
+ model_config = ConfigDict(
+ extra="forbid",
+ populate_by_name=True,
+ strict=True,
+ allow_inf_nan=False,
+ json_schema_extra=_omit_inherited_descriptions,
+ )
+ value: (
+ GeographicFilterValue
+ | Annotated[list[GeographicFilterValue], Field(max_length=1000)]
+ | None
+ ) = Field(None, validation_alias=AliasChoices("value", "val"))
+
+
+class GeographicChartConfig(BaseChartConfig):
+ """Bounded controls shared by geographic visualizations."""
+
+ model_config = ConfigDict(extra="forbid", strict=True)
+ filters: list[GeographicFilterConfig] | None = Field(None, max_length=100)
+ row_limit: int = Field(10000, ge=1, le=10000)
+ time_range: str | None = Field(None, max_length=1000)
+
+ @field_validator("time_range")
+ @classmethod
+ def validate_geographic_time_range(cls, value: str | None) -> str | None:
+ """Reject time expressions the query parser would silently ignore."""
+ return validate_time_range(value)
+
+ @model_validator(mode="after")
+ def validate_geographic_roles(self) -> "GeographicChartConfig":
+ """Reject metric dimensions and unaggregated metric roles."""
+ for field in ("entity", "latitude", "longitude", "dimension"):
+ ref = getattr(self, field, None)
+ if ref is not None and (not ref.name or ref.is_metric):
+ raise ValueError(
+ f"{field} requires a named dataset column, not a metric"
+ )
+ for field in ("metric", "secondary_metric", "radius_metric"):
+ ref = getattr(self, field, None)
+ if ref is not None and not ref.is_metric:
+ raise ValueError(
+ f"{field} requires aggregate, saved_metric, or
sql_expression"
+ )
+ # These helpers import chart schemas, so defer to avoid a cycle.
+ from superset.mcp_service.chart.chart_utils import create_metric_object
+ from superset.mcp_service.chart.query_result import metric_result_label
+
+ dimensions = {
+ ref.name
+ for field in ("entity", "latitude", "longitude", "dimension")
+ if (ref := getattr(self, field, None)) is not None
+ }
+ seen_metrics: dict[str | None, ColumnRef] = {}
+ for field in ("metric", "secondary_metric", "radius_metric"):
+ ref = getattr(self, field, None)
+ if ref is None:
+ continue
+ label = metric_result_label(create_metric_object(ref))
+ if self.chart_type == "deck_scatter" and label in {
+ "position",
+ "weight",
+ "extraProps",
+ }:
+ raise ValueError(
+ f"Metric alias {label!r} conflicts with a native spatial
field"
+ )
+ if label in dimensions:
+ raise ValueError(
+ f"Metric alias {label!r} conflicts with a geographic
column"
+ )
+ if label in seen_metrics and seen_metrics[label] != ref:
+ raise ValueError(f"Distinct geographic metrics share alias
{label!r}")
+ seen_metrics[label] = ref
+ return self
+
+
+class CountryMapChartConfig(GeographicChartConfig):
+ """Regional choropleth joined against the bundled country boundaries."""
+
+ chart_type: Literal["country_map"]
+ country: Literal["usa", "canada", "australia", "japan", "uk"] = Field(
+ ..., description="Bundled boundary set. Only these countries are
supported."
+ )
+ region_format: Literal["name", "abbreviation", "iso_3166_2"] = Field(
+ ...,
+ description="Explicit source format; abbreviation means the ISO
suffix. "
+ "Names match bundled boundary names, not geocoding or fuzzy matching.",
+ )
+ entity: GeographicColumnRef
+ metric: GeographicColumnRef
+ linear_color_scheme: str = Field("schemeBlues", min_length=1,
max_length=100)
+ number_format: str = Field("SMART_NUMBER", min_length=1, max_length=100)
+
+
+class WorldMapChartConfig(GeographicChartConfig):
+ """Country choropleth with optional metric-sized bubbles."""
+
+ chart_type: Literal["world_map"]
+ entity: GeographicColumnRef
+ country_format: Literal["name", "cca2", "cca3", "cioc"]
+ metric: GeographicColumnRef
+ secondary_metric: GeographicColumnRef | None = None
+ show_bubbles: bool = False
+ max_bubble_size: int = Field(25, ge=1, le=100)
+ sort_by_metric: bool = True
+ linear_color_scheme: str = Field("schemeBlues", min_length=1,
max_length=100)
+
+ @model_validator(mode="after")
+ def validate_bubble_metric(self) -> "WorldMapChartConfig":
+ """Require an explicit size metric when bubbles are requested."""
+ if self.show_bubbles and self.secondary_metric is None:
+ raise ValueError(
+ "show_bubbles requires secondary_metric (may equal metric)"
+ )
+ return self
+
+
+class DeckScatterChartConfig(GeographicChartConfig):
+ """Geographic points from numeric longitude and latitude columns."""
+
+ chart_type: Literal["deck_scatter"]
+ latitude: GeographicColumnRef
+ longitude: GeographicColumnRef
+ dimension: GeographicColumnRef | None = None
+ radius_metric: GeographicColumnRef | None = None
+ radius: int = Field(1000, ge=1, le=1000000)
+ point_unit: Literal[
+ "square_m",
+ "square_km",
+ "square_miles",
+ "radius_m",
+ "radius_km",
+ "radius_miles",
+ ] = "radius_m"
+
+ @model_validator(mode="after")
+ def validate_coordinates(self) -> "DeckScatterChartConfig":
+ """Require distinct coordinates and protect computed spatial fields."""
+ if self.latitude.name == self.longitude.name:
+ raise ValueError("latitude and longitude must reference different
columns")
+ if self.dimension is not None and self.dimension.name in {
Review Comment:
Confirmed and fixed in 8578ad486ed416128a99ec9a55213c33efc42343.
`DeckScatterChartConfig` now rejects a `dimension` that names the latitude or
longitude column ("reuses a coordinate column; choose a separate category
column"), alongside the existing position/weight/extraProps check. Covered by
`test_point_dimension_cannot_reuse_a_coordinate_column[latitude|longitude]`,
which failed before the change.
##########
superset-frontend/plugins/preset-chart-deckgl/src/layers/Scatter/transformProps.ts:
##########
@@ -106,16 +170,32 @@ export default function transformProps(chartProps:
ChartProps) {
const { spatial, point_radius_fixed, dimension } =
formData as DeckScatterFormData;
- // Check if this is a fixed value or metric
- const fixedRadiusValue = isFixedValue(point_radius_fixed)
- ? getFixedValue(point_radius_fixed)
- : null;
+ // Typed compatibility accepts preserved numeric-string fixed radii.
+ // Native charts interpret bare strings as saved metric names.
+ const legacyFixedRadius =
Review Comment:
Confirmed and fixed in 8578ad486ed416128a99ec9a55213c33efc42343, both ways:
- Backend: typed updates (`merge_update_form_data` when `mcp_geographic` is
set) now save a preserved numeric-string radius as `{"type": "fix", "value":
100}` (`2.5` stays a float). Non-numeric strings like `count` are kept as
saved-metric keys. Covered by
`test_update_normalizes_preserved_numeric_string_radius`, which failed before.
- Frontend: a shared `getTypedFixedRadius` helper is used by both Scatter
`buildQuery` and `transformProps`, so a typed chart with `"100"` emits no
metric/orderby and renders a fixed radius. Without `mcp_geographic`,
`buildQuery` still emits `metrics: ["100"]` as before. New jest cases in
`buildQuery.test.ts` cover both; the typed ones failed before.
About a native saved metric literally named `100`: the MCP backend already
treats numeric strings as non-metrics in `_is_metric_ref` on master
(`chart_helpers.py`), so MCP could not query such a metric before this PR
either. The merge has no dataset metadata to tell the two apart, so this change
follows that existing rule and makes query and render agree.
##########
superset/mcp_service/chart/plugins/geographic.py:
##########
@@ -0,0 +1,717 @@
+# 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.
+
+"""Native geographic chart mapping shared by the three public map types."""
+
+from __future__ import annotations
+
+import logging
+import math
+import re
+from collections.abc import Mapping
+from decimal import Decimal
+from functools import lru_cache
+from numbers import Real
+from typing import Any, ClassVar, TypeGuard
+
+from superset.mcp_service.chart.chart_utils import (
+ _add_adhoc_filters,
+ create_metric_object,
+ merge_geographic_update_form_data,
+)
+from superset.mcp_service.chart.plugin import BaseChartPlugin
+from superset.mcp_service.chart.query_result import (
+ column_result_label,
+ metric_result_label,
+ query_result_failure,
+)
+from superset.mcp_service.chart.schemas import (
+ ChartError,
+ ColumnRef,
+ CountryMapChartConfig,
+ DeckScatterChartConfig,
+ VegaLitePreview,
+ WorldMapChartConfig,
+)
+from superset.mcp_service.chart.validation.dataset_validator import
DatasetValidator
+
+logger = logging.getLogger(__name__)
+
+GeographicConfig = CountryMapChartConfig | WorldMapChartConfig |
DeckScatterChartConfig
+ROLE_FIELDS = (
+ "entity",
+ "metric",
+ "secondary_metric",
+ "latitude",
+ "longitude",
+ "dimension",
+ "radius_metric",
+)
+MAX_GEOGRAPHIC_ROWS = 10000
+ASCII_BANNER = "Geographic source data (geometry not reproduced)"
+
+
+def _is_finite_geographic_number(value: object) -> TypeGuard[Real | Decimal]:
+ """Accept database NUMERIC/real scalars that remain finite in JSON.
+
+ Validation precedes JSON conversion; retain the original Decimal values for
+ data/export while rejecting booleans, complex numbers, and numeric strings.
+ """
+ if isinstance(value, bool) or not isinstance(value, (Real, Decimal)):
+ return False
+ if isinstance(value, Decimal) and not value.is_finite():
+ return False
+ try:
+ return math.isfinite(value)
+ except (OverflowError, ValueError):
+ return False
+
+
+def _decode_geographic_coordinates(value: object, spatial_type: str) ->
list[float]:
+ """Decode native encoded coordinates for validation, not result
replacement."""
+ if not isinstance(value, str) or not value:
+ raise ValueError("Spatial coordinates require a nonempty string")
+ if spatial_type == "geohash":
+ import pygeohash
+
+ try:
+ latitude, longitude = pygeohash.decode(value.lower())
+ except (ValueError, KeyError, TypeError) as ex:
+ raise ValueError("Invalid geographic geohash") from ex
+ return [longitude, latitude]
+ # Match the numeric prefix consumed by JavaScript parseFloat in the native
+ # Deck.gl spatial transform, including incomplete exponents and suffixes.
+ coordinates: list[float] = []
+ for part in value.split(","):
+ prefix = re.match(
+ r"[+-]?(?:Infinity|(?:[0-9]+(?:\.[0-9]*)?|\.[0-9]+)"
+ r"(?:[eE][+-]?[0-9]+)?)",
+ part.lstrip(
+ "\t\n\v\f\r \u00a0\u1680\u2000\u2001\u2002\u2003"
+ "\u2004\u2005\u2006\u2007\u2008\u2009\u200a\u2028"
+ "\u2029\u202f\u205f\u3000\ufeff"
+ ),
+ )
+ if prefix is None:
+ raise ValueError("Invalid delimited geographic coordinates")
+ coordinates.append(float(prefix[0]))
+ if len(coordinates) != 2:
+ raise ValueError("Delimited coordinates require two values")
+ return coordinates
+
+
+@lru_cache(maxsize=4)
+def _world_country_entries(field: str) -> tuple[tuple[str, str], ...]:
+ """Reuse immutable country aliases for the four supported world formats."""
+ from superset.examples.countries import countries
+
+ return tuple(
+ (country[field], country["cca3"]) for country in countries if
country[field]
+ )
+
+
+def _metric_expression(metric: object) -> tuple[object, ...] | None:
+ """Identify native metric expressions without UI labels or column
metadata."""
+ if isinstance(metric, str):
+ return ("saved", metric)
+ if not isinstance(metric, Mapping):
+ return None
+ expression_type = metric.get("expressionType")
+ if expression_type == "SQL" and isinstance(metric.get("sqlExpression"),
str):
+ return (expression_type, metric["sqlExpression"])
+ column = metric.get("column")
+ if expression_type == "SIMPLE" and isinstance(column, Mapping):
+ name = column.get("column_name") or column.get("columnName")
+ if isinstance(name, str) and isinstance(metric.get("aggregate"), str):
+ return (expression_type, metric["aggregate"], name)
+ return None
+
+
+def _bind_time_column(qd: dict[str, Any], form_data: Mapping[str, Any]) ->
None:
+ """Bind the legacy SQL time column so a saved ``time_range`` stays applied.
+
+ Native map viz classes filter ``time_range`` on ``granularity_sqla``; a
+ normalized dashboard ``granularity`` override takes precedence.
+ """
+ granularity = form_data.get("granularity",
form_data.get("granularity_sqla"))
+ if granularity:
+ qd["granularity"] = granularity
+
+
+def _typed_row_limit(form_data: Mapping[str, Any]) -> int:
+ """Bound the typed map row limit, matching the frontend's full-map
query."""
+ try:
+ limit = int(form_data.get("row_limit") or MAX_GEOGRAPHIC_ROWS)
+ except (TypeError, ValueError, OverflowError):
+ return MAX_GEOGRAPHIC_ROWS
+ return min(MAX_GEOGRAPHIC_ROWS, max(1, limit))
+
+
+class GeographicChartPlugin(BaseChartPlugin):
+ """Translate explicit geographic roles into native plugin controls.
+
+ Typed MCP charts carry ``mcp_geographic`` in their saved form_data. Only
+ those charts get the full-map row limits and strict result contract;
+ charts built in Explore keep their native behavior.
+ """
+
+ native_viz_types: ClassVar[Mapping[str, str]] = {}
+ requires_compile_check = True
+ requires_config_for_dataset_rebind = True
+ dataset_rebind_roles = "geographic/metric roles"
+ strict_dataset_rebind = True
+ normalize_data_results = True
+ supports_vega_lite_preview = False
+ invalid_result_error_code = "INVALID_GEOGRAPHIC_RESULT"
+ invalid_result_message = "Geographic query returned invalid values"
+ invalid_result_suggestions: ClassVar[tuple[str, ...]] = (
+ "Match country and value format to the source identifiers",
+ "Correct source values or filter other geographies",
+ "Use finite numeric metrics and valid latitude/longitude",
+ )
+
+ # ------------------------------------------------------------------
+ # Result contract
+ # ------------------------------------------------------------------
+
+ def result_metrics(self, form_data: Mapping[str, Any]) -> list[Any]:
+ """Return the metrics whose values every result row must carry."""
+ raise NotImplementedError
+
+ def size_metric_labels(
+ self, form_data: Mapping[str, Any], labels: list[str]
+ ) -> set[str]:
+ """Return the metric labels that size marks and must be nonnegative."""
+ return set()
+
+ def row_identifier(
+ self, row: Mapping[str, Any], form_data: Mapping[str, Any]
+ ) -> str | None:
+ """Resolve the row's geography, or validate it and return None."""
+ raise NotImplementedError
+
+ def build_query_dicts(
+ self,
+ form_data: dict[str, Any],
+ *,
+ viz_type: str,
+ engine: str,
+ row_limit: int | None,
+ order_desc: bool | None,
+ ) -> list[dict[str, Any]] | None:
+ """Build the shared single query, keeping the saved time column
bound."""
+ from superset.mcp_service.chart.chart_helpers import
build_single_query_dict
+
+ fields = self.resolve_query_fields(form_data, viz_type)
+ if fields is None:
+ return None
+ metrics, groupby = fields
+ qd = build_single_query_dict(
+ form_data, groupby, metrics, row_limit=row_limit,
order_desc=order_desc
+ )
+ _bind_time_column(qd, form_data)
+ return [qd]
+
+ def _metric_labels(self, form_data: Mapping[str, Any]) -> list[str]:
+ labels = [metric_result_label(m) for m in
self.result_metrics(form_data)]
+ if any(label is None for label in labels):
+ raise ValueError("Geographic metric has no resolvable result
label")
+ return [label for label in labels if label is not None]
+
+ def _validate_rows(self, result: Any, form_data: Mapping[str, Any]) ->
None:
+ if (
+ not isinstance(result, Mapping)
+ or not isinstance(result.get("queries"), list)
+ or len(result["queries"]) != 1
+ ):
+ raise ValueError("Expected exactly one geographic query result")
+ query = result["queries"][0]
+ if not isinstance(query, Mapping) or not isinstance(query.get("data"),
list):
+ raise ValueError("Expected geographic query data to be a list of
records")
+ labels = self._metric_labels(form_data)
+ size_labels = self.size_metric_labels(form_data, labels)
+ seen: set[str] = set()
+ for row in query["data"]:
+ if not isinstance(row, Mapping):
+ raise ValueError("Expected geographic rows to be records")
+ for label in labels:
+ value = row.get(label)
+ if not _is_finite_geographic_number(value):
+ raise ValueError(
+ f"Geographic metric {label!r} must be a finite number"
+ )
+ if value < 0 and label in size_labels:
+ raise ValueError("Geographic size metrics must be
nonnegative")
+ if identifier := self.row_identifier(row, form_data):
+ if identifier in seen:
+ raise ValueError(
+ f"Multiple result rows resolve to {identifier}; "
+ "normalize source values before aggregation"
+ )
+ seen.add(identifier)
+
+ def normalize_query_result(self, result: Any, form_data: Mapping[str,
Any]) -> Any:
+ """Reject unresolved regions and malformed or nonfinite results.
+
+ Source values are preserved for exports and filtering; the native
+ transform owns display-only ISO mapping using the same bundled
+ boundary identifiers. Charts without the typed MCP marker keep their
+ native behavior.
+ """
+ if failure := query_result_failure(result):
+ return failure
+ if not form_data.get("mcp_geographic"):
+ return result
+ try:
+ self._validate_rows(result, form_data)
+ except (ValueError, TypeError, KeyError) as exc:
+ return ChartError(error=str(exc),
error_type="InvalidGeographicResult")
+ return result
+
+ # ------------------------------------------------------------------
+ # Query limits and previews
+ # ------------------------------------------------------------------
+
+ def compile_row_limit(self, form_data: Mapping[str, Any]) -> int:
+ """Validate the full bounded map, not only its first rows."""
+ if form_data.get("mcp_geographic"):
+ return _typed_row_limit(form_data)
+ return super().compile_row_limit(form_data)
+
+ def preview_row_limit(self, form_data: Mapping[str, Any], fallback: int)
-> int:
+ """Preview the same bounded rows the native map renders."""
+ if form_data.get("mcp_geographic"):
+ return _typed_row_limit(form_data)
+ return fallback
+
+ def ascii_preview(
+ self, data: list[Any], form_data: dict[str, Any], width: int
+ ) -> str | ChartError | None:
+ """Render the source rows; map geometry has no ASCII form."""
+ from superset.mcp_service.chart.ascii_charts import
generate_ascii_table
+
+ try:
+ return f"{ASCII_BANNER}\n" + generate_ascii_table(data, max(width,
21))
+ except (TypeError, ValueError, KeyError, IndexError) as exc:
+ logger.error("ASCII chart generation failed: %s", exc,
exc_info=True)
+ return "ASCII chart generation failed"
+
+ def vega_lite_preview(
+ self, data: list[Any], form_data: dict[str, Any]
+ ) -> VegaLitePreview | ChartError | None:
+ """Never fabricate a non-geographic Vega-Lite chart for a map."""
+ return ChartError(
+ error=(
+ "Geographic Vega previews are not supported. Use table/ascii
for "
+ "source data, or open Explore for native geography."
+ ),
+ error_type="UnsupportedGeographicPreview",
+ )
+
+ # ------------------------------------------------------------------
+ # Updates
+ # ------------------------------------------------------------------
+
+ def merge_update_form_data(
+ self,
+ existing_form_data: dict[str, Any],
+ new_form_data: dict[str, Any],
+ config: Any,
+ *,
+ dataset_rebind: bool,
+ ) -> dict[str, Any] | None:
+ """Preserve omitted native controls; a rebind keeps presentation
only."""
+ merged = merge_geographic_update_form_data(
+ existing_form_data, new_form_data, config,
dataset_rebind=dataset_rebind
+ )
+ # Native query and compile validation must cover the same bounded rows.
+ if merged.get("mcp_geographic"):
+ merged["row_limit"] = _typed_row_limit(merged)
+ return merged
+
+ def extract_column_refs(self, config: GeographicConfig) -> list[ColumnRef]:
+ """Include all spatial, metric, and filter references."""
+ refs = [
+ ref
+ for field in ROLE_FIELDS
+ if (ref := getattr(config, field, None)) is not None
+ ]
+ refs.extend(ColumnRef(name=f.column) for f in config.filters or [])
+ return refs
+
+ def normalize_column_refs(
+ self, config: GeographicConfig, dataset_context: Any
+ ) -> GeographicConfig:
+ """Canonicalize dataset identifiers without losing explicit field
sets."""
+ patch = config.model_dump(exclude_unset=True)
+ for field in ROLE_FIELDS:
+ ref = patch.get(field)
+ if ref and ref.get("name"):
+ canonical = (
+ DatasetValidator.get_canonical_metric_name
+ if ref.get("saved_metric")
+ else DatasetValidator.get_canonical_column_name
+ )
+ ref["name"] = canonical(ref["name"], dataset_context)
+ DatasetValidator.normalize_filters(patch, dataset_context)
+ return type(config).model_validate(patch)
+
+ def resolve_viz_type(self, config: GeographicConfig) -> str:
+ """Public names match frontend registration keys."""
+ return config.chart_type
+
+ def generate_name(
+ self, config: GeographicConfig, dataset_name: str | None = None
+ ) -> str:
+ """Name geographic charts without guessing the source geography."""
+ return self._with_context(self.display_name, dataset_name)
+
+ def to_form_data(
+ self, config: GeographicConfig, dataset_id: int | str | None = None
+ ) -> dict[str, Any]:
+ """Mirror native country/world/deck Scatter controls and defaults."""
+ result: dict[str, Any] = {
+ "viz_type": config.chart_type,
+ "row_limit": config.row_limit,
+ "mcp_geographic": True,
+ }
+ if config.time_range is not None:
+ result["time_range"] = config.time_range
+ _add_adhoc_filters(result, config.filters)
+ if isinstance(config, CountryMapChartConfig):
+ result.update(
+ entity=config.entity.name,
+ metric=create_metric_object(config.metric),
+ select_country=config.country,
+ region_format=config.region_format,
+ linear_color_scheme=config.linear_color_scheme,
+ number_format=config.number_format,
+ )
+ elif isinstance(config, WorldMapChartConfig):
+ result.update(
+ entity=config.entity.name,
+ metric=create_metric_object(config.metric),
+ country_fieldtype=config.country_format,
+ secondary_metric=create_metric_object(config.secondary_metric)
+ if config.secondary_metric
+ else None,
+ show_bubbles=config.show_bubbles,
+ max_bubble_size=config.max_bubble_size,
+ sort_by_metric=config.sort_by_metric,
+ linear_color_scheme=config.linear_color_scheme,
+ color_by="metric",
+ color_picker={"r": 0, "g": 122, "b": 135, "a": 1},
+ color_scheme="supersetColors",
+ y_axis_format="SMART_NUMBER",
+ )
+ else:
+ radius: dict[str, Any] = {"type": "fix", "value": config.radius}
+ if config.radius_metric:
+ radius = {
+ "type": "metric",
+ "value": create_metric_object(config.radius_metric),
+ }
+ result.update(
+ spatial={
+ "type": "latlong",
+ "latCol": config.latitude.name,
+ "lonCol": config.longitude.name,
+ },
+ dimension=config.dimension.name if config.dimension else None,
+ point_radius_fixed=radius,
+ point_unit=config.point_unit,
+ multiplier=1,
+ min_radius=2,
+ max_radius=250,
+ color_picker={"r": 0, "g": 122, "b": 135, "a": 1},
+ color_scheme="supersetColors",
+ map_renderer="maplibre",
+
maplibre_style="https://basemaps.cartocdn.com/gl/positron-gl-style/style.json",
+ viewport={
+ "longitude": 0,
+ "latitude": 0,
+ "zoom": 1,
+ "bearing": 0,
+ "pitch": 0,
+ },
+ autozoom=True,
+ filter_nulls=True,
+ )
+ return result
+
+
+class CountryMapChartPlugin(GeographicChartPlugin):
+ """Register the Country Map visualization."""
+
+ chart_type = "country_map"
+ display_name = "Country Map"
+ native_viz_types: ClassVar[Mapping[str, str]] = {"country_map": "Country
Map"}
+
+ def resolve_query_fields(
+ self, form_data: Mapping[str, Any], viz_type: str
+ ) -> tuple[list[Any], list[Any]] | None:
+ """Query the region entity and its single metric."""
+ metric = form_data.get("metric")
+ entity = form_data.get("entity")
+ return [metric] if metric else [], [entity] if entity else []
+
+ def result_metrics(self, form_data: Mapping[str, Any]) -> list[Any]:
+ return [form_data.get("metric")]
+
+ def row_identifier(
+ self, row: Mapping[str, Any], form_data: Mapping[str, Any]
+ ) -> str | None:
+ """Resolve the row to one bundled region boundary."""
+ from superset.utils.geographic import resolve_region
+
+ entity = column_result_label(form_data.get("entity"))
+ if entity is None:
+ raise ValueError("Geographic maps require an entity column")
+ # Explore's legacy format uses full boundary ISO codes without
normalization.
+ return resolve_region(
+ row.get(entity),
+ form_data.get("select_country", ""),
+ form_data.get("region_format") or "iso_3166_2",
Review Comment:
Confirmed and fixed in 8578ad486ed416128a99ec9a55213c33efc42343.
`resolve_region` / `resolve_geographic_value` take an `exact` flag, and
`CountryMapChartPlugin.row_identifier` sets it when `region_format` is empty.
Legacy-format validation now accepts only the exact boundary id the renderer
joins on (`US-CA` passes; `us-ca`, `Us-Ca`, `US-ca` are rejected as
unrecognized). Covered by
`test_legacy_country_format_requires_exact_boundary_iso`, which failed before.
The existing legacy export test with `US-CA` still passes.
##########
tests/unit_tests/mcp_service/chart/test_geographic_chart.py:
##########
@@ -0,0 +1,1858 @@
+# 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.
+
+"""Typed geographic contracts, native query semantics, and boundary parity."""
+
+from copy import deepcopy
+from decimal import Decimal
+from pathlib import Path
+from typing import Any
+from unittest.mock import Mock, patch
+
+import pytest
+from fastmcp import Client
+from pydantic import TypeAdapter, ValidationError
+from sqlalchemy.exc import SQLAlchemyError
+
+from superset.mcp_service.app import mcp
+from superset.mcp_service.chart.chart_helpers import
build_query_dicts_from_form_data
+from superset.mcp_service.chart.chart_utils import (
+ analyze_chart_capabilities,
+ map_config_to_form_data,
+ merge_chart_form_data,
+)
+from superset.mcp_service.chart.compile import _compile_chart
+from superset.mcp_service.chart.plugins.geographic import (
+ CountryMapChartPlugin,
+ DeckScatterChartPlugin,
+)
+from superset.mcp_service.chart.preview_utils import (
+ _generate_ascii_preview_from_data,
+ _generate_vega_lite_preview_from_data,
+)
+from superset.mcp_service.chart.query_result import (
+ metric_result_label,
+ normalize_chart_query_result,
+)
+from superset.mcp_service.chart.registry import get_registry
+from superset.mcp_service.chart.schemas import (
+ ChartConfig,
+ ChartError,
+ GenerateChartRequest,
+ GenerateExploreLinkRequest,
+ UpdateChartRequest,
+)
+from superset.mcp_service.chart.tool.get_chart_type_schema import (
+ _CHART_EXAMPLES,
+ _get_chart_type_schema_impl,
+)
+from superset.utils import json
+from superset.utils.geographic import resolve_geographic_value, resolve_region
+from superset.utils.geographic_regions import REGIONS
+
+KINDS = ("country_map", "world_map", "deck_scatter")
+# Reuse the compiled union schema; each validation still creates a fresh
config.
+CHART_CONFIG_ADAPTER = TypeAdapter(ChartConfig)
+
+
+def config_for(kind: str) -> Any:
+ """Parse the published example rather than duplicating a private
contract."""
+ return CHART_CONFIG_ADAPTER.validate_python(_CHART_EXAMPLES[kind][0])
+
+
+def form_for(kind: str) -> dict[str, Any]:
+ """Map the same config consumed by the three public tools."""
+ return map_config_to_form_data(config_for(kind))
+
+
+def result_for(kind: str) -> dict[str, Any]:
+ """Native query results before frontend display transforms."""
+ row = (
+ {"state": "CA", "SUM(sales)": 10}
+ if kind == "country_map"
+ else {"country": "US", "SUM(sales)": 10}
+ if kind == "world_map"
+ else {"latitude": 37.8, "longitude": -122.4}
+ )
+ return {"queries": [{"data": [row]}]}
+
+
+def invalid_result_for(kind: str) -> dict[str, Any]:
+ """Keep valid metrics while failing the actual geographic value
contract."""
+ result = result_for(kind)
+ row = result["queries"][0]["data"][0]
+ if kind == "country_map":
+ row["state"] = "BC"
+ elif kind == "world_map":
+ row["country"] = "not-a-country"
+ else:
+ row["latitude"] = 91
+ return result
+
+
[email protected]("kind", KINDS)
+def test_geographic_example_configs_are_independent(kind: str) -> None:
+ """Sharing a compiled schema must not share mutable config instances."""
+ first = config_for(kind)
+ second = config_for(kind)
+ assert first is not second
+ first.row_limit = 1
+ assert second.row_limit == 10000
+ assert config_for(kind).row_limit == 10000
+
+
[email protected]("kind", KINDS)
+def test_geographic_schema_examples_and_all_request_unions(kind: str) -> None:
+ """Each entry point uses the required, bounded shared discriminator."""
+ example = _CHART_EXAMPLES[kind][0]
+ schema = _get_chart_type_schema_impl(kind)["schema"]
+ assert "chart_type" in schema["required"]
+ assert schema["additionalProperties"] is False
+ assert schema["properties"]["row_limit"]["maximum"] == 10000
+ for model, identity in (
+ (GenerateChartRequest, {"dataset_id": 3}),
+ (GenerateExploreLinkRequest, {"dataset_id": 3}),
+ (UpdateChartRequest, {"identifier": 1}),
+ ):
+ assert (
+ model.model_validate({**identity, "config":
example}).config.chart_type
+ == kind
+ )
+ for patch_ in (
+ {"row_limit": 10001},
+ {"row_limit": True},
+ {"row_limit": "100"},
+ {"bogus": 1},
+ ):
+ with pytest.raises(ValidationError):
+ CHART_CONFIG_ADAPTER.validate_python({**example, **patch_})
+ with pytest.raises(ValidationError):
+ CHART_CONFIG_ADAPTER.validate_python(
+ {k: v for k, v in example.items() if k != "chart_type"}
+ )
+
+
[email protected](
+ "country,value,format_,expected",
+ [
+ ("usa", "CA", "abbreviation", "US-CA"),
+ ("usa", "ca", "abbreviation", "US-CA"),
+ ("usa", "California", "name", "US-CA"),
+ ("usa", "us-ca", "iso_3166_2", "US-CA"),
+ ("canada", "BC", "abbreviation", "CA-BC"),
+ ("australia", "Victoria", "name", "AU-VIC"),
+ ("australia", "NSW", "abbreviation", "AU-NSW"),
+ ("australia", "Queensland", "name", "AU-QLD"),
+ ("japan", "Tokyo", "name", "JP-13"),
+ ("japan", "Osaka", "name", "JP-27"),
+ ("uk", "Isle of Wight", "name", "GB-IOW"),
+ ],
+)
+def test_region_resolution(
+ country: str, value: str, format_: str, expected: str
+) -> None:
+ """All formats resolve only to identifiers present in the chosen
geometry."""
+ assert resolve_region(value, country, format_) == expected
+
+
[email protected](
+ "value",
+ [
+ "BC",
+ "Victoria",
+ "NSW",
+ "Queensland",
+ "Tokyo",
+ "Osaka",
+ "Isle of Wight",
+ None,
+ 1,
+ "",
+ "CA ",
+ ],
+)
+def test_us_rejects_non_us_and_malformed_values(value: object) -> None:
+ """Cross-country values never silently disappear from a map."""
+ with pytest.raises(ValueError, match="country=usa"):
+ resolve_region(value, "usa", "abbreviation")
+
+
+def test_exact_first_and_ambiguous_folded_names() -> None:
+ """Do not let a case-insensitive dictionary overwrite distinct names."""
+ pairs = [("Region", "A"), ("REGION", "B")]
+ assert resolve_geographic_value("Region", pairs) == "A"
+ with pytest.raises(ValueError, match="ambiguous"):
+ resolve_geographic_value("region", pairs)
+ with pytest.raises(ValueError, match="ambiguous"):
+ resolve_geographic_value("Region", pairs + [("Region", "C")])
+
+
[email protected]("country", REGIONS)
+def test_region_data_matches_frontend_geometry(country: str) -> None:
+ """Updating geometry requires updating its bounded backend lookup too."""
+ root = Path(__file__).resolve().parents[4]
+ path = (
+ root
+ / "superset-frontend/plugins/plugin-chart-country-map/src/countries"
+ / f"{country}.geojson"
+ )
+ expected = sorted(
+ {
+ (
+ f["properties"]["ISO"],
+ f["properties"].get("NAME_2") or f["properties"]["NAME_1"],
+ )
+ for f in json.loads(path.read_text())["features"]
+ }
+ )
+ assert REGIONS[country] == expected
+
+
[email protected]("kind", KINDS)
+def test_geographic_native_query_and_filters(kind: str) -> None:
+ """Query roles and ordering match the frontend's buildQuery contract."""
+ config = CHART_CONFIG_ADAPTER.validate_python(
+ {
+ **_CHART_EXAMPLES[kind][0],
+ "filters": [{"column": "segment", "op": "IN", "value":
["Retail"]}],
+ }
+ )
+ form = map_config_to_form_data(config)
+ with patch(
+ "superset.mcp_service.chart.chart_helpers.resolve_datasource_engine",
+ return_value="sqlite",
+ ):
+ query = build_query_dicts_from_form_data(form, 3, "table")[0]
+ assert {"col": "segment", "op": "IN", "val": ["Retail"]} in
query["filters"]
+ assert query["row_limit"] == 10000
+ if kind == "deck_scatter":
+ assert set(query["columns"]) == {"latitude", "longitude"}
+ assert query["metrics"] == []
+ assert query["is_timeseries"] is False
+ assert query["orderby"] == []
+ assert {"col": "latitude", "op": "IS NOT NULL", "val": ""} in
query["filters"]
+ else:
+ assert query["columns"] == [form["entity"]]
+ assert query["metrics"] == [form["metric"]]
+ if kind == "world_map":
+ assert query["orderby"] == [(form["metric"], False)]
+
+
[email protected]("same", [True, False])
+def test_world_bubble_metrics_deduplicate_by_label(same: bool) -> None:
+ """The secondary bubble-size metric is queried unless its alias is
shared."""
+ config = CHART_CONFIG_ADAPTER.validate_python(
+ {
+ **_CHART_EXAMPLES["world_map"][0],
+ "show_bubbles": True,
+ "secondary_metric": {
+ "name": "sales" if same else "population",
+ "aggregate": "SUM",
+ },
+ }
+ )
+ form = map_config_to_form_data(config)
+ with patch(
+ "superset.mcp_service.chart.chart_helpers.resolve_datasource_engine",
+ return_value="sqlite",
+ ):
+ query = build_query_dicts_from_form_data(form, 3, "table")[0]
+ assert len(query["metrics"]) == (1 if same else 2)
+
+
+def _size_metric_case(kind: str) -> tuple[dict[str, Any], dict[str, Any], str]:
+ """Return form data, a valid result and the size metric's result label."""
+ if kind == "world_map":
+ config = CHART_CONFIG_ADAPTER.validate_python(
+ {
+ **_CHART_EXAMPLES["world_map"][0],
+ "show_bubbles": True,
+ "secondary_metric": {"name": "population", "aggregate": "SUM"},
+ }
+ )
+ else:
+ config = CHART_CONFIG_ADAPTER.validate_python(
+ {
+ **_CHART_EXAMPLES["deck_scatter"][0],
+ "radius_metric": {"name": "orders", "aggregate": "COUNT"},
+ }
+ )
+ form = map_config_to_form_data(config)
+ label = (
+ metric_result_label(form["secondary_metric"])
+ if kind == "world_map"
+ else metric_result_label(form["point_radius_fixed"]["value"])
+ )
+ assert label is not None
+ result = result_for(kind)
+ result["queries"][0]["data"][0][label] = 5
+ return form, result, label
+
+
[email protected]("kind", ["world_map", "deck_scatter"])
[email protected]("value", [-1, -0.5, Decimal("-0.001")])
+def test_negative_size_metric_is_rejected(kind: str, value: object) -> None:
+ """Bubble and point radius metrics size marks, so they must be
nonnegative."""
+ form, result, label = _size_metric_case(kind)
+ assert not isinstance(normalize_chart_query_result(result, form),
ChartError)
+ result["queries"][0]["data"][0][label] = value
+ failure = normalize_chart_query_result(result, form)
+ assert isinstance(failure, ChartError)
+ assert failure.error_type == "InvalidGeographicResult"
+ assert "size metrics must be nonnegative" in failure.error
+
+
[email protected]("kind", ["world_map", "deck_scatter"])
+def test_zero_size_metric_is_accepted(kind: str) -> None:
+ """Zero is a valid, if invisible, mark size."""
+ form, result, label = _size_metric_case(kind)
+ result["queries"][0]["data"][0][label] = 0
+ assert not isinstance(normalize_chart_query_result(result, form),
ChartError)
+
+
+def test_world_map_negative_color_metric_is_accepted_with_bubbles() -> None:
+ """Only the bubble metric is size-constrained; the color metric may be <
0."""
+ form, result, _ = _size_metric_case("world_map")
+ result["queries"][0]["data"][0][metric_result_label(form["metric"])] = -10
+ assert not isinstance(normalize_chart_query_result(result, form),
ChartError)
+
+
[email protected]("value", [-1, None, float("nan"), "n/a"])
+def test_world_map_unused_secondary_metric_does_not_gate_choropleth(
+ value: object,
+) -> None:
+ """A secondary metric kept with show_bubbles=False is not rendered, so its
+ values cannot fail the color choropleth."""
+ form, result, label = _size_metric_case("world_map")
+ form = {**form, "show_bubbles": False}
+ result["queries"][0]["data"][0][label] = value
+ assert not isinstance(normalize_chart_query_result(result, form),
ChartError)
+
+ form["show_bubbles"] = True
+ assert isinstance(normalize_chart_query_result(result, form), ChartError)
+
+
+def test_world_map_bubbles_without_secondary_metric_is_rejected() -> None:
+ """Bubbles still require the size metric at result validation."""
+ form = {**form_for("world_map"), "show_bubbles": True}
+ form.pop("secondary_metric", None)
+ failure = normalize_chart_query_result(result_for("world_map"), form)
+ assert isinstance(failure, ChartError)
+ assert "show_bubbles requires secondary_metric" in failure.error
+
+
[email protected]("kind", KINDS)
+def test_result_validation_preserves_source_and_cache(kind: str) -> None:
+ """Validation cannot mutate cache records or export identifiers."""
+ result = result_for(kind)
+ before = deepcopy(result)
+ assert normalize_chart_query_result(result, form_for(kind)) == before
+ assert result == before
+ assert normalize_chart_query_result(
+ {"queries": [{"data": []}]}, form_for(kind)
+ ) == {"queries": [{"data": []}]}
+ bad: dict[str, Any]
+ for bad in (
+ {},
+ {"queries": []},
+ {"queries": [{"data": [{}]}]},
+ {"queries": [{"data": "bad"}]},
+ ):
+ assert isinstance(normalize_chart_query_result(bad, form_for(kind)),
ChartError)
+
+
[email protected]("kind", KINDS)
+def test_invalid_values_and_geometry_preview(kind: str) -> None:
+ """Never return fabricated Vega bars for geographic requests."""
+ form = form_for(kind)
+ result = result_for(kind)
+ row = result["queries"][0]["data"][0]
+ for column in list(row):
+ bad = deepcopy(result)
+ bad["queries"][0]["data"][0][column] = "not valid"
+ error = normalize_chart_query_result(bad, form)
+ assert isinstance(error, ChartError)
+ assert error.error_type == "InvalidGeographicResult"
+ assert (
+ "geometry not reproduced"
+ in _generate_ascii_preview_from_data([row], form).ascii_content
+ )
+ preview = _generate_vega_lite_preview_from_data([row], form)
+ assert isinstance(preview, ChartError)
+ assert preview.error_type == "UnsupportedGeographicPreview"
+
+
[email protected]("kind", KINDS)
+def test_advertised_preview_formats_can_actually_be_produced(kind: str) ->
None:
+ """Capabilities must not offer a Vega-Lite preview the generator
rejects."""
+ form = form_for(kind)
+ row = result_for(kind)["queries"][0]["data"][0]
+ capabilities = analyze_chart_capabilities(form["viz_type"],
config_for(kind))
+
+ assert "vega_lite" not in capabilities.optimal_formats
+ assert isinstance(_generate_vega_lite_preview_from_data([row], form),
ChartError)
+
+ # A non-geographic interactive type still advertises what it can produce.
+ scatter_form = {"viz_type": "echarts_timeseries_scatter", "x_axis": "x"}
+ scatter_capabilities = analyze_chart_capabilities(
+ scatter_form["viz_type"], config_for(kind)
+ )
+ assert "vega_lite" in scatter_capabilities.optimal_formats
+ assert not isinstance(
+ _generate_vega_lite_preview_from_data([{"x": "a", "y": 1}],
scatter_form),
+ ChartError,
+ )
+
+
+def test_alias_collisions_fail_instead_of_losing_aggregates() -> None:
+ """CA and ca grouped separately cannot safely be added (e.g. AVG)."""
+ result = {
+ "queries": [
+ {"data": [{"state": value, "SUM(sales)": 1} for value in ("CA",
"ca")]}
+ ]
+ }
+ error = normalize_chart_query_result(result, form_for("country_map"))
+ assert isinstance(error, ChartError)
+ assert "normalize source values before aggregation" in error.error
+
+
[email protected]("kind", KINDS)
+def test_update_omissions_explicit_clearing_and_rebind(kind: str) -> None:
+ """Unspecified controls survive updates, not old dataset roles on
rebind."""
+ config = config_for(kind)
+ old = {
+ **form_for(kind),
+ "row_limit": 12,
+ "time_range": "Last week",
+ "template_params": {"old": 1},
+ "adhoc_filters": [
+ {
+ "subject": "segment",
+ "operator": "IN",
+ "comparator": ["Retail"],
+ "expressionType": "SIMPLE",
+ "clause": "WHERE",
+ }
+ ],
+ }
+ merged = merge_chart_form_data(old, form_for(kind), config)
+ assert merged["row_limit"] == 12
+ assert merged["time_range"] == "Last week"
+ assert merged["adhoc_filters"] == old["adhoc_filters"]
+ cleared = CHART_CONFIG_ADAPTER.validate_python(
+ {**_CHART_EXAMPLES[kind][0], "filters": [], "time_range": None}
+ )
+ merged = merge_chart_form_data(old, map_config_to_form_data(cleared),
cleared)
+ assert not merged.get("adhoc_filters")
+ assert merged.get("time_range") is None
+ rebound = merge_chart_form_data(old, form_for(kind), config,
dataset_rebind=True)
+ assert not rebound.get("adhoc_filters")
+ assert "template_params" not in rebound
+
+
[email protected]("kind", KINDS)
+def test_compile_checks_full_bounded_map_result(kind: str) -> None:
+ """An invalid region outside the first two rows must block generation."""
+ result = result_for(kind)
+ with (
+ patch(
+
"superset.mcp_service.chart.chart_helpers.build_query_context_from_form_data",
+ return_value=Mock(),
+ ) as build,
+ patch(
+ "superset.commands.chart.data.get_data_command.ChartDataCommand"
+ ) as command,
+ ):
+ command.return_value.run.return_value = result
+ assert _compile_chart(form_for(kind), 3).success
+ assert build.call_args.kwargs["row_limit"] == 10000
Review Comment:
Good catch, fixed in 8578ad486ed416128a99ec9a55213c33efc42343. The test now
adds a second, distinct valid row (`TX` / `FR` / a second coordinate pair,
depending on kind), checks that the two-row result compiles, then appends the
invalid third row and checks for `INVALID_GEOGRAPHIC_RESULT`. A regression that
only validates `data[:2]` would now fail.
--
This is an automated message from the Apache Git Service.
To respond to the message, please log on to GitHub and use the
URL above to go to the specific comment.
To unsubscribe, e-mail: [email protected]
For queries about this service, please contact Infrastructure at:
[email protected]
---------------------------------------------------------------------
To unsubscribe, e-mail: [email protected]
For additional commands, e-mail: [email protected]