Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
15 changes: 9 additions & 6 deletions superset/mcp_service/chart/schemas.py
Original file line number Diff line number Diff line change
Expand Up @@ -78,6 +78,7 @@
sanitize_user_input,
sanitize_user_input_with_changes,
)
from superset.mcp_service.utils.schema_utils import OmittedMeansUnchanged
from superset.mcp_service.utils.serialization import (
JsonSafeMapping,
JsonSafeRows,
Expand Down Expand Up @@ -799,7 +800,7 @@ def check_unknown_fields(cls, data: Any) -> Any:
return _check_unknown_fields(data, cls)


class BaseChartConfig(UnknownFieldCheckMixin):
class BaseChartConfig(UnknownFieldCheckMixin, OmittedMeansUnchanged):
Comment thread
tien238lnd marked this conversation as resolved.
Comment thread
tien238lnd marked this conversation as resolved.
"""Fields shared by every MCP chart configuration."""

temporal_column: str | None = Field(
Expand All @@ -826,7 +827,7 @@ def sanitize_temporal_column(cls, v: str | None) -> str | None:
)


class ColumnRef(UnknownFieldCheckMixin):
class ColumnRef(UnknownFieldCheckMixin, OmittedMeansUnchanged):
model_config = ConfigDict(extra="ignore", populate_by_name=True)

name: str | None = Field(
Expand Down Expand Up @@ -952,7 +953,7 @@ def sanitize_sql(cls, v: str | None) -> str | None:
)


class AxisConfig(UnknownFieldCheckMixin):
class AxisConfig(UnknownFieldCheckMixin, OmittedMeansUnchanged):
model_config = ConfigDict(extra="ignore")

title: str | None = Field(None, max_length=200)
Expand Down Expand Up @@ -990,7 +991,7 @@ def to_form_data(self) -> Dict[str, str]:
LEGEND_POSITION_LITERAL = Literal["top", "bottom", "left", "right"]


class FilterConfig(UnknownFieldCheckMixin):
class FilterConfig(UnknownFieldCheckMixin, OmittedMeansUnchanged):
model_config = ConfigDict(extra="ignore", populate_by_name=True)

column: str = Field(
Expand Down Expand Up @@ -2359,7 +2360,7 @@ def validate_metric_aggregate(self) -> Self:
return self


class TableColumnConfig(UnknownFieldCheckMixin):
class TableColumnConfig(UnknownFieldCheckMixin, OmittedMeansUnchanged):
"""Display formatting supported by the MCP table-chart schema."""

model_config = ConfigDict(
Expand Down Expand Up @@ -4043,7 +4044,9 @@ class GenerateExploreLinkRequest(ChartRequestNormalizerMixin, FormDataCacheContr
)


class UpdateChartRequest(ChartRequestNormalizerMixin, QueryCacheControl):
class UpdateChartRequest(
ChartRequestNormalizerMixin, OmittedMeansUnchanged, QueryCacheControl
):
model_config = ConfigDict(populate_by_name=True)

identifier: int | str = Field(
Expand Down
3 changes: 2 additions & 1 deletion superset/mcp_service/dashboard/schemas.py
Original file line number Diff line number Diff line change
Expand Up @@ -122,6 +122,7 @@
sanitize_user_input,
sanitize_user_input_with_changes,
)
from superset.mcp_service.utils.schema_utils import OmittedMeansUnchanged
from superset.mcp_service.utils.serialization import JsonSafeRows, OptionalRowCount
from superset.mcp_service.utils.url_utils import get_superset_base_url
from superset.utils.core import DatasourceType
Expand Down Expand Up @@ -826,7 +827,7 @@ def sanitize_dashboard_title(cls, v: str | None) -> str | None:
)


class UpdateDashboardRequest(BaseModel):
class UpdateDashboardRequest(OmittedMeansUnchanged):
Comment thread
tien238lnd marked this conversation as resolved.
"""Request schema for updating an existing dashboard's layout/theme/style.

All fields are optional; only the fields explicitly passed are applied.
Expand Down
7 changes: 4 additions & 3 deletions superset/mcp_service/dataset/schemas.py
Original file line number Diff line number Diff line change
Expand Up @@ -60,6 +60,7 @@
TagInfo,
)
from superset.mcp_service.utils.response_utils import humanize_timestamp
from superset.mcp_service.utils.schema_utils import OmittedMeansUnchanged
from superset.mcp_service.utils.serialization import (
JsonSafeRows,
OptionalRowCount,
Expand Down Expand Up @@ -635,7 +636,7 @@ class CreateVirtualDatasetResponse(BaseModel):
)


class MetricCurrency(BaseModel):
class MetricCurrency(OmittedMeansUnchanged):
"""Currency formatting configuration for a metric."""

symbol: str | None = Field(
Expand All @@ -648,7 +649,7 @@ class MetricCurrency(BaseModel):
)


class DatasetMetricProperties(BaseModel):
class DatasetMetricProperties(OmittedMeansUnchanged):
Comment thread
tien238lnd marked this conversation as resolved.
"""Dataset identifier and writable saved-metric properties."""

model_config = ConfigDict(populate_by_name=True)
Expand Down Expand Up @@ -946,7 +947,7 @@ class RestoreDatasetResponse(BaseModel):
)


class UpdateDatasetRequest(BaseModel):
class UpdateDatasetRequest(OmittedMeansUnchanged):
"""Request schema for update_dataset."""

model_config = ConfigDict(populate_by_name=True)
Expand Down
27 changes: 26 additions & 1 deletion superset/mcp_service/utils/schema_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -27,13 +27,38 @@
import logging
from typing import Any, Callable, List, Type, TypeVar

from pydantic import BaseModel, ValidationError
from pydantic import BaseModel, GetJsonSchemaHandler, ValidationError

logger = logging.getLogger(__name__)

T = TypeVar("T")


class OmittedMeansUnchanged(BaseModel):
"""Base for models that tell an omitted field from an explicit ``null``.

These models read ``model_fields_set``, so leaving a field out is not the
same as passing ``null``. Pydantic advertises ``"default": null`` for every
optional field, and a client that materialises those defaults then sends
nulls the caller never named, which the model reads as deliberate input.

Dropping the advertised default keeps the fields optional without handing
clients a value to fill in. Nothing else changes: an omitted field is still
unset, and an explicit ``null`` still means whatever the tool already made
it mean.
"""

@classmethod
def __get_pydantic_json_schema__(
cls, core_schema: Any, handler: GetJsonSchemaHandler
) -> dict[str, Any]:
schema = handler(core_schema)
for field in schema.get("properties", {}).values():
Comment thread
tien238lnd marked this conversation as resolved.
if "default" in field and field["default"] is None:
del field["default"]
return schema


class JSONParseError(ValueError):
"""Raised when JSON parsing fails with helpful context."""

Expand Down
106 changes: 106 additions & 0 deletions tests/unit_tests/mcp_service/test_mcp_tool_registration.py
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,7 @@

from superset.mcp_service.app import get_default_instructions, init_fastmcp_server, mcp
from superset.mcp_service.utils.response_size_utils import COMMITTED_WRITE_SPECS
from superset.mcp_service.utils.schema_utils import OmittedMeansUnchanged

# Patch target for the feature_flag_manager imported inside _apply_config_guards
_FFM_PATH = "superset.extensions.feature_flag_manager"
Expand Down Expand Up @@ -263,6 +264,111 @@ def _run(coro):
return asyncio.run(coro)


# Tools whose request model tells "field omitted" from "field set to null":
# omitting leaves the stored value alone, an explicit null is deliberate input.
# Their optional fields must not advertise a default, or a client that
# materialises defaults sends nulls for everything the caller never named.
OMITTED_MEANS_UNCHANGED_TOOLS = (
"update_chart",
"update_dashboard",
"update_dataset",
"update_dataset_metric",
)


def _request_model_schema(tool: Any) -> dict[str, Any]:
"""Return the JSON Schema of a tool's ``request`` argument."""
schema = tool.parameters or {}
request = schema.get("properties", {}).get("request", {})
for candidate in (request, *request.get("allOf", [])):
if reference := candidate.get("$ref"):
name = reference.rpartition("/")[2]
for container in ("$defs", "definitions"):
if name in schema.get(container, {}):
return schema[container][name]
raise AssertionError(
f"{tool.name}: cannot resolve request schema reference {reference!r}"
)
if "properties" in request:
return request
raise AssertionError(
f"{tool.name}: unrecognised request schema shape {sorted(request)}; "
"the null-default check below would silently pass"
)


def _null_defaults(schema: Any, path: str = "") -> list[str]:
"""Return every field at or under ``schema`` that defaults to null.

The walk goes all the way down on purpose. A tool's advertised parameters
arrive fully inlined, while ``model_json_schema()`` puts nested models in
``$defs``, and a null default is a hazard wherever it sits: the fields of
a nested model are merged into stored state the same way.
"""
hits: list[str] = []
if isinstance(schema, dict):
for name, spec in (schema.get("properties") or {}).items():
if isinstance(spec, dict) and "default" in spec and spec["default"] is None:
hits.append(f"{path}/{name}")
for key, value in schema.items():
hits.extend(_null_defaults(value, f"{path}/{key}"))
elif isinstance(schema, list):
for index, value in enumerate(schema):
hits.extend(_null_defaults(value, f"{path}/{index}"))
return sorted(hits)


def _omitted_means_unchanged_models() -> list[type[OmittedMeansUnchanged]]:
"""Return every model built on ``OmittedMeansUnchanged``."""
models: list[type[OmittedMeansUnchanged]] = []
pending: list[type[OmittedMeansUnchanged]] = [OmittedMeansUnchanged]
while pending:
for subclass in pending.pop().__subclasses__():
if subclass not in models:
models.append(subclass)
pending.append(subclass)
return models


def test_partial_update_tools_advertise_no_null_default() -> None:
"""No optional field of a partial-update tool offers null as its default."""
registered = {tool.name: tool for tool in _run(mcp.list_tools())}
advertised = {}
for name in OMITTED_MEANS_UNCHANGED_TOOLS:
tool: Any = registered.get(name)
if tool is None:
raise AssertionError(
f"{name} is not registered, so its schema cannot be checked"
)
if offenders := _null_defaults(_request_model_schema(tool)):
advertised[name] = offenders

assert not advertised, (
"Partial-update tools must not advertise null defaults, or a client "
f"filling them clears values the caller never named: {advertised}"
)


def test_omitted_means_unchanged_models_advertise_no_null_default() -> None:
"""Every model on the base keeps null out of its advertised defaults.

``update_chart`` reads ``model_fields_set`` on the nested chart config
rather than on the request, so the config models carry the base too. This
covers them, and any model that joins them later.
"""
models = _omitted_means_unchanged_models()
assert models, "OmittedMeansUnchanged has no subclasses; is the import stale?"

offenders = {
model.__name__: null_defaults
for model in models
if (null_defaults := _null_defaults(model.model_json_schema()))
}
assert not offenders, (
f"Models on OmittedMeansUnchanged must not advertise null defaults: {offenders}"
)


def test_mcp_app_imports_successfully():
"""Test that the MCP app can be imported without errors."""
assert mcp is not None
Expand Down
Loading