aminghadersohi commented on code in PR #43771:
URL: https://github.com/apache/superset/pull/43771#discussion_r4190030428
##########
superset/mcp_service/chart/query_result.py:
##########
@@ -15,67 +15,1416 @@
# specific language governing permissions and limitations
# under the License.
-"""Helpers for interpreting ChartDataCommand result envelopes."""
+"""Canonicalize and validate ``ChartDataCommand`` result envelopes."""
import math
+import time as system_time
+from bisect import bisect_right
from collections.abc import Mapping
+from dataclasses import dataclass
+from datetime import date, datetime, time, timedelta, timezone
from decimal import Decimal
+from enum import Enum
from numbers import Real
from typing import Any, cast
+from uuid import UUID
+from zoneinfo import ZoneInfo
+import numpy as np
+import pandas as pd
+import pytz
+from dateutil import tz as dateutil_tz, zoneinfo as dateutil_zoneinfo
+from pydantic import BaseModel
+from pydantic_core import to_json
+
+from superset.common.chart_data import ChartDataResultFormat
+from superset.common.db_query_status import QueryStatus
+from superset.constants import CACHE_DISABLED_TIMEOUT
from superset.mcp_service.chart.schemas import ChartError
+from superset.utils.core import (
+ ExtraFiltersReasonType,
+ ExtraFiltersTimeColumnType,
+ GenericDataType,
+)
FAILED_QUERY_STATUSES = frozenset(
{"error", "failed", "stopped", "timed_out", "cancelled", "canceled"}
)
+# These are aggregate envelope limits, not per-query allowances. In particular,
+# splitting a result across the maximum number of queries must not multiply the
+# permitted rows, nodes, or encoded bytes.
+MAX_QUERY_RESULTS = 32
+MAX_QUERY_RESULT_ROWS_PER_QUERY = 50_000
+MAX_QUERY_RESULT_ROWS = 100_000
+MAX_QUERY_RESULT_COLUMNS = 4_096
+MAX_QUERY_RESULT_VALUES = 2_500_000
+MAX_QUERY_RESULT_VALUE_BYTES = 16 * 1024 * 1024
+MAX_QUERY_RESULT_METADATA_BYTES = 1024 * 1024
+MAX_QUERY_RESULT_METADATA_ITEMS = 32_768
+MAX_RESULT_VALUE_ITEMS = 4_096
+MAX_RESULT_VALUE_DEPTH = 32
+MAX_RESULT_STRING_LENGTH = 65_536
+MAX_RESULT_KEY_LENGTH = 4_096
+MAX_RESULT_INTEGER_BITS = 4_096
+MAX_RESULT_INTEGER_DIGITS = 1_234
+MAX_RESULT_DECIMAL_DIGITS = 1_024
+MAX_RESULT_DECIMAL_MAGNITUDE = 4_096
+MAX_RESULT_DECIMAL_STORAGE = 2_048
+MAX_QUERY_RESULT_ROWCOUNT = 2**63 - 1
+MAX_QUERY_RESULT_CACHE_TIMEOUT = 2**31 - 1
+MAX_QUERY_RESULT_TIMESTAMP_LENGTH = 64
+
+_ERROR_KEYS = ("error", "errors", "error_message", "message", "detail")
+_MAX_ERROR_TEXT_BYTES = 2_000
+_TRUSTED_TIMEZONE_TYPES = (timezone, ZoneInfo)
+_SAFE_RESULT_ENUM_TYPES = frozenset(
+ {
+ ChartDataResultFormat,
+ QueryStatus,
+ ExtraFiltersReasonType,
+ ExtraFiltersTimeColumnType,
+ GenericDataType,
+ }
+)
+_RESULT_FORMAT_VALUES = frozenset(
+ object.__getattribute__(member, "_value_") for member in
ChartDataResultFormat
+)
+_COLTYPE_VALUES = frozenset(
+ object.__getattribute__(member, "_value_") for member in GenericDataType
+)
+_NUMPY_INTEGER_TYPES = frozenset(
+ type(value)
+ for value in (
+ np.int8(0),
+ np.int16(0),
+ np.int32(0),
+ np.int64(0),
+ np.uint8(0),
+ np.uint16(0),
+ np.uint32(0),
+ np.uint64(0),
+ )
+)
+_NUMPY_FLOAT_TYPES = frozenset(
+ type(value)
+ for value in (np.float16(0), np.float32(0), np.float64(0),
np.longdouble(0))
+)
+_PANDAS_NAT_TYPE = type(pd.NaT)
+_PANDAS_NA_TYPE = type(pd.NA)
+_PANDAS_PERIOD_TYPE = type(pd.Period("2000-01", freq="M"))
+_PANDAS_INTERVAL_TYPE = type(pd.Interval(0, 1))
+_DATEUTIL_FIXED_TIMEZONE_TYPES = frozenset(
+ {type(dateutil_tz.tzoffset(None, 0)), type(dateutil_tz.tzutc())}
+)
+_DATEUTIL_NAMED_TIMEZONE_TYPES = frozenset(
+ {dateutil_tz.tzfile, dateutil_zoneinfo.tzfile}
+)
+_DATEUTIL_TTINFO_TYPE = type(
+ object.__getattribute__(dateutil_tz.gettz("UTC"),
"__dict__")["_ttinfo_std"]
+)
+_DATEUTIL_LOCAL_TIMEZONE_TYPE = type(dateutil_tz.tzlocal())
+_PYTZ_FIXED_TIMEZONE_TYPES = frozenset({type(pytz.FixedOffset(1))})
+_MAX_DATEUTIL_TRANSITIONS = 4_096
+_MAX_DATEUTIL_TTINFOS = 256
+_MAX_DATEUTIL_TRANSITION_MAGNITUDE = 10**12
+
+
+@dataclass
+class _ResultBudget:
+ """Aggregate counters shared by every query and metadata value."""
+
+ rows: int = 0
+ values: int = 0
+ json_bytes: int = 0
+ metadata_items: int = 0
+ metadata_bytes: int = 0
+
+
+@dataclass(frozen=True)
+class _DateutilTimezoneState:
+ """Hook-free subset of a validated exact dateutil tzfile transition
table."""
+
+ transitions: tuple[int, ...]
+ transition_offsets: tuple[int, ...]
+ standard_offset: int
+ before_offset: int | None
-def _query_error_text(value: Any) -> str | None:
- """Convert a bounded query error payload into a useful message."""
- if value is None or value is False:
+
+def _invalid_result(message: str) -> ChartError:
+ return ChartError(
+ error=f"Chart query returned {message}.",
+ error_type="InvalidQueryResult",
+ )
+
+
+def _invalid_metadata(label: str) -> ChartError:
+ return ChartError(
+ error=f"{label} returned hostile or malformed metadata.",
+ error_type="InvalidQueryResult",
+ )
+
+
+def _safe_enum_value(value: Any, expected: frozenset[type[Any]]) -> Any | None:
+ """Read trusted enum storage without invoking public conversion hooks."""
+ if type(value) not in expected or type(value) not in
_SAFE_RESULT_ENUM_TYPES:
return None
- if isinstance(value, Mapping):
- for key in ("error", "error_message", "message", "detail"):
- if text := _query_error_text(value.get(key)):
- return text
+ return object.__getattribute__(value, "_value_")
+
+
+def _bounded_utf8_length(value: str, maximum: int) -> int | None:
+ """Return the exact UTF-8 size while bounding pre-encoding work."""
+ if str.__len__(value) > maximum:
return None
- if isinstance(value, (list, tuple)):
- parts = [text for item in value if (text := _query_error_text(item))]
- return "; ".join(parts[:3]) or None
- text = str(value)
- return text[:2000] if text else None
+ try:
+ encoded = str.encode(value, "utf-8", errors="strict")
+ except UnicodeEncodeError:
+ return None
+ size = bytes.__len__(encoded)
+ return size if size <= maximum else None
-def _failure_for_query_payload(
- payload: Mapping[str, Any], label: str
-) -> ChartError | None:
- """Extract one failure from a top-level or per-query payload."""
+def _json_string_size(value: str, maximum: int) -> int | None:
+ """Return compact UTF-8 JSON string size without serializing the value."""
+ raw_size = _bounded_utf8_length(value, maximum)
+ if raw_size is None:
+ return None
+ escaped_size = raw_size + 2
+ for character in value:
+ codepoint = ord(character)
+ if character in {'"', "\\", "\b", "\t", "\n", "\f", "\r"}:
+ escaped_size += 1
+ elif codepoint < 0x20:
+ escaped_size += 5
+ return escaped_size
+
+
+def _integer_json_size(value: int) -> int:
+ """Return exact decimal JSON size without rendering the bounded integer."""
+ magnitude = -value if value < 0 else value
+ if magnitude == 0:
+ digits = 1
+ else:
+ bits = int.bit_length(magnitude)
+ digits = ((bits - 1) * 30103) // 100000 + 1
+ if magnitude >= 10**digits:
+ digits += 1
+ return digits + (value < 0)
+
+
+def _container_json_syntax_size(item_count: int, *, mapping: bool) -> int:
+ """Return braces/brackets plus compact separators and mapping colons."""
+ if item_count == 0:
+ return 2
+ return 2 + item_count - 1 + (item_count if mapping else 0)
+
+
+def _normalized_scalar_json_size(value: Any) -> int:
+ """Return exact compact JSON size for a normalized scalar."""
+ value_type = type(value)
+ if value is None:
+ return 4
+ if value_type is bool:
+ return 4 if value else 5
+ if value_type is str:
+ size = _json_string_size(value, MAX_RESULT_STRING_LENGTH)
+ assert size is not None
+ return size
+ if value_type is int:
+ return _integer_json_size(value)
+ if value_type is float:
+ return len(float.__repr__(value))
+ if value_type is Decimal:
+ # Pydantic serializes Decimal values as JSON strings so their exact
+ # finite value survives the wire projection without binary rounding.
+ text = Decimal.__str__(value)
+ size = _json_string_size(text, MAX_RESULT_STRING_LENGTH)
+ assert size is not None
+ return size
+ raise AssertionError("result scalar was not normalized")
+
+
+def _pydantic_scalar_json_size(value: Any) -> int:
+ """Return the scalar size emitted by Pydantic's JSON serializer."""
+ if type(value) is float:
+ # pydantic-core uses the shortest exponent (``1e-7``), while Python's
+ # repr retains a leading zero (``1e-07``).
+ return len(to_json(value))
+ return _normalized_scalar_json_size(value)
+
+
+def _charge_json_bytes(
+ budget: _ResultBudget, size: int, *, metadata: bool = False
+) -> str | None:
+ budget.json_bytes += size
+ if budget.json_bytes > MAX_QUERY_RESULT_VALUE_BYTES:
+ return "too many aggregate JSON bytes"
+ if metadata:
+ budget.metadata_bytes += size
+ if budget.metadata_bytes > MAX_QUERY_RESULT_METADATA_BYTES:
+ return "too many aggregate metadata JSON bytes"
+ return None
+
+
+def _charge_value(budget: _ResultBudget, *, metadata: bool = False) -> str |
None:
+ budget.values += 1
+ if budget.values > MAX_QUERY_RESULT_VALUES:
+ return "too many aggregate values"
+ if metadata:
+ budget.metadata_items += 1
+ if budget.metadata_items > MAX_QUERY_RESULT_METADATA_ITEMS:
+ return "too many aggregate metadata values"
+ return None
+
+
+def _charge_text(
+ value: str,
+ budget: _ResultBudget,
+ *,
+ key: bool = False,
+ metadata: bool = False,
+) -> str | None:
+ maximum = (
+ MAX_RESULT_KEY_LENGTH
+ if key
+ else MAX_QUERY_RESULT_METADATA_BYTES
+ if metadata
+ else MAX_RESULT_STRING_LENGTH
+ )
+ size = _json_string_size(value, maximum)
+ if size is None:
+ return "an invalid or oversized object key" if key else "invalid text
data"
+ return _charge_json_bytes(budget, size, metadata=metadata)
+
+
+def _integer_failure(value: int) -> str | None:
+ bits = int.bit_length(value)
+ if bits > MAX_RESULT_INTEGER_BITS:
+ return "an oversized integer"
+ digits = 1 if bits == 0 else ((bits - 1) * 30103) // 100000 + 1
+ if digits > MAX_RESULT_INTEGER_DIGITS:
+ return "an oversized integer"
+ return None
+
+
+def _decimal_failure(value: Decimal) -> str | None:
+ if Decimal.__sizeof__(value) > MAX_RESULT_DECIMAL_STORAGE:
+ return "an oversized Decimal"
+ if not Decimal.is_finite(value):
+ return "a non-finite Decimal"
+ parts = Decimal.as_tuple(value)
+ if tuple.__len__(parts.digits) > MAX_RESULT_DECIMAL_DIGITS:
+ return "an oversized Decimal"
+ exponent = parts.exponent
+ if type(exponent) is not int or abs(exponent) >
MAX_RESULT_DECIMAL_MAGNITUDE:
+ return "an oversized Decimal"
+ return None
+
+
+def _type_mro(value_type: type[Any]) -> tuple[type[Any], ...]:
+ """Read a concrete type's MRO without consulting metaclass overrides."""
+ try:
+ mro = type.__getattribute__(value_type, "__mro__")
+ except (AttributeError, TypeError): # pragma: no cover - defensive
metaclass
+ return ()
+ return mro if type(mro) is tuple else ()
+
+
+def _timezone_name_without_hooks(tzinfo: Any) -> str | None: # noqa: C901
+ """Read common pytz/dateutil zone state without dispatching timezone
hooks."""
+ value_mro = _type_mro(type(tzinfo))
+ if any(base is pytz.tzinfo.BaseTzInfo for base in value_mro):
+ for base in value_mro:
+ try:
+ namespace = type.__getattribute__(base, "__dict__")
+ except (AttributeError, TypeError): # pragma: no cover
+ continue
+ zone = namespace.get("zone")
+ if type(zone) is str and _bounded_utf8_length(zone, 256) is not
None:
+ try:
+ canonical = pytz.timezone(zone)
+ except (KeyError, ValueError):
+ return None
+ # Generated pytz types are trusted; arbitrary subclasses that
+ # inherit their internal fields are not.
+ return zone if type(canonical) is type(tzinfo) else None
+
+ if type(tzinfo) not in _DATEUTIL_NAMED_TIMEZONE_TYPES:
+ return None
+
+ try:
+ namespace = object.__getattribute__(tzinfo, "__dict__")
+ except (AttributeError, TypeError):
+ return None
+ if type(namespace) is not dict:
+ return None
+ filename = dict.get(namespace, "_filename")
+ if type(filename) is not str or _bounded_utf8_length(filename, 4_096) is
None:
+ return None
+ marker = "/zoneinfo/"
+ if (offset := str.find(filename, marker)) >= 0:
+ name = str.__getitem__(filename, slice(offset + len(marker), None))
+ elif not str.startswith(filename, "/") and str.find(filename, "\\") < 0:
+ name = filename
+ else:
+ return None
+ parts = str.split(name, "/")
+ if not parts or any(part in {"", ".", ".."} for part in parts):
+ return None
+ return name if _bounded_utf8_length(name, 256) is not None else None
+
+
+def _object_namespace(value: Any) -> dict[str, Any] | None:
+ """Read exact instance storage without descriptor dispatch."""
+ try:
+ namespace = object.__getattribute__(value, "__dict__")
+ except (AttributeError, TypeError):
+ return None
+ return namespace if type(namespace) is dict else None
+
+
+def _dateutil_ttinfo_offset_without_hooks(value: Any) -> int | None:
+ """Validate one exact dateutil transition record and return its offset."""
+ if type(value) is not _DATEUTIL_TTINFO_TYPE:
+ return None
+ try:
+ offset = object.__getattribute__(value, "offset")
+ delta = object.__getattribute__(value, "delta")
+ isdst = object.__getattribute__(value, "isdst")
+ abbreviation = object.__getattribute__(value, "abbr")
+ is_standard = object.__getattribute__(value, "isstd")
+ is_gmt = object.__getattribute__(value, "isgmt")
+ dst_offset = object.__getattribute__(value, "dstoffset")
+ except (AttributeError, TypeError):
+ return None
+ if type(offset) is not int or not -86_400 < offset < 86_400:
+ return None
+ if type(delta) is not timedelta or delta != timedelta(seconds=offset):
+ return None
+ if type(isdst) is not int or isdst not in {0, 1}:
+ return None
+ if abbreviation is not None and (
+ type(abbreviation) is not str or _bounded_utf8_length(abbreviation,
256) is None
+ ):
+ return None
+ if type(is_standard) is not bool or type(is_gmt) is not bool:
+ return None
+ if type(dst_offset) is not timedelta:
+ return None
+ if not -timedelta(days=1) < dst_offset < timedelta(days=1):
+ return None
+ return offset
+
+
+def _dateutil_named_state_without_hooks( # noqa: C901
+ tzinfo: Any,
+) -> _DateutilTimezoneState | None:
+ """Validate bounded exact dateutil tzfile state without timezone hooks."""
+ if type(tzinfo) not in _DATEUTIL_NAMED_TIMEZONE_TYPES:
+ return None
+ if _timezone_name_without_hooks(tzinfo) is None:
+ return None
+ namespace = _object_namespace(tzinfo)
+ if namespace is None:
+ return None
+ transitions = dict.get(namespace, "_trans_list")
+ utc_transitions = dict.get(namespace, "_trans_list_utc")
+ transition_info = dict.get(namespace, "_trans_idx")
+ info_list = dict.get(namespace, "_ttinfo_list")
+ standard_info = dict.get(namespace, "_ttinfo_std")
+ before_info = dict.get(namespace, "_ttinfo_before")
+ first_info = dict.get(namespace, "_ttinfo_first")
+ if (
+ type(transitions) is not tuple
+ or type(utc_transitions) is not tuple
+ or type(transition_info) is not tuple
+ or type(info_list) is not list
+ or tuple.__len__(transitions) > _MAX_DATEUTIL_TRANSITIONS
+ or tuple.__len__(utc_transitions) != tuple.__len__(transitions)
+ or tuple.__len__(transition_info) != tuple.__len__(transitions)
+ or list.__len__(info_list) == 0
+ or list.__len__(info_list) > _MAX_DATEUTIL_TTINFOS
+ ):
+ return None
+
+ previous_transition: int | None = None
+ previous_utc_transition: int | None = None
+ for index in range(tuple.__len__(transitions)):
+ transition = tuple.__getitem__(transitions, index)
+ utc_transition = tuple.__getitem__(utc_transitions, index)
+ if (
+ type(transition) is not int
+ or type(utc_transition) is not int
+ or abs(transition) > _MAX_DATEUTIL_TRANSITION_MAGNITUDE
+ or abs(utc_transition) > _MAX_DATEUTIL_TRANSITION_MAGNITUDE
+ or (previous_transition is not None and transition <=
previous_transition)
+ or (
+ previous_utc_transition is not None
+ and utc_transition <= previous_utc_transition
+ )
+ ):
+ return None
+ previous_transition = transition
+ previous_utc_transition = utc_transition
+
+ known_info_ids: set[int] = set()
+ for index in range(list.__len__(info_list)):
+ info = list.__getitem__(info_list, index)
+ if _dateutil_ttinfo_offset_without_hooks(info) is None:
+ return None
+ known_info_ids.add(id(info))
+ for info in (standard_info, before_info, first_info):
+ if info is not None and id(info) not in known_info_ids:
+ return None
+ for index in range(tuple.__len__(transition_info)):
+ if id(tuple.__getitem__(transition_info, index)) not in known_info_ids:
+ return None
+ if not transitions:
+ if (
+ standard_info is not list.__getitem__(info_list, 0)
+ or first_info is not standard_info
+ or before_info is not None
+ ):
+ return None
+ else:
+ expected_standard = None
+ expected_dst = None
+ for index in range(tuple.__len__(transition_info) - 1, -1, -1):
+ info = tuple.__getitem__(transition_info, index)
+ is_dst = object.__getattribute__(info, "isdst")
+ if expected_standard is None and not is_dst:
+ expected_standard = info
+ elif expected_dst is None and is_dst:
+ expected_dst = info
+ if expected_standard is not None and expected_dst is not None:
+ break
+ if expected_standard is None:
+ expected_standard = expected_dst
+ expected_before = None
+ for index in range(list.__len__(info_list)):
+ info = list.__getitem__(info_list, index)
+ if not object.__getattribute__(info, "isdst"):
+ expected_before = info
+ break
+ if expected_before is None:
+ expected_before = list.__getitem__(info_list, 0)
+ if standard_info is not expected_standard or before_info is not
expected_before:
+ return None
+ standard_offset = _dateutil_ttinfo_offset_without_hooks(standard_info)
+ if standard_offset is None:
+ return None
+ before_offset = (
+ _dateutil_ttinfo_offset_without_hooks(before_info)
+ if before_info is not None
+ else None
+ )
+ transition_offsets: list[int] = []
+ previous_offset: int | None = None
+ previous_base_offset: int | None = None
+ previous_is_dst: int | None = None
+ previous_dst_offset = 0
+ for index in range(tuple.__len__(transition_info)):
+ info = tuple.__getitem__(transition_info, index)
+ if id(info) not in known_info_ids:
+ return None
+ offset = _dateutil_ttinfo_offset_without_hooks(info)
+ if offset is None:
+ return None
+ is_dst = object.__getattribute__(info, "isdst")
+ dst_offset_seconds = 0
+ if previous_is_dst is not None and is_dst:
+ if not previous_is_dst:
+ assert previous_offset is not None
+ dst_offset_seconds = offset - previous_offset
+ if not dst_offset_seconds and previous_dst_offset:
+ dst_offset_seconds = previous_dst_offset
+ previous_dst_offset = dst_offset_seconds
+ base_offset = offset - dst_offset_seconds
+ adjustment = base_offset
+ if (
+ previous_base_offset is not None
+ and base_offset != previous_base_offset
+ and is_dst != previous_is_dst
+ ):
+ adjustment = previous_base_offset
+ if (
+ tuple.__getitem__(transitions, index)
+ != tuple.__getitem__(utc_transitions, index) + adjustment
+ ):
+ return None
+ transition_offsets.append(offset)
+ previous_offset = offset
+ previous_base_offset = base_offset
+ previous_is_dst = is_dst
+ if transitions and before_offset is None:
+ return None
+ return _DateutilTimezoneState(
+ transitions=transitions,
+ transition_offsets=tuple(transition_offsets),
+ standard_offset=standard_offset,
+ before_offset=before_offset,
+ )
+
+
+def _dateutil_named_offset_without_hooks( # noqa: C901
+ value: datetime, tzinfo: Any
+) -> timezone | None:
+ """Preserve dateutil's source-selected wall offset from validated state."""
+ state = _dateutil_named_state_without_hooks(tzinfo)
+ if state is None:
+ return None
+ epoch_ordinal = date.toordinal(date(1970, 1, 1))
+ wall_timestamp = (
+ (datetime.toordinal(value) - epoch_ordinal) * 86_400
+ + value.hour * 3_600
+ + value.minute * 60
+ + value.second
+ )
+ transitions = state.transitions
+ selected_offset: int | None
+ if not transitions:
+ selected_offset = state.standard_offset
+ else:
+ index = bisect_right(transitions, wall_timestamp) - 1
+
+ def offset_at(transition_index: int | None) -> int | None:
+ if transition_index is None or transition_index + 1 >=
len(transitions):
+ return state.standard_offset
+ if transition_index < 0:
+ return state.before_offset
+ return state.transition_offsets[transition_index]
+
+ if index > 0:
+ selected_offset = offset_at(index)
+ previous_offset = offset_at(index - 1)
+ if selected_offset is None or previous_offset is None:
+ return None
+ is_ambiguous = wall_timestamp < transitions[index] + (
+ previous_offset - selected_offset
+ )
+ if not value.fold and is_ambiguous:
+ index -= 1
+ selected_offset = offset_at(index)
+ if selected_offset is None:
+ return None
+ try:
+ return timezone(timedelta(seconds=selected_offset))
+ except (OverflowError, ValueError):
+ return None
+
+
+def _pytz_named_offset_without_hooks(tzinfo: Any) -> timezone | None:
+ """Return a localized pytz zone's stored offset without calling hooks."""
+ if _timezone_name_without_hooks(tzinfo) is None:
+ return None
+ namespace = _object_namespace(tzinfo)
+ offset = dict.get(namespace, "_utcoffset") if namespace is not None else
None
+ if offset is None:
+ # Static pytz zones store their fixed offset on the generated class.
+ class_namespace = type.__getattribute__(type(tzinfo), "__dict__")
+ offset = class_namespace.get("_utcoffset")
+ if type(offset) is not timedelta:
+ return None
+ try:
+ return timezone(offset)
+ except ValueError:
+ return None
+
+
+def _dateutil_local_offset_without_hooks(
+ value: datetime, tzinfo: Any
+) -> timezone | None:
+ """Select a dateutil local offset using builtin system-time data."""
+ if type(tzinfo) is not _DATEUTIL_LOCAL_TIMEZONE_TYPE:
+ return None
+ namespace = _object_namespace(tzinfo)
+ if namespace is None:
+ return None
+ standard_offset = dict.get(namespace, "_std_offset")
+ daylight_offset = dict.get(namespace, "_dst_offset")
+ has_daylight = dict.get(namespace, "_hasdst")
+ if (
+ type(standard_offset) is not timedelta
+ or type(daylight_offset) is not timedelta
+ or type(has_daylight) is not bool
+ ):
+ return None
+ selected_offset = standard_offset
+ if has_daylight:
+ epoch = datetime(1970, 1, 1)
+ naive = datetime(
+ value.year,
+ value.month,
+ value.day,
+ value.hour,
+ value.minute,
+ value.second,
+ value.microsecond,
+ )
+ timestamp = (naive - epoch).total_seconds()
+ try:
+ is_daylight = bool(
+ system_time.localtime(timestamp +
system_time.timezone).tm_isdst
+ )
+ daylight_saved = daylight_offset - standard_offset
+ previous_is_daylight = bool(
+ system_time.localtime(
+ timestamp
+ - timedelta.total_seconds(daylight_saved)
+ + system_time.timezone
+ ).tm_isdst
+ )
+ except (OverflowError, OSError, ValueError):
+ return None
+ if not is_daylight and is_daylight != previous_is_daylight:
+ is_daylight = not bool(value.fold)
+ selected_offset = daylight_offset if is_daylight else standard_offset
+ try:
+ return timezone(selected_offset)
+ except ValueError:
+ return None
+
+
+def _canonical_timezone(tzinfo: Any) -> timezone | ZoneInfo | None: # noqa:
C901
+ if any(type(tzinfo) is trusted for trusted in _TRUSTED_TIMEZONE_TYPES):
+ return tzinfo
+ if tzinfo is pytz.UTC:
+ return timezone.utc
+ if type(tzinfo) in _DATEUTIL_FIXED_TIMEZONE_TYPES:
+ try:
+ namespace = object.__getattribute__(tzinfo, "__dict__")
+ except (AttributeError, TypeError):
+ return timezone.utc if type(tzinfo) is type(dateutil_tz.tzutc())
else None
+ if type(namespace) is not dict:
+ return None
+ offset = dict.get(namespace, "_offset")
+ if type(offset) is not timedelta:
+ return timezone.utc if type(tzinfo) is type(dateutil_tz.tzutc())
else None
+ if abs(offset) >= timedelta(days=1):
+ return None
+ return timezone(offset)
+ if type(tzinfo) in _PYTZ_FIXED_TIMEZONE_TYPES:
+ try:
+ namespace = object.__getattribute__(tzinfo, "__dict__")
+ except (AttributeError, TypeError):
+ return None
+ if type(namespace) is not dict:
+ return None
+ minutes = dict.get(namespace, "_minutes")
+ if type(minutes) is not int or not -1_440 < minutes < 1_440:
+ return None
+ return timezone(timedelta(minutes=minutes))
+ return None
+
+
+def _timestamp_offset_without_hooks(value: pd.Timestamp) -> timezone | None:
+ """Recover a timestamp's stored wall-clock offset without timezone
hooks."""
+ multipliers = {"s": 1_000_000_000, "ms": 1_000_000, "us": 1_000, "ns": 1}
+ multiplier = multipliers.get(value.unit)
+ if multiplier is None:
+ return None
+ try:
+ instant_ns = int(value.asm8.view("i8")) * multiplier
+ epoch_ordinal = date.toordinal(date(1970, 1, 1))
+ wall_ns = (
+ (
+ (datetime.toordinal(value) - epoch_ordinal) * 86_400
+ + value.hour * 3600
+ + value.minute * 60
+ + value.second
+ )
+ * 1_000_000_000
+ + value.microsecond * 1000
+ + value.nanosecond
+ )
+ offset_ns = wall_ns - instant_ns
+ if offset_ns % 1000:
+ return None
+ return timezone(timedelta(microseconds=offset_ns // 1000))
+ except (OverflowError, TypeError, ValueError):
+ return None
+
+
+def _canonical_datetime(value: datetime) -> tuple[str | None, str | None]:
+ """Serialize an exact datetime through trusted timezone state only."""
+ tzinfo = value.tzinfo
+ canonical_value = value
+ if tzinfo is not None and not any(
+ type(tzinfo) is trusted for trusted in _TRUSTED_TIMEZONE_TYPES
+ ):
+ canonical_tz = (
+ _pytz_named_offset_without_hooks(tzinfo)
+ or _dateutil_local_offset_without_hooks(value, tzinfo)
+ or _dateutil_named_offset_without_hooks(value, tzinfo)
+ or _canonical_timezone(tzinfo)
+ )
+ if canonical_tz is None:
+ return None, "a datetime with an unsupported timezone"
+ canonical_value = datetime(
+ value.year,
+ value.month,
+ value.day,
+ value.hour,
+ value.minute,
+ value.second,
+ value.microsecond,
+ tzinfo=canonical_tz,
+ fold=value.fold,
+ )
+ try:
+ return datetime.isoformat(canonical_value), None
+ except (OverflowError, TypeError, ValueError):
+ return None, "an invalid datetime"
+
+
+def _canonical_time(value: time) -> tuple[str | None, str | None]:
+ """Serialize an exact time through trusted timezone state only."""
+ tzinfo = value.tzinfo
+ canonical_value = value
+ if tzinfo is not None and not any(
+ type(tzinfo) is trusted for trusted in _TRUSTED_TIMEZONE_TYPES
+ ):
+ if _dateutil_named_state_without_hooks(tzinfo) is not None:
+ canonical_value = time(
+ value.hour,
+ value.minute,
+ value.second,
+ value.microsecond,
+ fold=value.fold,
+ )
+ else:
+ canonical_tz = _canonical_timezone(tzinfo)
+ if canonical_tz is None and type(tzinfo) is
_DATEUTIL_LOCAL_TIMEZONE_TYPE:
+ namespace = _object_namespace(tzinfo)
+ if namespace is not None and dict.get(namespace, "_hasdst") is
False:
+ offset = dict.get(namespace, "_std_offset")
+ if type(offset) is timedelta:
+ try:
+ canonical_tz = timezone(offset)
+ except ValueError:
+ canonical_tz = None
+ if canonical_tz is None:
+ return None, "a time with an unsupported timezone"
+ canonical_value = time(
+ value.hour,
+ value.minute,
+ value.second,
+ value.microsecond,
+ tzinfo=canonical_tz,
+ fold=value.fold,
+ )
+ try:
+ return time.isoformat(canonical_value), None
+ except (OverflowError, TypeError, ValueError):
+ return None, "an invalid time"
+
+
+def _canonical_timestamp(value: pd.Timestamp) -> tuple[str | None, str | None]:
+ """Preserve a trusted timestamp's instant, offset, nanoseconds, and
fold."""
+ try:
+ tzinfo = value.tzinfo
+ if tzinfo is not None and not any(
+ type(tzinfo) is trusted for trusted in _TRUSTED_TIMEZONE_TYPES
+ ):
+ supported_timezone = (
+ _canonical_timezone(tzinfo) is not None
+ or _pytz_named_offset_without_hooks(tzinfo) is not None
+ or type(tzinfo) is _DATEUTIL_LOCAL_TIMEZONE_TYPE
+ or _dateutil_named_state_without_hooks(tzinfo) is not None
+ )
+ if not supported_timezone:
+ return None, "a timestamp with an unsupported timezone"
+ canonical_tz = _timestamp_offset_without_hooks(value)
+ if canonical_tz is None:
+ return None, "an invalid timestamp"
+ raw_value = value.asm8.view("i8")
+ value = pd.Timestamp(raw_value, unit=value.unit,
tz="UTC").tz_convert(
+ canonical_tz
+ )
+ return pd.Timestamp.isoformat(value), None
+ except (KeyError, OverflowError, TypeError, ValueError):
+ return None, "an invalid timestamp"
+
+
+def _normalize_scalar(value: Any) -> tuple[Any, str | None]: # noqa: C901
+ """Convert one exact trusted producer scalar to a JSON-safe scalar."""
+ value_type = type(value)
+ if value is None or value_type is bool or value_type is str:
+ return value, None
+ if value_type is int:
+ return value, _integer_failure(value)
+ if value_type is float:
+ if math.isnan(value):
+ return None, None
+ return (value, None) if math.isfinite(value) else (None, "a non-finite
number")
+ if value_type is Decimal:
+ return value, _decimal_failure(value)
+ if value_type is datetime:
+ return _canonical_datetime(value)
+ if value_type is time:
+ return _canonical_time(value)
+ if value_type is date:
+ return date.isoformat(value), None
+ if value_type is timedelta:
+ try:
+ return pd.Timedelta(value).isoformat(), None
+ except (OverflowError, TypeError, ValueError):
+ return None, "an invalid duration"
+ if value_type is UUID:
+ return UUID.__str__(value), None
+
+ if value_type is _PANDAS_NAT_TYPE or value_type is _PANDAS_NA_TYPE:
+ return None, None
+ if value_type is pd.Timestamp:
+ return _canonical_timestamp(value)
+ if value_type is pd.Timedelta:
+ if pd.isna(value):
+ return None, None
+ try:
+ return pd.Timedelta.isoformat(value), None
+ except (OverflowError, TypeError, ValueError):
+ return None, "an invalid pandas duration"
+ if value_type is _PANDAS_PERIOD_TYPE or value_type is
_PANDAS_INTERVAL_TYPE:
+ # These concrete immutable pandas extension scalars are trusted. Exact
+ # type checks deliberately exclude subclasses with conversion hooks.
+ try:
+ normalized_text = str(value)
+ except (OverflowError, TypeError, ValueError):
+ return None, "an invalid pandas scalar"
+ if _bounded_utf8_length(normalized_text, MAX_RESULT_STRING_LENGTH) is
None:
+ return None, "an oversized pandas scalar"
+ return normalized_text, None
+ if value_type in _NUMPY_INTEGER_TYPES:
+ normalized = int(value)
+ return normalized, _integer_failure(normalized)
+ if value_type in _NUMPY_FLOAT_TYPES:
+ normalized_float = float(value)
+ if math.isnan(normalized_float):
+ return None, None
+ return (
+ (normalized_float, None)
+ if math.isfinite(normalized_float)
+ else (None, "a non-finite NumPy number")
+ )
+ if value_type is np.bool_:
+ return bool(value), None
+ if value_type is np.str_:
+ return str(value), None
+ if value_type is np.datetime64:
+ if np.isnat(value):
+ return None, None
+ try:
+ return _canonical_timestamp(pd.Timestamp(value))
+ except (OverflowError, TypeError, ValueError):
+ return None, "an invalid NumPy timestamp"
+ if value_type is np.timedelta64:
+ if np.isnat(value):
+ return None, None
+ try:
+ return pd.Timedelta(value).isoformat(), None
+ except (OverflowError, TypeError, ValueError):
+ return None, "an invalid NumPy duration"
+ return None, "an unsupported or subclassed value"
Review Comment:
Fixed in 0ba0030f39 (test follow-up e0c5b4badb). `_normalize_scalar` now
accepts exact `bytes`, `bytearray` and `memoryview`. It renders them the way
the MCP serializer does: UTF-8 text when the bytes decode, otherwise
`base64:`-prefixed text (`b"\xff"` becomes `"base64:/w=="`). The raw size is
bounded before copying, and the rendered string must fit the per-cell text
budget, so only oversized binary is still rejected. Tests:
`test_binary_cells_use_the_response_serializer_encoding`, plus the rejection
cases updated to oversized binary.
##########
superset/mcp_service/chart/query_result.py:
##########
@@ -15,67 +15,1416 @@
# specific language governing permissions and limitations
# under the License.
-"""Helpers for interpreting ChartDataCommand result envelopes."""
+"""Canonicalize and validate ``ChartDataCommand`` result envelopes."""
import math
+import time as system_time
+from bisect import bisect_right
from collections.abc import Mapping
+from dataclasses import dataclass
+from datetime import date, datetime, time, timedelta, timezone
from decimal import Decimal
+from enum import Enum
from numbers import Real
from typing import Any, cast
+from uuid import UUID
+from zoneinfo import ZoneInfo
+import numpy as np
+import pandas as pd
+import pytz
+from dateutil import tz as dateutil_tz, zoneinfo as dateutil_zoneinfo
+from pydantic import BaseModel
+from pydantic_core import to_json
+
+from superset.common.chart_data import ChartDataResultFormat
+from superset.common.db_query_status import QueryStatus
+from superset.constants import CACHE_DISABLED_TIMEOUT
from superset.mcp_service.chart.schemas import ChartError
+from superset.utils.core import (
+ ExtraFiltersReasonType,
+ ExtraFiltersTimeColumnType,
+ GenericDataType,
+)
FAILED_QUERY_STATUSES = frozenset(
{"error", "failed", "stopped", "timed_out", "cancelled", "canceled"}
)
+# These are aggregate envelope limits, not per-query allowances. In particular,
+# splitting a result across the maximum number of queries must not multiply the
+# permitted rows, nodes, or encoded bytes.
+MAX_QUERY_RESULTS = 32
+MAX_QUERY_RESULT_ROWS_PER_QUERY = 50_000
+MAX_QUERY_RESULT_ROWS = 100_000
+MAX_QUERY_RESULT_COLUMNS = 4_096
+MAX_QUERY_RESULT_VALUES = 2_500_000
+MAX_QUERY_RESULT_VALUE_BYTES = 16 * 1024 * 1024
+MAX_QUERY_RESULT_METADATA_BYTES = 1024 * 1024
+MAX_QUERY_RESULT_METADATA_ITEMS = 32_768
+MAX_RESULT_VALUE_ITEMS = 4_096
+MAX_RESULT_VALUE_DEPTH = 32
+MAX_RESULT_STRING_LENGTH = 65_536
+MAX_RESULT_KEY_LENGTH = 4_096
+MAX_RESULT_INTEGER_BITS = 4_096
+MAX_RESULT_INTEGER_DIGITS = 1_234
+MAX_RESULT_DECIMAL_DIGITS = 1_024
+MAX_RESULT_DECIMAL_MAGNITUDE = 4_096
+MAX_RESULT_DECIMAL_STORAGE = 2_048
+MAX_QUERY_RESULT_ROWCOUNT = 2**63 - 1
+MAX_QUERY_RESULT_CACHE_TIMEOUT = 2**31 - 1
+MAX_QUERY_RESULT_TIMESTAMP_LENGTH = 64
+
+_ERROR_KEYS = ("error", "errors", "error_message", "message", "detail")
+_MAX_ERROR_TEXT_BYTES = 2_000
+_TRUSTED_TIMEZONE_TYPES = (timezone, ZoneInfo)
+_SAFE_RESULT_ENUM_TYPES = frozenset(
+ {
+ ChartDataResultFormat,
+ QueryStatus,
+ ExtraFiltersReasonType,
+ ExtraFiltersTimeColumnType,
+ GenericDataType,
+ }
+)
+_RESULT_FORMAT_VALUES = frozenset(
+ object.__getattribute__(member, "_value_") for member in
ChartDataResultFormat
+)
+_COLTYPE_VALUES = frozenset(
+ object.__getattribute__(member, "_value_") for member in GenericDataType
+)
+_NUMPY_INTEGER_TYPES = frozenset(
+ type(value)
+ for value in (
+ np.int8(0),
+ np.int16(0),
+ np.int32(0),
+ np.int64(0),
+ np.uint8(0),
+ np.uint16(0),
+ np.uint32(0),
+ np.uint64(0),
+ )
+)
+_NUMPY_FLOAT_TYPES = frozenset(
+ type(value)
+ for value in (np.float16(0), np.float32(0), np.float64(0),
np.longdouble(0))
+)
+_PANDAS_NAT_TYPE = type(pd.NaT)
+_PANDAS_NA_TYPE = type(pd.NA)
+_PANDAS_PERIOD_TYPE = type(pd.Period("2000-01", freq="M"))
+_PANDAS_INTERVAL_TYPE = type(pd.Interval(0, 1))
+_DATEUTIL_FIXED_TIMEZONE_TYPES = frozenset(
+ {type(dateutil_tz.tzoffset(None, 0)), type(dateutil_tz.tzutc())}
+)
+_DATEUTIL_NAMED_TIMEZONE_TYPES = frozenset(
+ {dateutil_tz.tzfile, dateutil_zoneinfo.tzfile}
+)
+_DATEUTIL_TTINFO_TYPE = type(
+ object.__getattribute__(dateutil_tz.gettz("UTC"),
"__dict__")["_ttinfo_std"]
+)
+_DATEUTIL_LOCAL_TIMEZONE_TYPE = type(dateutil_tz.tzlocal())
+_PYTZ_FIXED_TIMEZONE_TYPES = frozenset({type(pytz.FixedOffset(1))})
+_MAX_DATEUTIL_TRANSITIONS = 4_096
+_MAX_DATEUTIL_TTINFOS = 256
+_MAX_DATEUTIL_TRANSITION_MAGNITUDE = 10**12
+
+
+@dataclass
+class _ResultBudget:
+ """Aggregate counters shared by every query and metadata value."""
+
+ rows: int = 0
+ values: int = 0
+ json_bytes: int = 0
+ metadata_items: int = 0
+ metadata_bytes: int = 0
+
+
+@dataclass(frozen=True)
+class _DateutilTimezoneState:
+ """Hook-free subset of a validated exact dateutil tzfile transition
table."""
+
+ transitions: tuple[int, ...]
+ transition_offsets: tuple[int, ...]
+ standard_offset: int
+ before_offset: int | None
-def _query_error_text(value: Any) -> str | None:
- """Convert a bounded query error payload into a useful message."""
- if value is None or value is False:
+
+def _invalid_result(message: str) -> ChartError:
+ return ChartError(
+ error=f"Chart query returned {message}.",
+ error_type="InvalidQueryResult",
+ )
+
+
+def _invalid_metadata(label: str) -> ChartError:
+ return ChartError(
+ error=f"{label} returned hostile or malformed metadata.",
+ error_type="InvalidQueryResult",
+ )
+
+
+def _safe_enum_value(value: Any, expected: frozenset[type[Any]]) -> Any | None:
+ """Read trusted enum storage without invoking public conversion hooks."""
+ if type(value) not in expected or type(value) not in
_SAFE_RESULT_ENUM_TYPES:
return None
- if isinstance(value, Mapping):
- for key in ("error", "error_message", "message", "detail"):
- if text := _query_error_text(value.get(key)):
- return text
+ return object.__getattribute__(value, "_value_")
+
+
+def _bounded_utf8_length(value: str, maximum: int) -> int | None:
+ """Return the exact UTF-8 size while bounding pre-encoding work."""
+ if str.__len__(value) > maximum:
return None
- if isinstance(value, (list, tuple)):
- parts = [text for item in value if (text := _query_error_text(item))]
- return "; ".join(parts[:3]) or None
- text = str(value)
- return text[:2000] if text else None
+ try:
+ encoded = str.encode(value, "utf-8", errors="strict")
+ except UnicodeEncodeError:
+ return None
+ size = bytes.__len__(encoded)
+ return size if size <= maximum else None
-def _failure_for_query_payload(
- payload: Mapping[str, Any], label: str
-) -> ChartError | None:
- """Extract one failure from a top-level or per-query payload."""
+def _json_string_size(value: str, maximum: int) -> int | None:
+ """Return compact UTF-8 JSON string size without serializing the value."""
+ raw_size = _bounded_utf8_length(value, maximum)
+ if raw_size is None:
+ return None
+ escaped_size = raw_size + 2
+ for character in value:
+ codepoint = ord(character)
+ if character in {'"', "\\", "\b", "\t", "\n", "\f", "\r"}:
+ escaped_size += 1
+ elif codepoint < 0x20:
+ escaped_size += 5
+ return escaped_size
+
+
+def _integer_json_size(value: int) -> int:
+ """Return exact decimal JSON size without rendering the bounded integer."""
+ magnitude = -value if value < 0 else value
+ if magnitude == 0:
+ digits = 1
+ else:
+ bits = int.bit_length(magnitude)
+ digits = ((bits - 1) * 30103) // 100000 + 1
+ if magnitude >= 10**digits:
+ digits += 1
+ return digits + (value < 0)
+
+
+def _container_json_syntax_size(item_count: int, *, mapping: bool) -> int:
+ """Return braces/brackets plus compact separators and mapping colons."""
+ if item_count == 0:
+ return 2
+ return 2 + item_count - 1 + (item_count if mapping else 0)
+
+
+def _normalized_scalar_json_size(value: Any) -> int:
+ """Return exact compact JSON size for a normalized scalar."""
+ value_type = type(value)
+ if value is None:
+ return 4
+ if value_type is bool:
+ return 4 if value else 5
+ if value_type is str:
+ size = _json_string_size(value, MAX_RESULT_STRING_LENGTH)
+ assert size is not None
+ return size
+ if value_type is int:
+ return _integer_json_size(value)
+ if value_type is float:
+ return len(float.__repr__(value))
+ if value_type is Decimal:
+ # Pydantic serializes Decimal values as JSON strings so their exact
+ # finite value survives the wire projection without binary rounding.
+ text = Decimal.__str__(value)
+ size = _json_string_size(text, MAX_RESULT_STRING_LENGTH)
+ assert size is not None
+ return size
+ raise AssertionError("result scalar was not normalized")
+
+
+def _pydantic_scalar_json_size(value: Any) -> int:
+ """Return the scalar size emitted by Pydantic's JSON serializer."""
+ if type(value) is float:
+ # pydantic-core uses the shortest exponent (``1e-7``), while Python's
+ # repr retains a leading zero (``1e-07``).
+ return len(to_json(value))
+ return _normalized_scalar_json_size(value)
+
+
+def _charge_json_bytes(
+ budget: _ResultBudget, size: int, *, metadata: bool = False
+) -> str | None:
+ budget.json_bytes += size
+ if budget.json_bytes > MAX_QUERY_RESULT_VALUE_BYTES:
+ return "too many aggregate JSON bytes"
+ if metadata:
+ budget.metadata_bytes += size
+ if budget.metadata_bytes > MAX_QUERY_RESULT_METADATA_BYTES:
+ return "too many aggregate metadata JSON bytes"
+ return None
+
+
+def _charge_value(budget: _ResultBudget, *, metadata: bool = False) -> str |
None:
+ budget.values += 1
+ if budget.values > MAX_QUERY_RESULT_VALUES:
+ return "too many aggregate values"
+ if metadata:
+ budget.metadata_items += 1
+ if budget.metadata_items > MAX_QUERY_RESULT_METADATA_ITEMS:
+ return "too many aggregate metadata values"
+ return None
+
+
+def _charge_text(
+ value: str,
+ budget: _ResultBudget,
+ *,
+ key: bool = False,
+ metadata: bool = False,
+) -> str | None:
+ maximum = (
+ MAX_RESULT_KEY_LENGTH
+ if key
+ else MAX_QUERY_RESULT_METADATA_BYTES
+ if metadata
+ else MAX_RESULT_STRING_LENGTH
+ )
+ size = _json_string_size(value, maximum)
+ if size is None:
+ return "an invalid or oversized object key" if key else "invalid text
data"
+ return _charge_json_bytes(budget, size, metadata=metadata)
+
+
+def _integer_failure(value: int) -> str | None:
+ bits = int.bit_length(value)
+ if bits > MAX_RESULT_INTEGER_BITS:
+ return "an oversized integer"
+ digits = 1 if bits == 0 else ((bits - 1) * 30103) // 100000 + 1
+ if digits > MAX_RESULT_INTEGER_DIGITS:
+ return "an oversized integer"
+ return None
+
+
+def _decimal_failure(value: Decimal) -> str | None:
+ if Decimal.__sizeof__(value) > MAX_RESULT_DECIMAL_STORAGE:
+ return "an oversized Decimal"
+ if not Decimal.is_finite(value):
+ return "a non-finite Decimal"
+ parts = Decimal.as_tuple(value)
+ if tuple.__len__(parts.digits) > MAX_RESULT_DECIMAL_DIGITS:
+ return "an oversized Decimal"
+ exponent = parts.exponent
+ if type(exponent) is not int or abs(exponent) >
MAX_RESULT_DECIMAL_MAGNITUDE:
+ return "an oversized Decimal"
+ return None
+
+
+def _type_mro(value_type: type[Any]) -> tuple[type[Any], ...]:
+ """Read a concrete type's MRO without consulting metaclass overrides."""
+ try:
+ mro = type.__getattribute__(value_type, "__mro__")
+ except (AttributeError, TypeError): # pragma: no cover - defensive
metaclass
+ return ()
+ return mro if type(mro) is tuple else ()
+
+
+def _timezone_name_without_hooks(tzinfo: Any) -> str | None: # noqa: C901
+ """Read common pytz/dateutil zone state without dispatching timezone
hooks."""
+ value_mro = _type_mro(type(tzinfo))
+ if any(base is pytz.tzinfo.BaseTzInfo for base in value_mro):
+ for base in value_mro:
+ try:
+ namespace = type.__getattribute__(base, "__dict__")
+ except (AttributeError, TypeError): # pragma: no cover
+ continue
+ zone = namespace.get("zone")
+ if type(zone) is str and _bounded_utf8_length(zone, 256) is not
None:
+ try:
+ canonical = pytz.timezone(zone)
+ except (KeyError, ValueError):
+ return None
+ # Generated pytz types are trusted; arbitrary subclasses that
+ # inherit their internal fields are not.
+ return zone if type(canonical) is type(tzinfo) else None
+
+ if type(tzinfo) not in _DATEUTIL_NAMED_TIMEZONE_TYPES:
+ return None
+
+ try:
+ namespace = object.__getattribute__(tzinfo, "__dict__")
+ except (AttributeError, TypeError):
+ return None
+ if type(namespace) is not dict:
+ return None
+ filename = dict.get(namespace, "_filename")
+ if type(filename) is not str or _bounded_utf8_length(filename, 4_096) is
None:
+ return None
+ marker = "/zoneinfo/"
+ if (offset := str.find(filename, marker)) >= 0:
+ name = str.__getitem__(filename, slice(offset + len(marker), None))
+ elif not str.startswith(filename, "/") and str.find(filename, "\\") < 0:
+ name = filename
+ else:
+ return None
+ parts = str.split(name, "/")
+ if not parts or any(part in {"", ".", ".."} for part in parts):
+ return None
+ return name if _bounded_utf8_length(name, 256) is not None else None
+
+
+def _object_namespace(value: Any) -> dict[str, Any] | None:
+ """Read exact instance storage without descriptor dispatch."""
+ try:
+ namespace = object.__getattribute__(value, "__dict__")
+ except (AttributeError, TypeError):
+ return None
+ return namespace if type(namespace) is dict else None
+
+
+def _dateutil_ttinfo_offset_without_hooks(value: Any) -> int | None:
+ """Validate one exact dateutil transition record and return its offset."""
+ if type(value) is not _DATEUTIL_TTINFO_TYPE:
+ return None
+ try:
+ offset = object.__getattribute__(value, "offset")
+ delta = object.__getattribute__(value, "delta")
+ isdst = object.__getattribute__(value, "isdst")
+ abbreviation = object.__getattribute__(value, "abbr")
+ is_standard = object.__getattribute__(value, "isstd")
+ is_gmt = object.__getattribute__(value, "isgmt")
+ dst_offset = object.__getattribute__(value, "dstoffset")
+ except (AttributeError, TypeError):
+ return None
+ if type(offset) is not int or not -86_400 < offset < 86_400:
+ return None
+ if type(delta) is not timedelta or delta != timedelta(seconds=offset):
+ return None
+ if type(isdst) is not int or isdst not in {0, 1}:
+ return None
+ if abbreviation is not None and (
+ type(abbreviation) is not str or _bounded_utf8_length(abbreviation,
256) is None
+ ):
+ return None
+ if type(is_standard) is not bool or type(is_gmt) is not bool:
+ return None
+ if type(dst_offset) is not timedelta:
+ return None
+ if not -timedelta(days=1) < dst_offset < timedelta(days=1):
+ return None
+ return offset
+
+
+def _dateutil_named_state_without_hooks( # noqa: C901
+ tzinfo: Any,
+) -> _DateutilTimezoneState | None:
+ """Validate bounded exact dateutil tzfile state without timezone hooks."""
+ if type(tzinfo) not in _DATEUTIL_NAMED_TIMEZONE_TYPES:
+ return None
+ if _timezone_name_without_hooks(tzinfo) is None:
+ return None
+ namespace = _object_namespace(tzinfo)
+ if namespace is None:
+ return None
+ transitions = dict.get(namespace, "_trans_list")
+ utc_transitions = dict.get(namespace, "_trans_list_utc")
+ transition_info = dict.get(namespace, "_trans_idx")
+ info_list = dict.get(namespace, "_ttinfo_list")
+ standard_info = dict.get(namespace, "_ttinfo_std")
+ before_info = dict.get(namespace, "_ttinfo_before")
+ first_info = dict.get(namespace, "_ttinfo_first")
+ if (
+ type(transitions) is not tuple
+ or type(utc_transitions) is not tuple
+ or type(transition_info) is not tuple
+ or type(info_list) is not list
+ or tuple.__len__(transitions) > _MAX_DATEUTIL_TRANSITIONS
+ or tuple.__len__(utc_transitions) != tuple.__len__(transitions)
+ or tuple.__len__(transition_info) != tuple.__len__(transitions)
+ or list.__len__(info_list) == 0
+ or list.__len__(info_list) > _MAX_DATEUTIL_TTINFOS
+ ):
+ return None
+
+ previous_transition: int | None = None
+ previous_utc_transition: int | None = None
+ for index in range(tuple.__len__(transitions)):
+ transition = tuple.__getitem__(transitions, index)
+ utc_transition = tuple.__getitem__(utc_transitions, index)
+ if (
+ type(transition) is not int
+ or type(utc_transition) is not int
+ or abs(transition) > _MAX_DATEUTIL_TRANSITION_MAGNITUDE
+ or abs(utc_transition) > _MAX_DATEUTIL_TRANSITION_MAGNITUDE
+ or (previous_transition is not None and transition <=
previous_transition)
+ or (
+ previous_utc_transition is not None
+ and utc_transition <= previous_utc_transition
+ )
+ ):
+ return None
+ previous_transition = transition
+ previous_utc_transition = utc_transition
+
+ known_info_ids: set[int] = set()
+ for index in range(list.__len__(info_list)):
+ info = list.__getitem__(info_list, index)
+ if _dateutil_ttinfo_offset_without_hooks(info) is None:
+ return None
+ known_info_ids.add(id(info))
+ for info in (standard_info, before_info, first_info):
+ if info is not None and id(info) not in known_info_ids:
+ return None
+ for index in range(tuple.__len__(transition_info)):
+ if id(tuple.__getitem__(transition_info, index)) not in known_info_ids:
+ return None
+ if not transitions:
+ if (
+ standard_info is not list.__getitem__(info_list, 0)
+ or first_info is not standard_info
+ or before_info is not None
+ ):
+ return None
+ else:
+ expected_standard = None
+ expected_dst = None
+ for index in range(tuple.__len__(transition_info) - 1, -1, -1):
+ info = tuple.__getitem__(transition_info, index)
+ is_dst = object.__getattribute__(info, "isdst")
+ if expected_standard is None and not is_dst:
+ expected_standard = info
+ elif expected_dst is None and is_dst:
+ expected_dst = info
+ if expected_standard is not None and expected_dst is not None:
+ break
+ if expected_standard is None:
+ expected_standard = expected_dst
+ expected_before = None
+ for index in range(list.__len__(info_list)):
+ info = list.__getitem__(info_list, index)
+ if not object.__getattribute__(info, "isdst"):
+ expected_before = info
+ break
+ if expected_before is None:
+ expected_before = list.__getitem__(info_list, 0)
+ if standard_info is not expected_standard or before_info is not
expected_before:
+ return None
+ standard_offset = _dateutil_ttinfo_offset_without_hooks(standard_info)
+ if standard_offset is None:
+ return None
+ before_offset = (
+ _dateutil_ttinfo_offset_without_hooks(before_info)
+ if before_info is not None
+ else None
+ )
+ transition_offsets: list[int] = []
+ previous_offset: int | None = None
+ previous_base_offset: int | None = None
+ previous_is_dst: int | None = None
+ previous_dst_offset = 0
+ for index in range(tuple.__len__(transition_info)):
+ info = tuple.__getitem__(transition_info, index)
+ if id(info) not in known_info_ids:
+ return None
+ offset = _dateutil_ttinfo_offset_without_hooks(info)
+ if offset is None:
+ return None
+ is_dst = object.__getattribute__(info, "isdst")
+ dst_offset_seconds = 0
+ if previous_is_dst is not None and is_dst:
+ if not previous_is_dst:
+ assert previous_offset is not None
+ dst_offset_seconds = offset - previous_offset
+ if not dst_offset_seconds and previous_dst_offset:
+ dst_offset_seconds = previous_dst_offset
+ previous_dst_offset = dst_offset_seconds
+ base_offset = offset - dst_offset_seconds
+ adjustment = base_offset
+ if (
+ previous_base_offset is not None
+ and base_offset != previous_base_offset
+ and is_dst != previous_is_dst
+ ):
+ adjustment = previous_base_offset
+ if (
+ tuple.__getitem__(transitions, index)
+ != tuple.__getitem__(utc_transitions, index) + adjustment
+ ):
+ return None
+ transition_offsets.append(offset)
+ previous_offset = offset
+ previous_base_offset = base_offset
+ previous_is_dst = is_dst
+ if transitions and before_offset is None:
+ return None
+ return _DateutilTimezoneState(
+ transitions=transitions,
+ transition_offsets=tuple(transition_offsets),
+ standard_offset=standard_offset,
+ before_offset=before_offset,
+ )
+
+
+def _dateutil_named_offset_without_hooks( # noqa: C901
+ value: datetime, tzinfo: Any
+) -> timezone | None:
+ """Preserve dateutil's source-selected wall offset from validated state."""
+ state = _dateutil_named_state_without_hooks(tzinfo)
+ if state is None:
+ return None
+ epoch_ordinal = date.toordinal(date(1970, 1, 1))
+ wall_timestamp = (
+ (datetime.toordinal(value) - epoch_ordinal) * 86_400
+ + value.hour * 3_600
+ + value.minute * 60
+ + value.second
+ )
+ transitions = state.transitions
+ selected_offset: int | None
+ if not transitions:
+ selected_offset = state.standard_offset
+ else:
+ index = bisect_right(transitions, wall_timestamp) - 1
+
+ def offset_at(transition_index: int | None) -> int | None:
+ if transition_index is None or transition_index + 1 >=
len(transitions):
+ return state.standard_offset
+ if transition_index < 0:
+ return state.before_offset
+ return state.transition_offsets[transition_index]
+
+ if index > 0:
+ selected_offset = offset_at(index)
+ previous_offset = offset_at(index - 1)
+ if selected_offset is None or previous_offset is None:
+ return None
+ is_ambiguous = wall_timestamp < transitions[index] + (
+ previous_offset - selected_offset
+ )
+ if not value.fold and is_ambiguous:
+ index -= 1
+ selected_offset = offset_at(index)
+ if selected_offset is None:
+ return None
+ try:
+ return timezone(timedelta(seconds=selected_offset))
+ except (OverflowError, ValueError):
+ return None
+
+
+def _pytz_named_offset_without_hooks(tzinfo: Any) -> timezone | None:
+ """Return a localized pytz zone's stored offset without calling hooks."""
+ if _timezone_name_without_hooks(tzinfo) is None:
+ return None
+ namespace = _object_namespace(tzinfo)
+ offset = dict.get(namespace, "_utcoffset") if namespace is not None else
None
+ if offset is None:
+ # Static pytz zones store their fixed offset on the generated class.
+ class_namespace = type.__getattribute__(type(tzinfo), "__dict__")
+ offset = class_namespace.get("_utcoffset")
+ if type(offset) is not timedelta:
+ return None
+ try:
+ return timezone(offset)
+ except ValueError:
+ return None
+
+
+def _dateutil_local_offset_without_hooks(
+ value: datetime, tzinfo: Any
+) -> timezone | None:
+ """Select a dateutil local offset using builtin system-time data."""
+ if type(tzinfo) is not _DATEUTIL_LOCAL_TIMEZONE_TYPE:
+ return None
+ namespace = _object_namespace(tzinfo)
+ if namespace is None:
+ return None
+ standard_offset = dict.get(namespace, "_std_offset")
+ daylight_offset = dict.get(namespace, "_dst_offset")
+ has_daylight = dict.get(namespace, "_hasdst")
+ if (
+ type(standard_offset) is not timedelta
+ or type(daylight_offset) is not timedelta
+ or type(has_daylight) is not bool
+ ):
+ return None
+ selected_offset = standard_offset
+ if has_daylight:
+ epoch = datetime(1970, 1, 1)
+ naive = datetime(
+ value.year,
+ value.month,
+ value.day,
+ value.hour,
+ value.minute,
+ value.second,
+ value.microsecond,
+ )
+ timestamp = (naive - epoch).total_seconds()
+ try:
+ is_daylight = bool(
+ system_time.localtime(timestamp +
system_time.timezone).tm_isdst
+ )
+ daylight_saved = daylight_offset - standard_offset
+ previous_is_daylight = bool(
+ system_time.localtime(
+ timestamp
+ - timedelta.total_seconds(daylight_saved)
+ + system_time.timezone
+ ).tm_isdst
+ )
+ except (OverflowError, OSError, ValueError):
+ return None
+ if not is_daylight and is_daylight != previous_is_daylight:
+ is_daylight = not bool(value.fold)
+ selected_offset = daylight_offset if is_daylight else standard_offset
+ try:
+ return timezone(selected_offset)
+ except ValueError:
+ return None
+
+
+def _canonical_timezone(tzinfo: Any) -> timezone | ZoneInfo | None: # noqa:
C901
+ if any(type(tzinfo) is trusted for trusted in _TRUSTED_TIMEZONE_TYPES):
+ return tzinfo
+ if tzinfo is pytz.UTC:
+ return timezone.utc
+ if type(tzinfo) in _DATEUTIL_FIXED_TIMEZONE_TYPES:
+ try:
+ namespace = object.__getattribute__(tzinfo, "__dict__")
+ except (AttributeError, TypeError):
+ return timezone.utc if type(tzinfo) is type(dateutil_tz.tzutc())
else None
+ if type(namespace) is not dict:
+ return None
+ offset = dict.get(namespace, "_offset")
+ if type(offset) is not timedelta:
+ return timezone.utc if type(tzinfo) is type(dateutil_tz.tzutc())
else None
+ if abs(offset) >= timedelta(days=1):
+ return None
+ return timezone(offset)
+ if type(tzinfo) in _PYTZ_FIXED_TIMEZONE_TYPES:
+ try:
+ namespace = object.__getattribute__(tzinfo, "__dict__")
+ except (AttributeError, TypeError):
+ return None
+ if type(namespace) is not dict:
+ return None
+ minutes = dict.get(namespace, "_minutes")
+ if type(minutes) is not int or not -1_440 < minutes < 1_440:
+ return None
+ return timezone(timedelta(minutes=minutes))
+ return None
+
+
+def _timestamp_offset_without_hooks(value: pd.Timestamp) -> timezone | None:
+ """Recover a timestamp's stored wall-clock offset without timezone
hooks."""
+ multipliers = {"s": 1_000_000_000, "ms": 1_000_000, "us": 1_000, "ns": 1}
+ multiplier = multipliers.get(value.unit)
+ if multiplier is None:
+ return None
+ try:
+ instant_ns = int(value.asm8.view("i8")) * multiplier
+ epoch_ordinal = date.toordinal(date(1970, 1, 1))
+ wall_ns = (
+ (
+ (datetime.toordinal(value) - epoch_ordinal) * 86_400
+ + value.hour * 3600
+ + value.minute * 60
+ + value.second
+ )
+ * 1_000_000_000
+ + value.microsecond * 1000
+ + value.nanosecond
+ )
+ offset_ns = wall_ns - instant_ns
+ if offset_ns % 1000:
+ return None
+ return timezone(timedelta(microseconds=offset_ns // 1000))
+ except (OverflowError, TypeError, ValueError):
+ return None
+
+
+def _canonical_datetime(value: datetime) -> tuple[str | None, str | None]:
+ """Serialize an exact datetime through trusted timezone state only."""
+ tzinfo = value.tzinfo
+ canonical_value = value
+ if tzinfo is not None and not any(
+ type(tzinfo) is trusted for trusted in _TRUSTED_TIMEZONE_TYPES
+ ):
+ canonical_tz = (
+ _pytz_named_offset_without_hooks(tzinfo)
+ or _dateutil_local_offset_without_hooks(value, tzinfo)
+ or _dateutil_named_offset_without_hooks(value, tzinfo)
+ or _canonical_timezone(tzinfo)
+ )
+ if canonical_tz is None:
+ return None, "a datetime with an unsupported timezone"
+ canonical_value = datetime(
+ value.year,
+ value.month,
+ value.day,
+ value.hour,
+ value.minute,
+ value.second,
+ value.microsecond,
+ tzinfo=canonical_tz,
+ fold=value.fold,
+ )
+ try:
+ return datetime.isoformat(canonical_value), None
+ except (OverflowError, TypeError, ValueError):
+ return None, "an invalid datetime"
+
+
+def _canonical_time(value: time) -> tuple[str | None, str | None]:
+ """Serialize an exact time through trusted timezone state only."""
+ tzinfo = value.tzinfo
+ canonical_value = value
+ if tzinfo is not None and not any(
+ type(tzinfo) is trusted for trusted in _TRUSTED_TIMEZONE_TYPES
+ ):
+ if _dateutil_named_state_without_hooks(tzinfo) is not None:
+ canonical_value = time(
+ value.hour,
+ value.minute,
+ value.second,
+ value.microsecond,
+ fold=value.fold,
+ )
+ else:
+ canonical_tz = _canonical_timezone(tzinfo)
+ if canonical_tz is None and type(tzinfo) is
_DATEUTIL_LOCAL_TIMEZONE_TYPE:
+ namespace = _object_namespace(tzinfo)
+ if namespace is not None and dict.get(namespace, "_hasdst") is
False:
+ offset = dict.get(namespace, "_std_offset")
+ if type(offset) is timedelta:
+ try:
+ canonical_tz = timezone(offset)
+ except ValueError:
+ canonical_tz = None
+ if canonical_tz is None:
+ return None, "a time with an unsupported timezone"
+ canonical_value = time(
+ value.hour,
+ value.minute,
+ value.second,
+ value.microsecond,
+ tzinfo=canonical_tz,
+ fold=value.fold,
+ )
+ try:
+ return time.isoformat(canonical_value), None
+ except (OverflowError, TypeError, ValueError):
+ return None, "an invalid time"
+
+
+def _canonical_timestamp(value: pd.Timestamp) -> tuple[str | None, str | None]:
+ """Preserve a trusted timestamp's instant, offset, nanoseconds, and
fold."""
+ try:
+ tzinfo = value.tzinfo
+ if tzinfo is not None and not any(
+ type(tzinfo) is trusted for trusted in _TRUSTED_TIMEZONE_TYPES
+ ):
+ supported_timezone = (
+ _canonical_timezone(tzinfo) is not None
+ or _pytz_named_offset_without_hooks(tzinfo) is not None
+ or type(tzinfo) is _DATEUTIL_LOCAL_TIMEZONE_TYPE
+ or _dateutil_named_state_without_hooks(tzinfo) is not None
+ )
+ if not supported_timezone:
+ return None, "a timestamp with an unsupported timezone"
+ canonical_tz = _timestamp_offset_without_hooks(value)
+ if canonical_tz is None:
+ return None, "an invalid timestamp"
+ raw_value = value.asm8.view("i8")
+ value = pd.Timestamp(raw_value, unit=value.unit,
tz="UTC").tz_convert(
+ canonical_tz
+ )
+ return pd.Timestamp.isoformat(value), None
+ except (KeyError, OverflowError, TypeError, ValueError):
+ return None, "an invalid timestamp"
+
+
+def _normalize_scalar(value: Any) -> tuple[Any, str | None]: # noqa: C901
+ """Convert one exact trusted producer scalar to a JSON-safe scalar."""
+ value_type = type(value)
+ if value is None or value_type is bool or value_type is str:
+ return value, None
+ if value_type is int:
+ return value, _integer_failure(value)
+ if value_type is float:
+ if math.isnan(value):
+ return None, None
+ return (value, None) if math.isfinite(value) else (None, "a non-finite
number")
+ if value_type is Decimal:
+ return value, _decimal_failure(value)
+ if value_type is datetime:
+ return _canonical_datetime(value)
+ if value_type is time:
+ return _canonical_time(value)
+ if value_type is date:
+ return date.isoformat(value), None
+ if value_type is timedelta:
+ try:
+ return pd.Timedelta(value).isoformat(), None
+ except (OverflowError, TypeError, ValueError):
+ return None, "an invalid duration"
+ if value_type is UUID:
+ return UUID.__str__(value), None
+
+ if value_type is _PANDAS_NAT_TYPE or value_type is _PANDAS_NA_TYPE:
+ return None, None
+ if value_type is pd.Timestamp:
+ return _canonical_timestamp(value)
+ if value_type is pd.Timedelta:
+ if pd.isna(value):
+ return None, None
+ try:
+ return pd.Timedelta.isoformat(value), None
+ except (OverflowError, TypeError, ValueError):
+ return None, "an invalid pandas duration"
+ if value_type is _PANDAS_PERIOD_TYPE or value_type is
_PANDAS_INTERVAL_TYPE:
+ # These concrete immutable pandas extension scalars are trusted. Exact
+ # type checks deliberately exclude subclasses with conversion hooks.
+ try:
+ normalized_text = str(value)
+ except (OverflowError, TypeError, ValueError):
+ return None, "an invalid pandas scalar"
+ if _bounded_utf8_length(normalized_text, MAX_RESULT_STRING_LENGTH) is
None:
+ return None, "an oversized pandas scalar"
+ return normalized_text, None
+ if value_type in _NUMPY_INTEGER_TYPES:
+ normalized = int(value)
+ return normalized, _integer_failure(normalized)
+ if value_type in _NUMPY_FLOAT_TYPES:
+ normalized_float = float(value)
+ if math.isnan(normalized_float):
+ return None, None
+ return (
+ (normalized_float, None)
+ if math.isfinite(normalized_float)
+ else (None, "a non-finite NumPy number")
+ )
+ if value_type is np.bool_:
+ return bool(value), None
+ if value_type is np.str_:
+ return str(value), None
+ if value_type is np.datetime64:
+ if np.isnat(value):
+ return None, None
+ try:
+ return _canonical_timestamp(pd.Timestamp(value))
+ except (OverflowError, TypeError, ValueError):
+ return None, "an invalid NumPy timestamp"
+ if value_type is np.timedelta64:
+ if np.isnat(value):
+ return None, None
+ try:
+ return pd.Timedelta(value).isoformat(), None
+ except (OverflowError, TypeError, ValueError):
+ return None, "an invalid NumPy duration"
+ return None, "an unsupported or subclassed value"
+
+
+def _normalize_value( # noqa: C901
+ value: Any,
+ budget: _ResultBudget,
+ *,
+ enum_types: frozenset[type[Any]] = frozenset(),
+ metadata: bool = False,
+) -> tuple[Any, str | None]:
+ """Iteratively normalize one bounded exact-container value tree."""
+ stack: list[
+ tuple[Any, list[Any] | dict[str, Any] | None, int | str | None, int,
bool]
+ ] = [(value, None, None, 0, False)]
+ active_containers: set[int] = set()
+ root = value
+
+ while stack:
+ item, parent, slot, depth, leaving = stack.pop()
+ if leaving:
+ active_containers.remove(id(item))
+ continue
+ if depth > MAX_RESULT_VALUE_DEPTH:
+ return None, "excessively nested data"
+ if reason := _charge_value(budget, metadata=metadata):
+ return None, reason
+
+ if type(item) is list:
+ identity = id(item)
+ if identity in active_containers:
+ return None, "cyclic containers"
+ active_containers.add(identity)
+ width = list.__len__(item)
+ if width > MAX_RESULT_VALUE_ITEMS:
Review Comment:
Fixed in b0d5676e2d. `indexnames` no longer goes through the metadata
normalizer. `_normalize_index_names` bounds it by
`MAX_QUERY_RESULT_ROWS_PER_QUERY` and charges each label to the shared value
and byte budgets like a data cell. MultiIndex tuple labels are normalized as
arrays. Coverage uses a real FULL payload: the new
`full_producer_command_result` fixture runs a DataFrame through
`query_actions._materialize_full_payload` and `ChartDataCommand.run`. Tests:
`test_full_payload_index_metadata_uses_the_row_budget` (4,097 rows) and
`test_full_payload_index_metadata_stays_bounded_per_query`.
##########
superset/mcp_service/chart/tool/get_chart_data.py:
##########
@@ -90,6 +102,70 @@ class _ChartFacts(NamedTuple):
datasource_type: str | None = None
+def _is_expected_json_load_error(exc: Exception) -> bool:
+ """Recognize only the exact exceptions emitted for ordinary JSON
failures."""
+ return type(exc) in {TypeError, ValueError, JSONDecodeError}
+
+
+def _requested_filter_columns(extra_form_data: dict[str, Any] | None) ->
set[str]:
+ """Return simple column names explicitly requested through extra form
data."""
+ if not extra_form_data:
+ return set()
+
+ columns: set[str] = set()
+ for filter_ in extra_form_data.get("filters", []):
Review Comment:
Fixed in da9e5ba876. `get_chart_data` had a stale local copy of the
rejected-filter helpers, and that copy iterated a null `filters` list. The tool
now uses the shared `rejected_requested_filter_columns`, as master does, which
treats null lists as absent. `adhoc_filters: null` also crashed earlier, in the
core extra-form-data merge, so `merge_extra_form_data` now drops null filter
lists before that merge. Regression test:
`test_null_extra_form_data_filter_lists_return_data`, parametrized over
`filters` and `adhoc_filters` through the real tool.
##########
superset/mcp_service/chart/query_result.py:
##########
@@ -15,67 +15,1416 @@
# specific language governing permissions and limitations
# under the License.
-"""Helpers for interpreting ChartDataCommand result envelopes."""
+"""Canonicalize and validate ``ChartDataCommand`` result envelopes."""
import math
+import time as system_time
+from bisect import bisect_right
from collections.abc import Mapping
+from dataclasses import dataclass
+from datetime import date, datetime, time, timedelta, timezone
from decimal import Decimal
+from enum import Enum
from numbers import Real
from typing import Any, cast
+from uuid import UUID
+from zoneinfo import ZoneInfo
+import numpy as np
+import pandas as pd
+import pytz
+from dateutil import tz as dateutil_tz, zoneinfo as dateutil_zoneinfo
+from pydantic import BaseModel
+from pydantic_core import to_json
+
+from superset.common.chart_data import ChartDataResultFormat
+from superset.common.db_query_status import QueryStatus
+from superset.constants import CACHE_DISABLED_TIMEOUT
from superset.mcp_service.chart.schemas import ChartError
+from superset.utils.core import (
+ ExtraFiltersReasonType,
+ ExtraFiltersTimeColumnType,
+ GenericDataType,
+)
FAILED_QUERY_STATUSES = frozenset(
{"error", "failed", "stopped", "timed_out", "cancelled", "canceled"}
)
+# These are aggregate envelope limits, not per-query allowances. In particular,
+# splitting a result across the maximum number of queries must not multiply the
+# permitted rows, nodes, or encoded bytes.
+MAX_QUERY_RESULTS = 32
+MAX_QUERY_RESULT_ROWS_PER_QUERY = 50_000
+MAX_QUERY_RESULT_ROWS = 100_000
+MAX_QUERY_RESULT_COLUMNS = 4_096
+MAX_QUERY_RESULT_VALUES = 2_500_000
+MAX_QUERY_RESULT_VALUE_BYTES = 16 * 1024 * 1024
Review Comment:
Fixed in ce7c1f1274. The "Response too large" troubleshooting section now
separates the two limits. The configurable `ResponseSizeGuardMiddleware`
(`max_bytes` / `enabled`) is one. The other is the mandatory chart query result
limits, reported as `InvalidQueryResult`: rows per query and in total, columns,
values, JSON bytes and per-cell text. The section states that
`MCP_RESPONSE_SIZE_CONFIG` does not affect them and how to resolve them.
`test_troubleshooting_guide_separates_mandatory_result_limits` keeps the
documented numbers in sync with the constants.
##########
superset/mcp_service/chart/schemas.py:
##########
@@ -1634,6 +1652,1045 @@ def reject_metric_style_groupby(self) ->
"TreemapChartUpdateConfig":
return self
+class SunburstStandardizedControls(UnknownFieldCheckMixin):
+ """Bounded shared controls retained by Explore across viz changes."""
+
+ model_config = ConfigDict(extra="ignore")
+
+ metrics: list[JsonValue] = Field(default_factory=list)
+ columns: list[JsonValue] = Field(default_factory=list)
+
+
+class SunburstStandardizedFormData(UnknownFieldCheckMixin):
+ """Validated shape of Explore's cross-plugin UI memory."""
+
+ model_config = ConfigDict(extra="ignore", populate_by_name=True)
+
+ controls: SunburstStandardizedControls
+ memorized_form_data: list[tuple[str, dict[str, JsonValue]]] = Field(
+ default_factory=list,
+ validation_alias=AliasChoices("memorized_form_data",
"memorizedFormData"),
+ serialization_alias="memorizedFormData",
+ )
+
+
+_NATIVE_COLUMN_META_MAX_DEPTH = 8
+_NATIVE_COLUMN_META_MAX_CONTAINER_ITEMS = 128
+_NATIVE_COLUMN_META_MAX_TOTAL_VALUES = 512
+_NATIVE_COLUMN_META_MAX_KEY_BYTES = 1_024
+_NATIVE_COLUMN_META_MAX_STRING_BYTES = 16 * 1_024
+_NATIVE_COLUMN_META_MAX_TOTAL_STRING_BYTES = 64 * 1_024
+_NATIVE_COLUMN_META_MAX_INT_BITS = 64
+_NATIVE_METRIC_ALLOWED_KEYS = frozenset(
+ {
+ "aggregate",
+ "column",
+ "datasourceWarning",
+ "expressionType",
+ "hasCustomLabel",
+ "label",
+ "optionName",
+ "sqlExpression",
+ }
+)
+_NATIVE_METRIC_EXPRESSION_TYPES = frozenset({"SIMPLE", "SQL"})
+_NATIVE_METRIC_DISCRIMINATOR_KEYS = frozenset(
+ {
+ "column",
+ "datasourceWarning",
+ "hasCustomLabel",
+ "optionName",
+ "sqlExpression",
+ }
+)
+_NATIVE_METRIC_MISSING = object()
+_NATIVE_METRIC_INPUT_KEYS = frozenset(
+ {"metric", "metrics", "secondary_metric", "secondaryMetric"}
+)
+_NATIVE_METRIC_INVALID_VALUE = b"invalid native metric value"
+_NATIVE_METRIC_INVALID_KEY = "__invalid_native_metric_key__"
+_NATIVE_METRIC_SANITIZE_MAX_DEPTH = 16
+_NATIVE_METRIC_SANITIZE_MAX_VALUES = 1_024
+_NATIVE_METRIC_ERROR_STRING_LENGTH = 128
+
+
+def _native_column_meta_string_bytes(value: str, *, key: bool = False) -> int:
+ """Return exact UTF-8 size for a trusted built-in string."""
+ kind = "key" if key else "string"
+ try:
+ size = len(str.encode(value, "utf-8"))
+ except UnicodeEncodeError as ex:
+ raise ValueError(
+ f"Sunburst native metric column metadata {kind} is invalid"
+ ) from ex
+ limit = (
+ _NATIVE_COLUMN_META_MAX_KEY_BYTES
+ if key
+ else _NATIVE_COLUMN_META_MAX_STRING_BYTES
+ )
+ if size > limit:
+ raise ValueError(f"Sunburst native metric column metadata {kind} is
too long")
+ return size
+
+
+def _sanitize_native_metric_dict(
+ value: dict[Any, Any],
+ depth: int,
+ remaining: list[int],
+ ancestors: set[int],
+) -> tuple[Any, bool]:
+ """Project one exact dictionary through built-in operations only."""
+ value_id = id(value)
+ if value_id in ancestors:
+ return _NATIVE_METRIC_INVALID_VALUE, True
+ if dict.__len__(value) > _NATIVE_COLUMN_META_MAX_CONTAINER_ITEMS:
+ return {
+ f"invalid_{index}": _NATIVE_METRIC_INVALID_VALUE
+ for index in range(_NATIVE_COLUMN_META_MAX_CONTAINER_ITEMS + 1)
+ }, True
+
+ ancestors.add(value_id)
+ projected: dict[str, Any] = {}
+ requires_rejection = False
+ for key, item in dict.items(value):
+ if type(key) is not str:
+ projected[_NATIVE_METRIC_INVALID_KEY] =
_NATIVE_METRIC_INVALID_VALUE
+ requires_rejection = True
+ continue
+ try:
+ key_size = len(str.encode(key, "utf-8"))
+ except UnicodeEncodeError:
+ key_size = _NATIVE_COLUMN_META_MAX_KEY_BYTES + 1
+ if (
+ str.__len__(key) > _NATIVE_COLUMN_META_MAX_KEY_BYTES
+ or key_size > _NATIVE_COLUMN_META_MAX_KEY_BYTES
+ ):
+ projected["x" * (_NATIVE_COLUMN_META_MAX_KEY_BYTES + 1)] = (
+ _NATIVE_METRIC_INVALID_VALUE
+ )
+ requires_rejection = True
+ continue
+ projected_item, item_requires_rejection =
_sanitize_native_metric_value(
+ item,
+ depth=depth + 1,
+ remaining=remaining,
+ ancestors=ancestors,
+ )
+ projected[key] = projected_item
+ requires_rejection = requires_rejection or item_requires_rejection
+ ancestors.remove(value_id)
+ return projected, requires_rejection
+
+
+def _sanitize_native_metric_list(
+ value: list[Any],
+ depth: int,
+ remaining: list[int],
+ ancestors: set[int],
+) -> tuple[Any, bool]:
+ """Project one exact list through built-in operations only."""
+ value_id = id(value)
+ if value_id in ancestors:
+ return _NATIVE_METRIC_INVALID_VALUE, True
+ if list.__len__(value) > _NATIVE_COLUMN_META_MAX_CONTAINER_ITEMS:
+ return [_NATIVE_METRIC_INVALID_VALUE] * (
+ _NATIVE_COLUMN_META_MAX_CONTAINER_ITEMS + 1
+ ), True
+
+ ancestors.add(value_id)
+ projected_items: list[Any] = []
+ requires_rejection = False
+ for index in range(list.__len__(value)):
+ projected_item, item_requires_rejection =
_sanitize_native_metric_value(
+ list.__getitem__(value, index),
+ depth=depth + 1,
+ remaining=remaining,
+ ancestors=ancestors,
+ )
+ projected_items.append(projected_item)
+ requires_rejection = requires_rejection or item_requires_rejection
+ ancestors.remove(value_id)
+ return projected_items, requires_rejection
+
+
+def _sanitize_native_metric_scalar(value: Any) -> tuple[Any, bool] | None:
+ """Project an exact scalar, or signal that ``value`` is a container."""
+ value_type = type(value)
+ if value_type is str:
+ try:
+ encoded_size = len(str.encode(value, "utf-8"))
+ except UnicodeEncodeError:
+ return _NATIVE_METRIC_INVALID_VALUE, True
+ if (
+ str.__len__(value) > _NATIVE_COLUMN_META_MAX_STRING_BYTES
+ or encoded_size > _NATIVE_COLUMN_META_MAX_STRING_BYTES * 4
+ ):
+ return "x" * (_NATIVE_COLUMN_META_MAX_STRING_BYTES + 1), True
+ return (
+ str.__getitem__(value, slice(0,
_NATIVE_METRIC_ERROR_STRING_LENGTH)),
+ False,
+ )
+ if value_type is int and int.bit_length(value) >
_NATIVE_COLUMN_META_MAX_INT_BITS:
+ return 1 << (_NATIVE_COLUMN_META_MAX_INT_BITS + 1), True
+ if value_type in (int, float, bool, type(None), ColumnRef):
+ return value, False
+ return None
+
+
+def _sanitize_native_metric_value(
+ value: Any,
+ *,
+ depth: int = 0,
+ remaining: list[int] | None = None,
+ ancestors: set[int] | None = None,
+) -> tuple[Any, bool]:
+ """Build a bounded exact-type projection without dispatching object
hooks."""
+ if remaining is None:
+ remaining = [_NATIVE_METRIC_SANITIZE_MAX_VALUES]
+ if ancestors is None:
+ ancestors = set()
+ remaining[0] -= 1
+ if remaining[0] < 0 or depth > _NATIVE_METRIC_SANITIZE_MAX_DEPTH:
+ return _NATIVE_METRIC_INVALID_VALUE, True
+
+ if (scalar := _sanitize_native_metric_scalar(value)) is not None:
+ return scalar
+ value_type = type(value)
+ if value_type is dict:
+ return _sanitize_native_metric_dict(value, depth, remaining, ancestors)
+ if value_type is list:
+ return _sanitize_native_metric_list(value, depth, remaining, ancestors)
+ return _NATIVE_METRIC_INVALID_VALUE, True
+
+
+def _metric_paths_in_dict_require_safe_rejection(
+ value: dict[Any, Any], depth: int, remaining: list[int], seen: set[int]
+) -> bool:
+ """Inspect metric-bearing paths of one exact dictionary."""
+ value_id = id(value)
+ if value_id in seen:
+ return False
+ seen.add(value_id)
+ for key, nested in dict.items(value):
+ if type(key) is not str:
+ return True
+ if key in _NATIVE_METRIC_INPUT_KEYS:
+ if _sanitize_native_metric_value(nested)[1]:
+ return True
+ elif _native_metric_paths_require_safe_rejection(
+ nested, depth=depth + 1, remaining=remaining, seen=seen
+ ):
+ return True
+ return False
+
+
+def _metric_paths_in_list_require_safe_rejection(
+ value: list[Any], depth: int, remaining: list[int], seen: set[int]
+) -> bool:
+ """Inspect metric-bearing paths of one exact list."""
+ value_id = id(value)
+ if value_id in seen:
+ return False
+ seen.add(value_id)
+ return any(
+ _native_metric_paths_require_safe_rejection(
+ list.__getitem__(value, index),
+ depth=depth + 1,
+ remaining=remaining,
+ seen=seen,
+ )
+ for index in range(list.__len__(value))
+ )
+
+
+def _native_metric_paths_require_safe_rejection(
+ value: Any,
+ *,
+ depth: int = 0,
+ remaining: list[int] | None = None,
+ seen: set[int] | None = None,
+) -> bool:
+ """Find unsafe metric subtrees without consuming non-exact containers."""
+ if remaining is None:
+ remaining = [_NATIVE_METRIC_SANITIZE_MAX_VALUES]
+ if seen is None:
+ seen = set()
+ remaining[0] -= 1
+ if remaining[0] < 0 or depth > _NATIVE_METRIC_SANITIZE_MAX_DEPTH:
+ return False
+ if type(value) is dict:
+ return _metric_paths_in_dict_require_safe_rejection(
+ value, depth, remaining, seen
+ )
+ if type(value) is list:
+ return _metric_paths_in_list_require_safe_rejection(
+ value, depth, remaining, seen
+ )
+ return False
+
+
+def _replace_exact_dict_contents(
+ value: dict[Any, Any], replacement: dict[str, Any]
+) -> None:
+ """Replace exact-dict input through built-in operations only."""
+ dict.clear(value)
+ for key, item in dict.items(replacement):
+ dict.__setitem__(value, key, item)
+
+
+def _sanitize_rejected_native_metric_config(
+ data: dict[str, Any], original_data: dict[str, Any] | None
+) -> dict[str, Any]:
+ """Project a rejected config and replace Pydantic's retained input."""
+ projected, _ = _sanitize_native_metric_value(data)
+ if type(projected) is not dict:
+ projected = {_NATIVE_METRIC_INVALID_KEY: _NATIVE_METRIC_INVALID_VALUE}
+ if original_data is not None:
+ _replace_exact_dict_contents(original_data, projected)
+ return projected
+
+
+def _validate_native_metric_string(
+ value: Any, field: str, max_length: int
+) -> str | None:
+ """Validate one optional exact string in the closed native metric
wrapper."""
+ if value is None:
+ return None
+ if type(value) is not str:
+ raise ValueError(f"Sunburst native metric {field} must be a string")
+ try:
+ encoded_size = len(str.encode(value, "utf-8"))
+ except UnicodeEncodeError as ex:
+ raise ValueError(f"Sunburst native metric {field} is invalid") from ex
+ if str.__len__(value) > max_length or encoded_size > max_length * 4:
+ raise ValueError(f"Sunburst native metric {field} is too long")
+ return value
+
+
+def _inspect_native_metric_expression_type(value: Any) -> tuple[object, bool]:
+ """Inspect an exact wrapper without hashing or comparing hostile
objects."""
+ if type(value) is not dict:
+ return _NATIVE_METRIC_MISSING, False
+ if dict.__len__(value) > len(_NATIVE_METRIC_ALLOWED_KEYS):
+ raise ValueError("Sunburst native metric has too many fields")
+ expression_type: object = _NATIVE_METRIC_MISSING
+ has_native_discriminator = False
+ for key, item in dict.items(value):
+ if type(key) is not str:
+ raise ValueError("Sunburst native metric keys must be strings")
+ _validate_native_metric_string(key, "field name", 255)
+ if key in _NATIVE_METRIC_DISCRIMINATOR_KEYS:
+ has_native_discriminator = True
+ if key == "expressionType":
+ if type(item) is not str:
+ raise ValueError(
+ "Sunburst native metric expressionType must be a string"
+ )
+ if item not in _NATIVE_METRIC_EXPRESSION_TYPES:
+ raise ValueError(
+ "Sunburst native metric expressionType must be SIMPLE or
SQL"
+ )
+ expression_type = item
+ return expression_type, has_native_discriminator
+
+
+def _validate_native_metric_column_value(value: Any) -> None:
+ """Validate a wrapper column without dispatching container or scalar
hooks."""
+ if type(value) is str:
+ _validate_native_metric_string(value, "column", 255)
+ elif type(value) is dict:
+ _validate_bounded_native_column_meta(value)
+ elif value is not None:
+ raise ValueError("Sunburst native metric column must be a string or
object")
+
+
+def _validate_native_metric_item(key: str, item: Any) -> object:
+ """Validate one value after its wrapper key is known to be an exact
string."""
+ if key == "expressionType":
+ if type(item) is not str:
+ raise ValueError("Sunburst native metric expressionType must be a
string")
+ _validate_native_metric_string(item, "expressionType", 10)
+ return item
+ if key in ("label", "optionName"):
+ _validate_native_metric_string(item, key, 500)
+ elif key == "sqlExpression":
+ _validate_native_metric_string(item, key, 2_000)
+ elif key == "aggregate":
+ _validate_native_metric_string(item, key, 32)
+ elif key in ("datasourceWarning", "hasCustomLabel"):
+ if item is not None and type(item) is not bool:
+ raise ValueError(f"Sunburst native metric {key} must be a boolean")
+ elif key == "column":
+ _validate_native_metric_column_value(item)
+ return _NATIVE_METRIC_MISSING
+
+
+def _validate_native_metric_wrapper(value: Any) -> dict[str, Any]:
+ """Validate and copy the closed native metric wrapper without calling
hooks."""
+ if type(value) is not dict:
+ raise ValueError("Sunburst native metric must be a string or object")
+ if dict.__len__(value) > len(_NATIVE_METRIC_ALLOWED_KEYS):
+ raise ValueError("Sunburst native metric has too many fields")
+
+ validated: dict[str, Any] = {}
+ expression_type: object = _NATIVE_METRIC_MISSING
+ for key, item in dict.items(value):
+ if type(key) is not str:
+ raise ValueError("Sunburst native metric keys must be strings")
+ _validate_native_metric_string(key, "field name", 255)
+ if key not in _NATIVE_METRIC_ALLOWED_KEYS:
+ raise ValueError(f"Unknown Sunburst native metric field: {key}")
+ item_expression_type = _validate_native_metric_item(key, item)
+ if item_expression_type is not _NATIVE_METRIC_MISSING:
+ expression_type = item_expression_type
+ validated[key] = item
+
+ if expression_type is _NATIVE_METRIC_MISSING:
+ raise ValueError("Sunburst native metric requires expressionType")
+ if expression_type not in _NATIVE_METRIC_EXPRESSION_TYPES:
+ raise ValueError("Sunburst native metric expressionType must be SIMPLE
or SQL")
+ return validated
+
+
+def _validate_non_dict_native_metric(value: Any) -> ColumnRef:
+ """Accept an existing model while rejecting hook-bearing dict
subclasses."""
+ if type(value) is ColumnRef:
+ return value
+ raise ValueError("Sunburst native metric must be a string or object")
+
+
+def _native_column_meta_dict_children(
+ value: dict[Any, Any], depth: int
+) -> tuple[list[tuple[Any, int]], int]:
+ """Read exact-dict children while validating bounded primitive keys."""
+ if dict.__len__(value) > _NATIVE_COLUMN_META_MAX_CONTAINER_ITEMS:
+ raise ValueError("Sunburst native metric column metadata object is too
large")
+ children: list[tuple[Any, int]] = []
+ string_bytes = 0
+ for key, nested in dict.items(value):
+ if type(key) is not str:
+ raise ValueError(
+ "Sunburst native metric column metadata keys must be strings"
+ )
+ string_bytes += _native_column_meta_string_bytes(key, key=True)
+ children.append((nested, depth + 1))
+ return children, string_bytes
+
+
+def _native_column_meta_list_children(
+ value: list[Any], depth: int
+) -> list[tuple[Any, int]]:
+ """Read exact-list children without dispatching subclass hooks."""
+ if list.__len__(value) > _NATIVE_COLUMN_META_MAX_CONTAINER_ITEMS:
+ raise ValueError("Sunburst native metric column metadata array is too
large")
+ return [
+ (list.__getitem__(value, index), depth + 1)
+ for index in range(list.__len__(value))
+ ]
+
+
+def _validate_native_column_meta_scalar(value: Any) -> int:
+ """Validate one exact JSON scalar and return its string-byte
contribution."""
+ value_type = type(value)
+ if value_type is str:
+ return _native_column_meta_string_bytes(value)
+ if value_type is int:
+ if int.bit_length(value) > _NATIVE_COLUMN_META_MAX_INT_BITS:
+ raise ValueError(
+ "Sunburst native metric column metadata integer is too large"
+ )
+ return 0
+ if value_type is float:
+ if not math.isfinite(value):
+ raise ValueError(
+ "Sunburst native metric column metadata number must be finite"
+ )
+ return 0
+ if value_type not in (bool, type(None)):
+ raise ValueError(
+ "Sunburst native metric column metadata contains an invalid value"
+ )
+ return 0
+
+
+def _validate_bounded_native_column_meta(value: Any) -> Any:
+ """Validate an extensible frontend ColumnMeta object without invoking
hooks."""
+ if type(value) is not dict:
+ raise ValueError("Sunburst native metric column must be an object")
+
+ total_values = 0
+ total_string_bytes = 0
+ stack: list[tuple[Any, int]] = [(value, 0)]
+ while stack:
+ item, depth = stack.pop()
+ total_values += 1
+ if total_values > _NATIVE_COLUMN_META_MAX_TOTAL_VALUES:
+ raise ValueError("Sunburst native metric column metadata is too
large")
+ if depth > _NATIVE_COLUMN_META_MAX_DEPTH:
+ raise ValueError(
+ "Sunburst native metric column metadata is too deeply nested"
+ )
+
+ item_type = type(item)
+ if item_type is dict:
+ children, key_bytes = _native_column_meta_dict_children(item,
depth)
+ total_string_bytes += key_bytes
+ stack.extend(children)
+ elif item_type is list:
+ stack.extend(_native_column_meta_list_children(item, depth))
+ else:
+ total_string_bytes += _validate_native_column_meta_scalar(item)
+
+ if total_string_bytes > _NATIVE_COLUMN_META_MAX_TOTAL_STRING_BYTES:
+ raise ValueError("Sunburst native metric column metadata is too
large")
+ return value
+
+
+class SunburstNativeMetricColumn(BaseModel):
+ """Owned projection of the frontend's intentionally extensible
ColumnMeta."""
+
+ model_config = ConfigDict(extra="ignore", populate_by_name=True)
+
+ column_name: StrictStr = Field(
+ ...,
+ min_length=1,
+ max_length=255,
+ validation_alias=AliasChoices("column_name", "columnName"),
+ )
+ type: StrictStr | None = Field(None, max_length=255)
+
+ @model_validator(mode="before")
+ @classmethod
+ def validate_bounded_column_meta(cls, value: Any) -> Any:
+ """Bound the full open object before projecting the fields we own."""
+ projected, requires_safe_rejection =
_sanitize_native_metric_value(value)
+ if requires_safe_rejection:
+ if type(value) is dict:
+ if type(projected) is not dict:
+ projected = {
+ _NATIVE_METRIC_INVALID_KEY:
_NATIVE_METRIC_INVALID_VALUE
+ }
+ _replace_exact_dict_contents(value, projected)
+ value = projected
+ else:
+ # Let Pydantic reject the safe scalar so the original subclass
+ # is not retained as the input of an in-validator exception.
+ return projected
+ value = _validate_bounded_native_column_meta(value)
+ snake_name = dict.get(value, "column_name", _NATIVE_METRIC_MISSING)
+ camel_name = dict.get(value, "columnName", _NATIVE_METRIC_MISSING)
+ if (
+ snake_name is not _NATIVE_METRIC_MISSING
+ and camel_name is not _NATIVE_METRIC_MISSING
+ ):
+ if type(snake_name) is not str or type(camel_name) is not str:
+ raise ValueError(
+ "column_name and columnName must both be exact strings"
+ )
+ if snake_name != camel_name:
+ raise ValueError(
+ "column_name and columnName must match when both are
provided"
+ )
+ return value
+
+
+class XYNativeMetricColumn(UnknownFieldCheckMixin):
+ """Closed legacy ColumnMeta projection retained for native XY axes."""
+
+ model_config = ConfigDict(extra="ignore", populate_by_name=True)
+
+ column_name: StrictStr = Field(
+ ...,
+ min_length=1,
+ max_length=255,
+ validation_alias=AliasChoices("column_name", "columnName"),
+ )
+ id: StrictInt | None = Field(None, ge=0)
+ type: StrictStr | None = Field(None, max_length=255)
+ type_generic: StrictInt | None = Field(None, ge=0, le=10)
+ groupby: StrictBool | None = None
+ is_dttm: StrictBool | None = None
+ filterable: StrictBool | None = None
+ verbose_name: StrictStr | None = Field(None, max_length=500)
+ description: StrictStr | None = Field(None, max_length=2000)
+ expression: StrictStr | None = Field(None, max_length=2000)
+ database_expression: StrictStr | None = Field(None, max_length=2000)
+ python_date_format: StrictStr | None = Field(None, max_length=255)
+ option_name: StrictStr | None = Field(
+ None,
+ max_length=500,
+ validation_alias=AliasChoices("option_name", "optionName"),
+ )
+ filter_by: StrictStr | None = Field(
+ None,
+ max_length=255,
+ validation_alias=AliasChoices("filter_by", "filterBy"),
+ )
+ value: StrictStr | None = Field(None, max_length=500)
+ advanced_data_type: StrictStr | None = Field(None, max_length=255)
+ uuid: StrictStr | None = Field(None, max_length=64)
+
+
+class SunburstLegacySavedFields(BaseModel):
+ """Known obsolete controls tolerated only while reading saved Sunbursts."""
+
+ model_config = ConfigDict(extra="forbid")
+
+ country_fieldtype: StrictStr | None = Field(None, max_length=255)
+ entity: StrictStr | None = Field(None, max_length=255)
+ granularity: StrictStr | None = Field(None, max_length=255)
+ limit: StrictInt | StrictStr | None = None
+ markup_type: Literal["markdown", "html"] | None = None
+ show_bubbles: StrictBool | None = None
+
+ @field_validator("limit")
+ @classmethod
+ def validate_legacy_limit(cls, value: int | str | None) -> int | str |
None:
+ """Bound the ignored legacy row limit and reject non-numeric
strings."""
+ if value is None:
+ return None
+ if isinstance(value, str):
+ if not value.isascii() or not value.isdigit():
+ raise ValueError("legacy Sunburst limit must be an integer")
+ numeric_value = int(value)
+ else:
+ numeric_value = value
+ if not 1 <= numeric_value <= 50000:
+ raise ValueError("legacy Sunburst limit must be between 1 and
50000")
+ return value
+
+
+_NATIVE_FORM_DATA_MARKER = "_mcp_native_form_data"
+_SUNBURST_IGNORED_LEGACY_FIELDS = frozenset(
+ {
+ "country_fieldtype",
+ "entity",
+ "granularity",
+ "limit",
+ "markup_type",
+ "show_bubbles",
+ }
+)
+
+
+class SunburstChartConfig(BaseChartConfig):
+ """Config for the ECharts Sunburst plugin (viz_type ``sunburst_v2``)."""
+
+ model_config = ConfigDict(extra="ignore", populate_by_name=True)
+
+ chart_type: Literal["sunburst"] = "sunburst"
+ viz_type: Literal["sunburst_v2"] = Field(
+ "sunburst_v2",
+ description="Exact Superset frontend visualization tag",
+ )
+ hierarchy: List[ColumnRef] = Field(
+ ...,
+ min_length=1,
+ description=(
+ "Hierarchy dimensions in ring order, from the innermost/root level
"
+ "to the outermost/leaf level"
+ ),
+ validation_alias=AliasChoices("hierarchy", "columns", "groupby"),
+ )
+ metric: ColumnRef = Field(
+ ...,
+ description=(
+ "Primary metric used to size arcs. Use aggregate for a SIMPLE "
+ "adhoc metric, saved_metric=True for a dataset metric, or "
+ "sql_expression plus label for a SQL metric."
+ ),
+ )
+ secondary_metric: ColumnRef | None = Field(
+ None,
+ description=(
+ "Optional metric whose ratio to the primary metric drives the "
+ "sequential color scale (for example, profit/revenue represents "
+ "margin). When omitted, arc colors are categorical."
+ ),
+ validation_alias=AliasChoices("secondary_metric", "secondaryMetric"),
+ )
+ filters: List[FilterConfig] | None = Field(
+ None,
+ description=(
+ "Structured WHERE filters (column/op/value/clause). An omitted
list is "
+ "preserved on updates; an explicit [] clears saved filters."
+ ),
+ )
+ time_range: str | None = Field(
+ None,
+ min_length=1,
+ max_length=1000,
+ description=(
+ "Superset time range, for example 'Last year', "
+ "'2025-01-01 : 2025-12-31', or 'No filter'"
+ ),
+ )
+ time_grain: TimeGrain | None = Field(
+ None,
+ description="Optional bucket for temporal hierarchy columns",
+ validation_alias=AliasChoices("time_grain", "time_grain_sqla"),
+ )
+ sort_by_metric: bool = Field(
+ False,
+ description=(
+ "Order hierarchy rows by the primary metric descending before "
+ "applying row_limit, matching the frontend buildQuery transform"
+ ),
+ )
+ row_limit: int = Field(10000, description="Maximum hierarchy rows", ge=1,
le=50000)
+ color_scheme: str | None = Field(
+ None,
+ max_length=100,
+ description="Categorical scheme used when secondary_metric is omitted",
+ )
+ linear_color_scheme: str | None = Field(
+ None,
+ max_length=100,
+ description="Sequential scheme used when secondary_metric is present",
+ )
+ show_labels: bool = False
+ show_labels_threshold: float = Field(
+ 5,
+ ge=0,
+ le=100,
+ description="Minimum arc size in percentage points for showing a
label",
+ )
+ show_total: bool = False
+ show_null_values: bool = Field(
+ True,
+ description="Keep null-valued hierarchy nodes in the rendered tree",
+ )
+ label_type: Literal["key", "value", "key_value"] = "key"
+ number_format: str = Field("SMART_NUMBER", min_length=1, max_length=50)
+ date_format: str = Field("smart_date", min_length=1, max_length=50)
+ currency_format: CurrencyFormat | None = None
+ # Bounded native Explore envelope/UI state. These keys occur in real saved
+ # form_data and are safe to round-trip, but remain typed so native mode
does
+ # not become an escape hatch for arbitrary plugin keys or malformed
nesting.
+ annotation_layers: list[dict[str, JsonValue]] | None = None
+ dashboard_id: int | None = Field(
+ None,
+ validation_alias=AliasChoices("dashboard_id", "dashboardId"),
+ serialization_alias="dashboardId",
+ )
+ dashboards: list[int] | None = None
+ datasource: str | None = Field(None, min_length=1, max_length=255)
+ extra_form_data: dict[str, JsonValue] | None = None
+ slice_id: int | None = None
+ url_params: dict[str, JsonValue] | None = None
+ standardized_form_data: SunburstStandardizedFormData | None = Field(
+ None,
+ validation_alias=AliasChoices("standardized_form_data",
"standardizedFormData"),
+ serialization_alias="standardizedFormData",
+ )
+ since: str | None = Field(None, max_length=1000)
+ until: str | None = Field(None, max_length=1000)
+ time_compare: str | list[str] | None = None
+ compare_lag: int | str | None = None
+ compare_suffix: str | None = Field(None, max_length=255)
+
+ @staticmethod
+ def _looks_like_native_form_data(data: Any) -> bool:
+ """Identify saved Explore payloads without weakening typed typo
checks."""
+ if type(data) is not dict:
+ return False
+ native_marker = dict.get(data, _NATIVE_FORM_DATA_MARKER)
+ viz_type = dict.get(data, "viz_type")
+ if not (
+ (type(native_marker) is bool and native_marker)
+ or (type(viz_type) is str and viz_type == "sunburst_v2")
+ ):
+ return False
+ metric = dict.get(data, "metric")
+ expression_type, has_native_discriminator = (
+ _inspect_native_metric_expression_type(metric)
+ )
+ return (
+ any(
+ key in data
+ for key in (
+ "adhoc_filters",
+ "annotation_layers",
+ "datasource",
+ "extra_form_data",
+ "since",
+ "slice_id",
+ "standardizedFormData",
+ "until",
+ )
+ )
+ or type(metric) is str
+ or expression_type is not _NATIVE_METRIC_MISSING
+ or has_native_discriminator
+ )
+
+ @staticmethod
+ def _coerce_native_metric(
+ value: Any, *, allow_extensible_column_meta: bool = True
+ ) -> Any:
+ """Accept saved metric names and native SIMPLE/SQL metric objects."""
+ if type(value) is str:
+ return {"name": value, "saved_metric": True}
+ if type(value) is dict:
+ expression_type, has_native_discriminator = (
+ _inspect_native_metric_expression_type(value)
+ )
+ if expression_type is _NATIVE_METRIC_MISSING:
+ if has_native_discriminator:
+ raise ValueError("Sunburst native metric requires
expressionType")
+ return value
+ value = _validate_native_metric_wrapper(value)
+ else:
+ return _validate_non_dict_native_metric(value)
+
+ expression_type = value.get("expressionType")
+ label = value.get("label")
+ has_custom_label = value.get("hasCustomLabel")
+ if expression_type == "SQL":
+ sql_expression = value.get("sqlExpression")
+ return {
+ "sql_expression": sql_expression,
+ # The frontend still assigns a result label when the label is
+ # not custom. Preserve a nonempty native label first, then use
+ # the SQL expression as the same effective fallback.
+ "label": label or sql_expression,
+ "has_custom_label": has_custom_label,
+ }
+ if expression_type == "SIMPLE":
+ column = value.get("column")
+ if type(column) is dict:
+ column_model = (
+ SunburstNativeMetricColumn
+ if allow_extensible_column_meta
+ else XYNativeMetricColumn
+ )
+ column_metadata = column_model.model_validate(column)
+ name = column_metadata.column_name
+ dtype = column_metadata.type
+ elif type(column) is str:
+ name = column
+ dtype = None
+ else:
+ raise ValueError(
+ "Sunburst SIMPLE metric column must be a column name or
object"
+ )
+ return {
+ "name": name,
+ "aggregate": value.get("aggregate"),
+ "label": label,
+ "has_custom_label": has_custom_label,
+ "dtype": dtype,
+ }
+ return value
+
+ @staticmethod
+ def _coerce_native_hierarchy(data: dict[str, Any]) -> None:
+ """Normalize native ``columns``/``groupby`` dimension shortcuts."""
+ hierarchy_key = next(
+ (key for key in ("hierarchy", "columns", "groupby") if key in
data),
+ None,
+ )
+ if hierarchy_key is None:
+ return
+ hierarchy = data[hierarchy_key]
+ if isinstance(hierarchy, (str, dict)):
+ hierarchy = [hierarchy]
+ if isinstance(hierarchy, list):
+ data[hierarchy_key] = [
+ {"name": item} if isinstance(item, str) else item for item in
hierarchy
+ ]
+ if hierarchy_key != "groupby":
+ # Saved Sunburst payloads sometimes retain an empty groupby from a
+ # previous plugin. The explicit columns/hierarchy control wins.
+ data.pop("groupby", None)
+
+ @staticmethod
+ def _coerce_native_filter(
+ native_filter: Any, data: dict[str, Any]
+ ) -> dict[str, Any] | None:
+ """Convert one native SIMPLE filter, extracting temporal controls."""
+ if not isinstance(native_filter, dict):
+ raise ValueError("Each Sunburst native adhoc_filter must be an
object")
+ allowed_keys = {
+ "clause",
+ "comparator",
+ "datasourceWarning",
+ "expressionType",
+ "filterOptionName",
+ "isExtra",
+ "operator",
+ "sqlExpression",
+ "subject",
+ }
+ if unknown := set(native_filter) - allowed_keys:
+ raise ValueError(
+ "Unknown Sunburst native adhoc_filter field(s): "
+ + ", ".join(sorted(unknown))
+ )
+ if native_filter.get("expressionType") != "SIMPLE":
+ raise ValueError(
+ "Sunburst native adhoc_filters round-trip supports SIMPLE
filters "
+ "only; express filters with the typed 'filters' field"
+ )
+ operator = native_filter.get("operator")
+ clause = native_filter.get("clause", "WHERE")
+ if not isinstance(clause, str) or clause != "WHERE":
+ raise ValueError(
+ "Sunburst native SIMPLE filter clause must be 'WHERE'; "
+ "SIMPLE HAVING is unsupported because the shared query mapper "
+ "cannot preserve HAVING semantics"
+ )
+ if operator == "TEMPORAL_RANGE":
+ if "temporal_column" not in data and "granularity_sqla" not in
data:
+ data["temporal_column"] = native_filter.get("subject")
+ data.setdefault("time_range", native_filter.get("comparator"))
+ return None
Review Comment:
Fixed in 6cf22d8bb1. Native TEMPORAL_RANGE predicates are now collected and
promoted together. Neutral ("No filter") predicates restrict nothing. Duplicate
identical ranges collapse into one. More than one distinct restricting range is
rejected, since the typed config holds a single `temporal_column`/`time_range`.
A range that conflicts with an explicit saved time subject or range is also
rejected. These now raise instead of losing the ship-date restriction. Tests:
`test_native_multiple_temporal_ranges_are_rejected_not_dropped`,
`test_native_restricting_temporal_range_wins_over_neutral_binding`,
`test_native_temporal_range_on_another_bound_column_is_rejected`.
--
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]