Compare commits

..
Author SHA1 Message Date
Amin Ghadersohi 16d136ef8d fix(mcp): context-aware recovery hints and sanitize identifier in not-found errors
- In get_chart_preview, when identifier looks like a form_data_key (long
  non-numeric string), suggest regenerating the explore link rather than
  always pointing to list_charts, which is only relevant for chart IDs.
- Truncate request.identifier to 200 chars before embedding in error
  messages across get_chart_preview, get_chart_data, and update_chart
  to prevent injection via oversized attacker-controlled identifiers.
2026-05-09 00:26:11 +00:00
Amin Ghadersohi c78658d852 fix(mcp): improve "not found" errors to suggest corresponding list_* tools
When MCP tools return "not found" errors for database, chart, dataset, or
dashboard IDs, include recovery guidance pointing to the appropriate list
tool (list_databases, list_charts, list_datasets, list_dashboards).

Affected tools: execute_sql, open_sql_lab_with_context, query_dataset,
get_chart_data, get_chart_preview, update_chart,
add_chart_to_existing_dashboard, generate_dashboard
2026-05-06 22:55:46 +00:00
5b5dd01028 fix(sqla): parenthesize calculated column expressions in WHERE clause (#39793)
Co-authored-by: Brian Donovan <briand@netflix.com>
Co-authored-by: Vitor Avila <96086495+Vitor-Avila@users.noreply.github.com>
2026-05-06 19:45:27 -03:00
bialkouandbito-code-review[bot] <188872107+bito-code-review[bot]@users.noreply.github.com> 4aa4415d8f fix(i18n): update Russian translations (#39589)
Co-authored-by: bito-code-review[bot] <188872107+bito-code-review[bot]@users.noreply.github.com>
2026-05-06 13:05:23 -04:00
Sebastian Mohr e667ceb6cf feat(themes): expose active theme mode via data-theme-mode attribute (#39063) 2026-05-06 18:17:54 +03:00
Enzo Martellucci 9aaa12c7d4 fix(reports): preserve urlParams in multi-tab report fan-out (#39884) 2026-05-06 16:29:45 +02:00
Alexandru Soare adfbbf1433 fix(sql): quote identifiers in transpile_to_dialect to fix case-sensitive column filters (#39521) 2026-05-06 10:53:09 +03:00
dependabot[bot]anddependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com> d7663a9a1c chore(deps-dev): update denodo-sqlalchemy requirement from ~=1.0.6 to >=1.0.6,<2.1.0 (#39832)
Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-05-05 22:17:21 -07:00
dependabot[bot]anddependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com> 7290d3c452 chore(deps-dev): update pyathena requirement from <3,>=2 to >=2,<4 (#39830)
Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-05-05 22:17:00 -07:00
dependabot[bot]anddependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com> d7beffcec1 chore(deps-dev): bump eslint-plugin-react-you-might-not-need-an-effect from 0.9.3 to 0.10.0 in /superset-frontend (#39853)
Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-05-05 22:15:10 -07:00
dependabot[bot]anddependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com> f018b67895 chore(deps-dev): update sqlalchemy-vertica-python requirement from <0.6,>=0.5.9 to >=0.5.9,<0.7 (#39831)
Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-05-05 22:14:08 -07:00
dependabot[bot]anddependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com> 5e2c6d8c9e chore(deps): bump nanoid from 5.1.9 to 5.1.11 in /superset-frontend (#39820)
Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-05-05 22:13:52 -07:00
dependabot[bot]anddependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com> b305c8681c chore(deps-dev): update impyla requirement from <0.17,>0.16.2 to >0.16.2,<0.23 (#39833)
Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-05-05 22:09:37 -07:00
dependabot[bot]dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>Đỗ Trọng Hải
d578fa1949 chore(deps): bump @deck.gl/mapbox from 9.3.1 to 9.3.2 in /superset-frontend (#39814)
Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
Co-authored-by: Đỗ Trọng Hải <41283691+hainenber@users.noreply.github.com>
2026-05-05 22:09:33 -07:00
dependabot[bot]anddependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com> 14d28c34fd chore(deps-dev): update cx-oracle requirement from <8.1,>8.0.0 to >8.0.0,<8.4 (#39753)
Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
2026-05-05 22:05:54 -07:00
dependabot[bot]dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>Đỗ Trọng Hải
c06aee8513 chore(deps-dev): bump jsdom from 29.1.0 to 29.1.1 in /superset-frontend (#39815)
Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
Co-authored-by: Đỗ Trọng Hải <41283691+hainenber@users.noreply.github.com>
2026-05-05 22:04:47 -07:00
dependabot[bot]dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>Evan Rusackas
d0ef19953a chore(deps): bump memoize-one from 5.2.1 to 6.0.0 in /superset-frontend/plugins/plugin-chart-ag-grid-table (#37910)
Signed-off-by: dependabot[bot] <support@github.com>
Co-authored-by: dependabot[bot] <49699333+dependabot[bot]@users.noreply.github.com>
Co-authored-by: Evan Rusackas <evan@rusackas.com>
2026-05-05 21:38:49 -07:00
Vitor Avila 3745e37182 fix(OAuth2): Support OAuth2 exception with legacy endpoint (#39897) 2026-05-05 21:21:48 -03:00
Joe LiandClaude Opus 4.6 4b17ac2629 fix(explore): add matrixify_enable guard to prevent stale validators on pre-revamp charts (#38765)
Co-authored-by: Claude Opus 4.6 <noreply@anthropic.com>
2026-05-05 16:45:38 -07:00
Amin Ghadersohi 4a21a5365f fix(mcp): validate column refs in generate_explore_link, update_chart_preview, and update_chart (#39797) 2026-05-05 19:12:31 -04:00
71 changed files with 4187 additions and 1229 deletions
+1 -18
View File
@@ -24,24 +24,6 @@ assists people when migrating to a new version.
## Next
### `SSH_TUNNEL_MANAGER_CLASS` replaced by `ENGINE_MANAGER_CLASS`
The `SSH_TUNNEL_MANAGER_CLASS` config setting, the `superset.extensions.ssh` module (containing `SSHManager` and `SSHManagerFactory`), and the `ssh_manager_factory` extension singleton have been removed. SQLAlchemy engine creation — including SSH tunnel construction and URL rewriting — is now centralized in `EngineManager` (`superset/engines/manager.py`), wired up via `EngineManagerExtension` (`superset/extensions/engine_manager.py`).
A new config setting, `ENGINE_MANAGER_CLASS` (default: `"superset.engines.manager.EngineManager"`), replaces `SSH_TUNNEL_MANAGER_CLASS` as the customization hook. Deployments that previously subclassed `SSHManager` (e.g. for bastion routing, audit logging, host-key policy, or custom credential handling) should subclass `EngineManager` instead and set `ENGINE_MANAGER_CLASS` to the dotted path of the subclass. Override the relevant methods:
| Old `SSHManager` method | New override point on `EngineManager` |
|---|---|
| `__init__(app)` reading `SSH_TUNNEL_*` configs | `__init__` — the same `SSH_TUNNEL_LOCAL_BIND_ADDRESS`, `SSH_TUNNEL_TIMEOUT_SEC`, and `SSH_TUNNEL_PACKET_TIMEOUT_SEC` configs are still loaded by `EngineManagerExtension.init_app` and passed in |
| `create_tunnel(ssh_tunnel, uri)` | `_get_tunnel_kwargs(ssh_tunnel, uri)` for parameter construction and `_create_tunnel(ssh_tunnel, uri)` for the `sshtunnel.open_tunnel` + `start()` call |
| `build_sqla_url(url, server)` | Inlined in `get_engine` as `uri.set(host=tunnel.local_bind_address[0], port=tunnel.local_bind_port)` |
**Behavioral note:** the old `SSHManager.create_tunnel` passed `debug_level=logging.getLogger("flask_appbuilder").level` to `sshtunnel.open_tunnel`. The new `_get_tunnel_kwargs` does not. Subclasses relying on that should add it back in their override.
### `Database.get_sqla_engine(nullpool=...)` deprecated
The `nullpool` keyword argument to `Database.get_sqla_engine` is deprecated and ignored — the engine manager always uses `NullPool`. The kwarg is still accepted (with a `DeprecationWarning`) so external callers passing `nullpool=False` won't fail with `TypeError`, but the resulting engine will use `NullPool` regardless. Remove the argument from your callers; it will be deleted in a future release.
### Granular Export Controls
A new feature flag `GRANULAR_EXPORT_CONTROLS` introduces three fine-grained permissions that replace the legacy `can_csv` permission:
@@ -132,6 +114,7 @@ DISTRIBUTED_COORDINATION_CONFIG = {
```
See `superset/config.py` for complete configuration options.
### WebSocket config for GAQ with Docker
[35896](https://github.com/apache/superset/pull/35896) and [37624](https://github.com/apache/superset/pull/37624) updated documentation on how to run and configure Superset with Docker. Specifically for the WebSocket configuration, a new `docker/superset-websocket/config.example.json` was added to the repo, so that users could copy it to create a `docker/superset-websocket/config.json` file. The existing `docker/superset-websocket/config.json` was removed and git-ignored, so if you're using GAQ / WebSocket make sure to:
+5 -5
View File
@@ -114,7 +114,7 @@ dependencies = [
[project.optional-dependencies]
athena = ["pyathena[pandas]>=2, <3"]
athena = ["pyathena[pandas]>=2, <4"]
aurora-data-api = ["preset-sqlalchemy-aurora-data-api>=0.2.8,<0.3"]
bigquery = [
"pandas-gbq>=0.19.1",
@@ -135,7 +135,7 @@ databricks = [
"databricks-sqlalchemy==1.0.5",
]
db2 = ["ibm-db-sa>0.3.8, <=0.4.0"]
denodo = ["denodo-sqlalchemy~=1.0.6"]
denodo = ["denodo-sqlalchemy>=1.0.6,<2.1.0"]
dremio = ["sqlalchemy-dremio>=1.2.1, <4"]
drill = ["sqlalchemy-drill>=1.1.4, <2"]
druid = ["pydruid>=0.6.5,<0.7"]
@@ -158,7 +158,7 @@ hive = [
"thrift>=0.14.1, <1.0.0",
"thrift_sasl>=0.4.3, < 1.0.0",
]
impala = ["impyla>0.16.2, <0.17"]
impala = ["impyla>0.16.2, <0.23"]
kusto = ["sqlalchemy-kusto>=3.0.0, <4"]
kylin = ["kylinpy>=2.8.1, <2.9"]
mssql = ["pymssql>=2.2.8, <3"]
@@ -171,7 +171,7 @@ ocient = [
"shapely",
"geojson",
]
oracle = ["cx-Oracle>8.0.0, <8.1"]
oracle = ["cx-Oracle>8.0.0, <8.4"]
parseable = ["sqlalchemy-parseable>=0.1.3,<0.2.0"]
pinot = ["pinotdb>=5.0.0, <6.0.0"]
playwright = ["playwright>=1.37.0, <2"]
@@ -197,7 +197,7 @@ tdengine = [
]
teradata = ["teradatasql>=16.20.0.23"]
thumbnails = [] # deprecated, will be removed in 7.0
vertica = ["sqlalchemy-vertica-python>=0.5.9, < 0.6"]
vertica = ["sqlalchemy-vertica-python>= 0.5.9, < 0.7"]
netezza = ["nzalchemy>=11.0.2"]
starrocks = ["starrocks>=1.0.0"]
doris = ["pydoris>=1.0.0, <2.0.0"]
+24 -18
View File
@@ -115,7 +115,7 @@
"memoize-one": "^5.2.1",
"mousetrap": "^1.6.5",
"mustache": "^4.2.0",
"nanoid": "^5.1.9",
"nanoid": "^5.1.11",
"ol": "^10.9.0",
"pretty-ms": "^9.3.0",
"query-string": "9.3.1",
@@ -249,7 +249,7 @@
"eslint-plugin-no-only-tests": "^3.4.0",
"eslint-plugin-prettier": "^5.5.5",
"eslint-plugin-react-prefer-function-component": "^5.0.0",
"eslint-plugin-react-you-might-not-need-an-effect": "^0.9.3",
"eslint-plugin-react-you-might-not-need-an-effect": "^0.10.0",
"eslint-plugin-storybook": "^0.8.0",
"eslint-plugin-testing-library": "^7.16.2",
"eslint-plugin-theme-colors": "file:eslint-rules/eslint-plugin-theme-colors",
@@ -264,7 +264,7 @@
"jest-html-reporter": "^4.4.0",
"jest-websocket-mock": "^2.5.0",
"js-yaml-loader": "^1.2.2",
"jsdom": "^29.1.0",
"jsdom": "^29.1.1",
"lerna": "^9.0.4",
"lightningcss": "^1.32.0",
"mini-css-extract-plugin": "^2.10.2",
@@ -22872,9 +22872,9 @@
"license": "MIT"
},
"node_modules/eslint-plugin-react-you-might-not-need-an-effect": {
"version": "0.9.3",
"resolved": "https://registry.npmjs.org/eslint-plugin-react-you-might-not-need-an-effect/-/eslint-plugin-react-you-might-not-need-an-effect-0.9.3.tgz",
"integrity": "sha512-44cce7LndBnpDRWBTQ8p7ircIdl2rJBP5+V9Ik64E935UB47uA9ZMU1Uv160lAMhtvoPYqXBjQ+tojr5JF3mFQ==",
"version": "0.10.0",
"resolved": "https://registry.npmjs.org/eslint-plugin-react-you-might-not-need-an-effect/-/eslint-plugin-react-you-might-not-need-an-effect-0.10.0.tgz",
"integrity": "sha512-a4pugbQc2zLiE2NZGuXdTjtMNvlP2984QFPDv71eskUYDzigLFYfBL4QjK+RnRtcboHoXRKOcQqEZKxiK6KegA==",
"dev": true,
"license": "MIT",
"dependencies": {
@@ -31681,9 +31681,9 @@
}
},
"node_modules/jsdom": {
"version": "29.1.0",
"resolved": "https://registry.npmjs.org/jsdom/-/jsdom-29.1.0.tgz",
"integrity": "sha512-YNUc7fB9QuvSSQWfrH0xF+TyABkxUwx8sswgIDaCrw4Hol8BghdZDkITtZheRJeMtzWlnTfsM3bBBusRvpO1wg==",
"version": "29.1.1",
"resolved": "https://registry.npmjs.org/jsdom/-/jsdom-29.1.1.tgz",
"integrity": "sha512-ECi4Fi2f7BdJtUKTflYRTiaMxIB0O6zfR1fX0GXpUrf6flp8QIYn1UT20YQqdSOfk2dfkCwS8LAFoJDEppNK5Q==",
"dev": true,
"license": "MIT",
"dependencies": {
@@ -36703,9 +36703,9 @@
}
},
"node_modules/nanoid": {
"version": "5.1.9",
"resolved": "https://registry.npmjs.org/nanoid/-/nanoid-5.1.9.tgz",
"integrity": "sha512-ZUvP7KeBLe3OZ1ypw6dI/TzYJuvHP77IM4Ry73waSQTLn8/g8rpdjfyVAh7t1/+FjBtG4lCP42MEbDxOsRpBMw==",
"version": "5.1.11",
"resolved": "https://registry.npmjs.org/nanoid/-/nanoid-5.1.11.tgz",
"integrity": "sha512-v+KEsUv2ps74PaSKv0gHTxTCgMXOIfBEbaqa6w6ISIGC7ZsvHN4N9oJ8d4cmf0n5oTzQz2SLmThbQWhjd/8eKg==",
"funding": [
{
"type": "github",
@@ -50575,7 +50575,7 @@
"classnames": "^2.5.1",
"d3-array": "^3.2.4",
"lodash": "^4.18.1",
"memoize-one": "^5.2.1",
"memoize-one": "^6.0.0",
"react-table": "^7.8.0",
"regenerator-runtime": "^0.14.1",
"xss": "^1.0.15"
@@ -50606,6 +50606,12 @@
"node": ">=12"
}
},
"plugins/plugin-chart-ag-grid-table/node_modules/memoize-one": {
"version": "6.0.0",
"resolved": "https://registry.npmjs.org/memoize-one/-/memoize-one-6.0.0.tgz",
"integrity": "sha512-rkpe71W0N0c0Xz6QD0eJETuWAJGnJ9afsl1srmwPrI+yBCkge5EycXXbYRyvL29zZVUWQCY7InPRCv3GDXuZNw==",
"license": "MIT"
},
"plugins/plugin-chart-cartodiagram": {
"name": "@superset-ui/plugin-chart-cartodiagram",
"version": "0.0.1",
@@ -50887,7 +50893,7 @@
"@deck.gl/extensions": "~9.2.9",
"@deck.gl/geo-layers": "~9.2.5",
"@deck.gl/layers": "~9.2.5",
"@deck.gl/mapbox": "~9.3.1",
"@deck.gl/mapbox": "^9.3.2",
"@deck.gl/mesh-layers": "~9.2.5",
"@luma.gl/constants": "~9.2.5",
"@luma.gl/core": "~9.2.5",
@@ -50935,16 +50941,16 @@
}
},
"plugins/preset-chart-deckgl/node_modules/@deck.gl/mapbox": {
"version": "9.3.1",
"resolved": "https://registry.npmjs.org/@deck.gl/mapbox/-/mapbox-9.3.1.tgz",
"integrity": "sha512-4SgpWMeZiqiZEiz9yPdr89cVRL8HFcvXLxXUA0ExhMreUdNuK/j2OIQHPhw6vp1xCFbJEEqRelQ0pJYkhGDkYw==",
"version": "9.3.2",
"resolved": "https://registry.npmjs.org/@deck.gl/mapbox/-/mapbox-9.3.2.tgz",
"integrity": "sha512-+T9pJwsOXwjUxyGN6oiBMfIs28VtDIG1V1Rqz4qqn4TjjNEFFw+xO0olJIg8FO5IAqw2OtePdsrMj0tX8tHdGQ==",
"license": "MIT",
"dependencies": {
"@math.gl/web-mercator": "^4.1.0"
},
"peerDependencies": {
"@deck.gl/core": "~9.3.0",
"@luma.gl/core": "~9.3.2",
"@luma.gl/core": "~9.3.3",
"@math.gl/web-mercator": "^4.1.0"
}
},
+3 -3
View File
@@ -196,7 +196,7 @@
"memoize-one": "^5.2.1",
"mousetrap": "^1.6.5",
"mustache": "^4.2.0",
"nanoid": "^5.1.9",
"nanoid": "^5.1.11",
"ol": "^10.9.0",
"pretty-ms": "^9.3.0",
"query-string": "9.3.1",
@@ -330,7 +330,7 @@
"eslint-plugin-no-only-tests": "^3.4.0",
"eslint-plugin-prettier": "^5.5.5",
"eslint-plugin-react-prefer-function-component": "^5.0.0",
"eslint-plugin-react-you-might-not-need-an-effect": "^0.9.3",
"eslint-plugin-react-you-might-not-need-an-effect": "^0.10.0",
"eslint-plugin-storybook": "^0.8.0",
"eslint-plugin-testing-library": "^7.16.2",
"eslint-plugin-theme-colors": "file:eslint-rules/eslint-plugin-theme-colors",
@@ -345,7 +345,7 @@
"jest-html-reporter": "^4.4.0",
"jest-websocket-mock": "^2.5.0",
"js-yaml-loader": "^1.2.2",
"jsdom": "^29.1.0",
"jsdom": "^29.1.1",
"lerna": "^9.0.4",
"lightningcss": "^1.32.0",
"mini-css-extract-plugin": "^2.10.2",
@@ -18,6 +18,7 @@
*/
import { isMatrixifyVisible } from './matrixifyControls';
import type { ControlStateMapping } from '../types';
/**
* Helper to build a controls object matching the shape used by
@@ -25,7 +26,7 @@ import { isMatrixifyVisible } from './matrixifyControls';
*/
function makeControls(
overrides: Record<string, unknown> = {},
): Record<string, { value: unknown }> {
): ControlStateMapping {
const defaults: Record<string, unknown> = {
matrixify_enable: false,
matrixify_mode_rows: 'disabled',
@@ -36,7 +37,7 @@ function makeControls(
const merged = { ...defaults, ...overrides };
return Object.fromEntries(
Object.entries(merged).map(([k, v]) => [k, { value: v }]),
);
) as ControlStateMapping;
}
// ── matrixify_enable guard ──────────────────────────────────────────
@@ -20,7 +20,7 @@
import { t } from '@apache-superset/core/translation';
import { validateNonEmpty } from '@superset-ui/core';
import { SharedControlConfig } from '../types';
import { ControlStateMapping, SharedControlConfig } from '../types';
import { dndAdhocMetricControl } from './dndControls';
import { defineSavedMetrics } from '../utils';
@@ -29,9 +29,12 @@ import { defineSavedMetrics } from '../utils';
* Controls for transforming charts into matrix/grid layouts
*/
// Utility function to check if matrixify controls should be visible
// Utility function to check if matrixify controls should be visible.
// Controls both visibility callbacks and validator injection via mapStateToProps.
// The matrixify_enable guard prevents hidden validators from firing on
// pre-revamp charts with stale matrixify_mode defaults (fix for #38519).
const isMatrixifyVisible = (
controls: any,
controls: ControlStateMapping | undefined,
axis: 'rows' | 'columns',
mode?: 'metrics' | 'dimensions',
selectionMode?: 'members' | 'topn' | 'all',
@@ -0,0 +1,238 @@
/**
* Licensed to the Apache Software Foundation (ASF) under one
* or more contributor license agreements. See the NOTICE file
* distributed with this work for additional information
* regarding copyright ownership. The ASF licenses this file
* to you under the Apache License, Version 2.0 (the
* "License"); you may not use this file except in compliance
* with the License. You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing,
* software distributed under the License is distributed on an
* "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
* KIND, either express or implied. See the License for the
* specific language governing permissions and limitations
* under the License.
*/
/**
* Tests for the matrixify_enable guard in isMatrixifyVisible() and
* validator injection via mapStateToProps on real matrixify control definitions.
*
* These are TDD tests for the fix to apache/superset#38519 regression:
* isMatrixifyVisible() must check matrixify_enable before evaluating mode,
* otherwise pre-revamp charts with stale matrixify_mode defaults trigger
* hidden validators that block save.
*/
import {
matrixifyControls,
isMatrixifyVisible,
} from '../../src/shared-controls/matrixifyControls';
import type { ControlPanelState, ControlStateMapping } from '../../src/types';
// Helper: build a minimal controls object for ControlPanelState
const buildControls = (
overrides: Record<string, any> = {},
): ControlStateMapping => {
const controls: Record<string, { value: any }> = {};
Object.entries(overrides).forEach(([key, value]) => {
controls[key] = { value };
});
return controls as ControlStateMapping;
};
// Helper: build a minimal ControlPanelState for mapStateToProps.
// Only provides fields that isMatrixifyVisible and mapStateToProps actually read.
const buildState = (
controlValues: Record<string, any> = {},
formData: Record<string, any> = {},
) =>
({
controls: buildControls(controlValues),
datasource: { columns: [], type: 'table' },
form_data: formData,
common: {},
metadata: {},
slice: { slice_id: 0 },
}) as unknown as ControlPanelState;
// ============================================================
// Validator injection tests via real mapStateToProps (rows)
// ============================================================
// --- matrixify_dimension_rows ---
test('matrixify_dimension_rows: validators empty when matrixify_enable is falsy', () => {
const control = matrixifyControls.matrixify_dimension_rows;
const state = buildState(
{
matrixify_enable: undefined,
matrixify_mode_rows: 'dimensions',
matrixify_dimension_selection_mode_rows: 'members',
},
{ matrixify_mode_rows: 'dimensions' },
);
const result = control.mapStateToProps!(state, {} as any);
expect(result.validators).toEqual([]);
});
test('matrixify_dimension_rows: validators present when matrixify_enable is true', () => {
const control = matrixifyControls.matrixify_dimension_rows;
const state = buildState(
{
matrixify_enable: true,
matrixify_mode_rows: 'dimensions',
matrixify_dimension_selection_mode_rows: 'members',
},
{ matrixify_mode_rows: 'dimensions' },
);
const result = control.mapStateToProps!(state, {} as any);
expect(result.validators.length).toBeGreaterThan(0);
});
// --- matrixify_topn_value_rows ---
test('matrixify_topn_value_rows: validators empty when matrixify_enable is falsy', () => {
const control = matrixifyControls.matrixify_topn_value_rows;
const state = buildState(
{
matrixify_enable: undefined,
matrixify_mode_rows: 'dimensions',
matrixify_dimension_selection_mode_rows: 'topn',
},
{ matrixify_mode_rows: 'dimensions' },
);
const result = control.mapStateToProps!(state, {} as any);
expect(result.validators).toEqual([]);
});
test('matrixify_topn_value_rows: validators present when matrixify_enable is true', () => {
const control = matrixifyControls.matrixify_topn_value_rows;
const state = buildState(
{
matrixify_enable: true,
matrixify_mode_rows: 'dimensions',
matrixify_dimension_selection_mode_rows: 'topn',
},
{ matrixify_mode_rows: 'dimensions' },
);
const result = control.mapStateToProps!(state, {} as any);
expect(result.validators.length).toBeGreaterThan(0);
});
// --- matrixify_topn_metric_rows ---
test('matrixify_topn_metric_rows: validators empty when matrixify_enable is falsy', () => {
const control = matrixifyControls.matrixify_topn_metric_rows;
const state = buildState(
{
matrixify_enable: undefined,
matrixify_mode_rows: 'dimensions',
matrixify_dimension_selection_mode_rows: 'topn',
},
{ matrixify_mode_rows: 'dimensions' },
);
const result = control.mapStateToProps!(state, {} as any);
expect(result.validators).toEqual([]);
});
test('matrixify_topn_metric_rows: validators present when matrixify_enable is true', () => {
const control = matrixifyControls.matrixify_topn_metric_rows;
const state = buildState(
{
matrixify_enable: true,
matrixify_mode_rows: 'dimensions',
matrixify_dimension_selection_mode_rows: 'topn',
},
{ matrixify_mode_rows: 'dimensions' },
);
const result = control.mapStateToProps!(state, {} as any);
expect(result.validators.length).toBeGreaterThan(0);
});
// ============================================================
// Validator injection tests via real mapStateToProps (columns)
// ============================================================
test('matrixify_dimension_columns: validators empty when matrixify_enable is falsy', () => {
const control = matrixifyControls.matrixify_dimension_columns;
const state = buildState(
{
matrixify_enable: undefined,
matrixify_mode_columns: 'dimensions',
matrixify_dimension_selection_mode_columns: 'members',
},
{ matrixify_mode_columns: 'dimensions' },
);
const result = control.mapStateToProps!(state, {} as any);
expect(result.validators).toEqual([]);
});
test('matrixify_dimension_columns: validators present when matrixify_enable is true', () => {
const control = matrixifyControls.matrixify_dimension_columns;
const state = buildState(
{
matrixify_enable: true,
matrixify_mode_columns: 'dimensions',
matrixify_dimension_selection_mode_columns: 'members',
},
{ matrixify_mode_columns: 'dimensions' },
);
const result = control.mapStateToProps!(state, {} as any);
expect(result.validators.length).toBeGreaterThan(0);
});
// ============================================================
// Direct isMatrixifyVisible guard tests
// ============================================================
test.each([
['undefined', undefined],
['null', null],
['false', false],
['0', 0],
])(
'isMatrixifyVisible returns false when matrixify_enable is %s',
(_, value) => {
const controls = buildControls({
matrixify_enable: value,
matrixify_mode_rows: 'dimensions',
});
expect(isMatrixifyVisible(controls, 'rows')).toBe(false);
},
);
test('isMatrixifyVisible returns true when matrixify_enable is true and mode matches', () => {
const controls = buildControls({
matrixify_enable: true,
matrixify_mode_rows: 'dimensions',
});
expect(isMatrixifyVisible(controls, 'rows', 'dimensions')).toBe(true);
});
test('isMatrixifyVisible returns false when matrixify_enable is true but mode is disabled', () => {
const controls = buildControls({
matrixify_enable: true,
matrixify_mode_rows: 'disabled',
});
expect(isMatrixifyVisible(controls, 'rows')).toBe(false);
});
test('isMatrixifyVisible returns true when matrixify_enable is true and any non-disabled mode (no mode filter)', () => {
const controls = buildControls({
matrixify_enable: true,
matrixify_mode_columns: 'metrics',
});
expect(isMatrixifyVisible(controls, 'columns')).toBe(true);
});
@@ -29,7 +29,7 @@
"classnames": "^2.5.1",
"d3-array": "^3.2.4",
"lodash": "^4.18.1",
"memoize-one": "^5.2.1",
"memoize-one": "^6.0.0",
"react-table": "^7.8.0",
"regenerator-runtime": "^0.14.1",
"xss": "^1.0.15"
@@ -0,0 +1,99 @@
/**
* Licensed to the Apache Software Foundation (ASF) under one
* or more contributor license agreements. See the NOTICE file
* distributed with this work for additional information
* regarding copyright ownership. The ASF licenses this file
* to you under the Apache License, Version 2.0 (the
* "License"); you may not use this file except in compliance
* with the License. You may obtain a copy of the License at
*
* http://www.apache.org/licenses/LICENSE-2.0
*
* Unless required by applicable law or agreed to in writing,
* software distributed under the License is distributed on an
* "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
* KIND, either express or implied. See the License for the
* specific language governing permissions and limitations
* under the License.
*/
/**
* Regression coverage for memoize-one v6 adoption.
*
* memoize-one v6 changed the signature of the (optional) custom `isEqual`
* callback from per-argument `(a, b) => bool` to arg-array
* `(newArgs, lastArgs) => bool`. Of the four memoizeOne callsites in
* `src/transformProps.ts` (`processComparisonDataRecords`,
* `processDataRecords`, `processColumns`, `getBasicColorFormatter`), only
* `processColumns` passes a custom comparator (`isEqualColumns`); its
* signature already takes arg-arrays and is compatible with v6. The other
* three rely on memoize-one's default referential-equality comparator, which
* is unchanged between v5 and v6.
*
* These tests lock those assumptions in by observing the memoization
* behavior through the public `transformProps` API: identical chart-props
* input references should produce referentially-equal `data` and `columns`
* arrays (cache hit), while inputs that differ on the sub-fields each
* memoizer actually compares should produce fresh arrays (cache miss).
*/
import transformProps from '../src/transformProps';
import testData from '../../plugin-chart-table/test/testData';
test('transformProps returns referentially-equal data/columns on identical input (cache hit)', () => {
// processColumns and processDataRecords are both wrapped by memoizeOne at
// module scope. Two consecutive calls with the same chartProps reference
// should hit both caches and yield the same output references.
const first = transformProps(testData.basic);
const second = transformProps(testData.basic);
expect(second.columns).toBe(first.columns);
expect(second.data).toBe(first.data);
});
test('transformProps busts its memoization caches when sub-field inputs change (cache miss)', () => {
const first = transformProps(testData.basic);
// `processColumns` is wrapped with a custom equality (`isEqualColumns`) that
// compares specific chartProps sub-fields by identity — mutating only the
// top-level props reference is NOT enough to bust it. Here we supply a fresh
// `datasource.columnFormats` reference, which `isEqualColumns` compares with
// `===`, forcing `processColumns` to recompute and return a new `columns`
// array.
//
// `processDataRecords` uses memoize-one's default referential equality on
// `(data, columns)`. We also hand it a fresh `queriesData[0].data` array, so
// together with the recomputed `columns` reference it too cache-misses.
const freshProps = {
...testData.basic,
datasource: {
...testData.basic.datasource,
columnFormats: {},
},
queriesData: [
{
...testData.basic.queriesData[0],
data: [...(testData.basic.queriesData[0].data || [])],
},
],
};
const second = transformProps(freshProps);
expect(second.columns).not.toBe(first.columns);
expect(second.data).not.toBe(first.data);
});
test('transformProps memoizes the comparison-mode data pipeline on identical input', () => {
// Exercises `processComparisonDataRecords` (the third of four memoizeOne
// callsites in transformProps.ts) via the `comparison` fixture, which has
// `time_compare` set and therefore flows through the comparison branch
// where `passedData = comparisonData`.
//
// Note: we don't assert reference equality on `columns` here because the
// comparison branch runs `comparisonColumns` through the non-memoized
// `processComparisonColumns` helper, which returns a fresh array on each
// call by design.
const first = transformProps(testData.comparison);
const second = transformProps(testData.comparison);
expect(second.data).toBe(first.data);
});
@@ -29,7 +29,7 @@
"@deck.gl/extensions": "~9.2.9",
"@deck.gl/geo-layers": "~9.2.5",
"@deck.gl/layers": "~9.2.5",
"@deck.gl/mapbox": "~9.3.1",
"@deck.gl/mapbox": "~9.3.2",
"@deck.gl/mesh-layers": "~9.2.5",
"@luma.gl/constants": "~9.2.5",
"@luma.gl/core": "~9.2.5",
@@ -17,7 +17,6 @@
* under the License.
*/
import { getExtensionsRegistry } from '@superset-ui/core';
import type { ComponentType, ReactNode } from 'react';
import { Provider as ReduxProvider } from 'react-redux';
import { QueryParamProvider } from 'use-query-params';
import { ReactRouter5Adapter } from 'use-query-params/adapters/react-router-5';
@@ -65,7 +64,7 @@ export const EmbeddedContextProviders: React.FC<{
}> = ({ children }) => {
const RootContextProviderExtension = extensionsRegistry.get(
'root.context.provider',
) as ComponentType<{ children?: ReactNode }> | undefined;
);
return (
<SupersetThemeProvider themeController={themeController}>
@@ -24,7 +24,7 @@ import {
ExplorePageState,
} from 'src/explore/types';
import { getChartKey } from 'src/explore/exploreUtils';
import { getControlsState } from 'src/explore/store';
import { getControlsState, handleDeprecatedControls } from 'src/explore/store';
import { Dispatch } from 'redux';
import {
Currency,
@@ -116,6 +116,12 @@ export const hydrateExplore =
]),
);
// Normalize deprecated controls (e.g., migrate old per-axis matrixify
// flags to matrixify_enable) before form_data is stored in Redux state.
// getControlsState also calls this on its own copy, but state.form_data
// must reflect the same migration so the two stay consistent.
handleDeprecatedControls(initialFormData);
const initialExploreState = {
form_data: initialFormData,
slice: initialSlice,
@@ -188,9 +188,7 @@ function CollectionControl({
// Two items can collide when keyAccessor returns falsy and the index
// fallback is used — breaking dnd-kit reordering and React reconciliation.
// Assign a stable nanoid per item ref when no key is available.
const generatedIdsRef = useRef<WeakMap<CollectionItem, string>>(
new WeakMap(),
);
const generatedIdsRef = useRef<WeakMap<CollectionItem, string>>(new WeakMap());
const itemIds = useMemo(
() =>
value.map(item => {
+352 -48
View File
@@ -17,55 +17,359 @@
* under the License.
*/
import { getChartControlPanelRegistry } from '@superset-ui/core';
import { applyDefaultFormData } from 'src/explore/store';
import {
applyDefaultFormData,
getControlsState,
handleDeprecatedControls,
} from 'src/explore/store';
// eslint-disable-next-line no-restricted-globals -- TODO: Migrate from describe blocks
describe('store', () => {
beforeAll(() => {
getChartControlPanelRegistry().registerValue('test-chart', {
controlPanelSections: [
{
label: 'Test section',
expanded: true,
controlSetRows: [['row_limit']],
},
],
});
});
// eslint-disable-next-line @typescript-eslint/no-explicit-any
(window as any).featureFlags = {};
afterAll(() => {
getChartControlPanelRegistry().remove('test-chart');
});
// eslint-disable-next-line no-restricted-globals -- TODO: Migrate from describe blocks
describe('applyDefaultFormData', () => {
// eslint-disable-next-line @typescript-eslint/no-explicit-any
(window as any).featureFlags = {};
test('applies default to formData if the key is missing', () => {
const inputFormData = {
datasource: '11_table',
viz_type: 'test-chart',
};
let outputFormData = applyDefaultFormData(inputFormData);
expect(outputFormData.row_limit).toEqual(10000);
const inputWithRowLimit = {
...inputFormData,
row_limit: 888,
};
outputFormData = applyDefaultFormData(inputWithRowLimit);
expect(outputFormData.row_limit).toEqual(888);
});
test('keeps null if key is defined with null', () => {
const inputFormData = {
datasource: '11_table',
viz_type: 'test-chart',
row_limit: null,
};
const outputFormData = applyDefaultFormData(inputFormData);
expect(outputFormData.row_limit).toBe(null);
});
beforeAll(() => {
getChartControlPanelRegistry().registerValue('test-chart', {
controlPanelSections: [
{
label: 'Test section',
expanded: true,
controlSetRows: [['row_limit']],
},
],
});
});
afterAll(() => {
getChartControlPanelRegistry().remove('test-chart');
});
// Helper: build ExploreState for getControlsState
const buildExploreState = (controlOverrides: Record<string, any> = {}) => ({
datasource: { type: 'table' },
controls: Object.fromEntries(
Object.entries(controlOverrides).map(([k, v]) => [k, { value: v }]),
),
});
// ============================================================
// Existing applyDefaultFormData tests
// ============================================================
test('applyDefaultFormData applies default to formData if the key is missing', () => {
const inputFormData = {
datasource: '11_table',
viz_type: 'test-chart',
};
let outputFormData = applyDefaultFormData(inputFormData);
expect(outputFormData.row_limit).toEqual(10000);
const inputWithRowLimit = {
...inputFormData,
row_limit: 888,
};
outputFormData = applyDefaultFormData(inputWithRowLimit);
expect(outputFormData.row_limit).toEqual(888);
});
test('applyDefaultFormData keeps null if key is defined with null', () => {
const inputFormData = {
datasource: '11_table',
viz_type: 'test-chart',
row_limit: null,
};
const outputFormData = applyDefaultFormData(inputFormData);
expect(outputFormData.row_limit).toBe(null);
});
// ============================================================
// Migration tests: handleDeprecatedControls normalizes stale matrixify modes
// (fix for apache/superset#38519 regression — guards validators AND
// downstream UI consumers that infer matrixify state from mode values)
// ============================================================
test('getControlsState resets stale matrixify_mode_rows to disabled when matrixify_enable key absent', () => {
const state = buildExploreState();
const formData = {
datasource: '1__table',
viz_type: 'test-chart',
matrixify_mode_rows: 'dimensions', // stale pre-revamp default
};
const result = getControlsState(state as any, formData as any);
const modeControl = result.matrixify_mode_rows as any;
expect(modeControl?.value).toBe('disabled');
});
test('getControlsState resets stale matrixify_mode_columns to disabled when matrixify_enable key absent', () => {
const state = buildExploreState();
const formData = {
datasource: '1__table',
viz_type: 'test-chart',
matrixify_mode_columns: 'metrics', // stale pre-revamp default
};
const result = getControlsState(state as any, formData as any);
const modeControl = result.matrixify_mode_columns as any;
expect(modeControl?.value).toBe('disabled');
});
test('getControlsState preserves matrixify mode values when matrixify_enable is true', () => {
const state = buildExploreState();
const formData = {
datasource: '1__table',
viz_type: 'test-chart',
matrixify_enable: true,
matrixify_mode_rows: 'dimensions',
};
const result = getControlsState(state as any, formData as any);
const modeControl = result.matrixify_mode_rows as any;
expect(modeControl?.value).toBe('dimensions');
});
test('getControlsState preserves matrixify mode values when matrixify_enable is explicitly false', () => {
const state = buildExploreState();
const formData = {
datasource: '1__table',
viz_type: 'test-chart',
matrixify_enable: false,
matrixify_mode_rows: 'dimensions',
};
const result = getControlsState(state as any, formData as any);
const modeControl = result.matrixify_mode_rows as any;
// matrixify_enable key IS present (just false) — migration does NOT fire
expect(modeControl?.value).toBe('dimensions');
});
test('getControlsState is idempotent when matrixify modes already disabled', () => {
const state = buildExploreState();
const formData = {
datasource: '1__table',
viz_type: 'test-chart',
matrixify_mode_rows: 'disabled',
matrixify_mode_columns: 'disabled',
};
const result = getControlsState(state as any, formData as any);
expect((result.matrixify_mode_rows as any)?.value).toBe('disabled');
expect((result.matrixify_mode_columns as any)?.value).toBe('disabled');
});
test('getControlsState handles form_data with no matrixify keys', () => {
const state = buildExploreState();
const formData = {
datasource: '1__table',
viz_type: 'test-chart',
};
const result = getControlsState(state as any, formData as any);
// Controls should get their defaults — matrixify_mode defaults to 'disabled'
expect((result.matrixify_mode_rows as any)?.value).toBe('disabled');
expect((result.matrixify_mode_columns as any)?.value).toBe('disabled');
});
test('getControlsState round-trip: pre-revamp form_data produces no matrixify validation errors', () => {
// Simulate a chart saved before #38519 with stale matrixify defaults
// Empty controls: on real first-load hydration, no pre-existing controls exist
const state = buildExploreState();
const preRevampFormData = {
datasource: '1__table',
viz_type: 'test-chart',
// Stale old defaults — no matrixify_enable key (legacy chart)
matrixify_mode_rows: 'dimensions',
matrixify_mode_columns: 'metrics',
};
const result = getControlsState(state as any, preRevampFormData as any);
// Every matrixify control should have zero validation errors
const matrixifyControlEntries = Object.entries(result).filter(([name]) =>
name.startsWith('matrixify_'),
);
const controlsWithErrors = matrixifyControlEntries.filter(
([, control]) => (control as any)?.validationErrors?.length > 0,
);
expect(controlsWithErrors).toEqual([]);
});
// ============================================================
// Dashboard hydration: applyDefaultFormData with stale form_data
// ============================================================
test('applyDefaultFormData normalizes stale matrixify modes for legacy charts', () => {
// Dashboard hydration now runs handleDeprecatedControls too, so stale
// matrixify modes from pre-revamp charts are normalized to 'disabled'.
// This protects downstream consumers (ChartContextMenu, DrillBySubmenu,
// ChartRenderer) that infer "matrixify is active" from mode values alone.
const preRevampFormData = {
datasource: '1__table',
viz_type: 'test-chart',
matrixify_mode_rows: 'dimensions',
matrixify_mode_columns: 'metrics',
// No matrixify_enable key — legacy chart that never used matrixify
};
const outputFormData = applyDefaultFormData(preRevampFormData as any);
// Stale values are now normalized to 'disabled'
expect(outputFormData.matrixify_mode_rows).toBe('disabled');
expect(outputFormData.matrixify_mode_columns).toBe('disabled');
expect(outputFormData.matrixify_enable).toBe(false);
});
// ============================================================
// P1: Pre-revamp charts that actually used matrixify via old per-axis flags
// (matrixify_enable_vertical_layout / matrixify_enable_horizontal_layout)
// ============================================================
test('getControlsState preserves modes and sets matrixify_enable when old vertical flag is true', () => {
const state = buildExploreState();
const formData = {
datasource: '1__table',
viz_type: 'test-chart',
matrixify_enable_vertical_layout: true,
matrixify_mode_rows: 'dimensions',
matrixify_mode_columns: 'metrics',
};
const result = getControlsState(state as any, formData as any);
// Vertical layout was enabled — rows mode preserved, matrixify_enable migrated
expect((result.matrixify_mode_rows as any)?.value).toBe('dimensions');
expect((result.matrixify_enable as any)?.value).toBe(true);
// Horizontal layout was NOT enabled — columns mode reset
expect((result.matrixify_mode_columns as any)?.value).toBe('disabled');
});
test('getControlsState preserves modes and sets matrixify_enable when old horizontal flag is true', () => {
const state = buildExploreState();
const formData = {
datasource: '1__table',
viz_type: 'test-chart',
matrixify_enable_horizontal_layout: true,
matrixify_mode_rows: 'dimensions',
matrixify_mode_columns: 'metrics',
};
const result = getControlsState(state as any, formData as any);
// Horizontal layout was enabled — columns mode preserved, matrixify_enable migrated
expect((result.matrixify_mode_columns as any)?.value).toBe('metrics');
expect((result.matrixify_enable as any)?.value).toBe(true);
// Vertical layout was NOT enabled — rows mode reset
expect((result.matrixify_mode_rows as any)?.value).toBe('disabled');
});
test('getControlsState preserves both modes when both old per-axis flags are true', () => {
const state = buildExploreState();
const formData = {
datasource: '1__table',
viz_type: 'test-chart',
matrixify_enable_vertical_layout: true,
matrixify_enable_horizontal_layout: true,
matrixify_mode_rows: 'dimensions',
matrixify_mode_columns: 'metrics',
};
const result = getControlsState(state as any, formData as any);
expect((result.matrixify_mode_rows as any)?.value).toBe('dimensions');
expect((result.matrixify_mode_columns as any)?.value).toBe('metrics');
expect((result.matrixify_enable as any)?.value).toBe(true);
});
test('getControlsState resets modes when old per-axis flags are explicitly false', () => {
const state = buildExploreState();
const formData = {
datasource: '1__table',
viz_type: 'test-chart',
matrixify_enable_vertical_layout: false,
matrixify_enable_horizontal_layout: false,
matrixify_mode_rows: 'dimensions',
matrixify_mode_columns: 'metrics',
};
const result = getControlsState(state as any, formData as any);
// Old flags present but false — chart never used matrixify, reset stale modes
expect((result.matrixify_mode_rows as any)?.value).toBe('disabled');
expect((result.matrixify_mode_columns as any)?.value).toBe('disabled');
});
// ============================================================
// P2: Dashboard hydration (applyDefaultFormData) with old per-axis flags
// ============================================================
test('applyDefaultFormData preserves modes when old vertical flag is true', () => {
const formData = {
datasource: '1__table',
viz_type: 'test-chart',
matrixify_enable_vertical_layout: true,
matrixify_mode_rows: 'dimensions',
matrixify_mode_columns: 'metrics',
};
const outputFormData = applyDefaultFormData(formData as any);
expect(outputFormData.matrixify_mode_rows).toBe('dimensions');
expect(outputFormData.matrixify_enable).toBe(true);
// Horizontal not enabled — columns reset
expect(outputFormData.matrixify_mode_columns).toBe('disabled');
});
test('applyDefaultFormData preserves modes when both old flags are true', () => {
const formData = {
datasource: '1__table',
viz_type: 'test-chart',
matrixify_enable_vertical_layout: true,
matrixify_enable_horizontal_layout: true,
matrixify_mode_rows: 'dimensions',
matrixify_mode_columns: 'metrics',
};
const outputFormData = applyDefaultFormData(formData as any);
expect(outputFormData.matrixify_mode_rows).toBe('dimensions');
expect(outputFormData.matrixify_mode_columns).toBe('metrics');
expect(outputFormData.matrixify_enable).toBe(true);
});
// ============================================================
// Direct handleDeprecatedControls tests: verify form_data mutation
// so callers (hydrateExplore) can propagate migrated fields into state
// ============================================================
test('handleDeprecatedControls sets matrixify_enable on form_data when old vertical flag is true', () => {
const formData: any = {
matrixify_enable_vertical_layout: true,
matrixify_mode_rows: 'dimensions',
matrixify_mode_columns: 'metrics',
};
handleDeprecatedControls(formData);
expect(formData.matrixify_enable).toBe(true);
expect(formData.matrixify_mode_rows).toBe('dimensions');
// Horizontal not enabled — columns reset
expect(formData.matrixify_mode_columns).toBe('disabled');
});
test('handleDeprecatedControls resets modes when no matrixify_enable and no old flags', () => {
const formData: any = {
matrixify_mode_rows: 'dimensions',
matrixify_mode_columns: 'metrics',
};
handleDeprecatedControls(formData);
expect(formData.matrixify_enable).toBeUndefined();
expect(formData.matrixify_mode_rows).toBe('disabled');
expect(formData.matrixify_mode_columns).toBe('disabled');
});
test('handleDeprecatedControls is idempotent — no-op when matrixify_enable already present', () => {
const formData: any = {
matrixify_enable: true,
matrixify_mode_rows: 'dimensions',
matrixify_mode_columns: 'metrics',
};
handleDeprecatedControls(formData);
// No mutation — matrixify_enable key is present
expect(formData.matrixify_enable).toBe(true);
expect(formData.matrixify_mode_rows).toBe('dimensions');
expect(formData.matrixify_mode_columns).toBe('metrics');
});
+51 -7
View File
@@ -41,9 +41,16 @@ type FormData = QueryFormData & {
y_axis_zero?: boolean;
y_axis_bounds?: [number | null, number | null];
datasource?: string;
matrixify_enable?: boolean;
matrixify_mode_rows?: string;
matrixify_mode_columns?: string;
// Pre-revamp per-axis enable flags (removed in #38519, may still exist in
// persisted form_data for charts that actually used matrixify)
matrixify_enable_vertical_layout?: boolean;
matrixify_enable_horizontal_layout?: boolean;
};
function handleDeprecatedControls(formData: FormData): void {
export function handleDeprecatedControls(formData: FormData): void {
// Reaffectation / handling of deprecated controls
/* eslint-disable no-param-reassign */
@@ -51,6 +58,37 @@ function handleDeprecatedControls(formData: FormData): void {
if (formData.y_axis_zero) {
formData.y_axis_bounds = [0, null];
}
// #38519: migrate pre-revamp matrixify controls to the new single-toggle
// system. Before the revamp, per-axis enable flags
// (matrixify_enable_vertical_layout / matrixify_enable_horizontal_layout)
// gated visibility, and matrixify_mode_rows/columns defaulted to
// non-disabled values ('dimensions'/'metrics'). The revamp replaced those
// with a single matrixify_enable toggle and mode default 'disabled'.
//
// Charts that actually used matrixify pre-revamp have the old per-axis
// flags set to true — we must preserve their modes and set
// matrixify_enable: true. Charts that never used matrixify (or predate it)
// need stale mode defaults reset to 'disabled' because 4 downstream UI
// consumers (ExploreChartPanel, ChartContextMenu, DrillBySubmenu,
// ChartRenderer) infer "matrixify is active" from mode values alone.
if (!('matrixify_enable' in formData)) {
const hadVerticalLayout =
formData.matrixify_enable_vertical_layout === true;
const hadHorizontalLayout =
formData.matrixify_enable_horizontal_layout === true;
if (hadVerticalLayout || hadHorizontalLayout) {
// Pre-revamp chart that genuinely used matrixify — migrate to new flag
formData.matrixify_enable = true;
if (!hadVerticalLayout) formData.matrixify_mode_rows = 'disabled';
if (!hadHorizontalLayout) formData.matrixify_mode_columns = 'disabled';
} else {
// Never used matrixify — reset stale defaults
formData.matrixify_mode_rows = 'disabled';
formData.matrixify_mode_columns = 'disabled';
}
}
}
export function getControlsState(
@@ -89,25 +127,31 @@ export function getControlsState(
export function applyDefaultFormData(
inputFormData: FormData,
): Record<string, unknown> {
const datasourceType = inputFormData.datasource?.split('__')[1] ?? '';
const vizType = inputFormData.viz_type;
// Normalize deprecated controls before building control state — ensures
// stale matrixify modes are cleaned on the dashboard hydration path too,
// not just the explore path (getControlsState).
const cleanedFormData = { ...inputFormData };
handleDeprecatedControls(cleanedFormData);
const datasourceType = cleanedFormData.datasource?.split('__')[1] ?? '';
const vizType = cleanedFormData.viz_type;
const controlsState = getAllControlsState(
vizType,
datasourceType as DatasourceType,
null,
inputFormData,
cleanedFormData,
);
// eslint-disable-next-line @typescript-eslint/no-explicit-any
const controlFormData = getFormDataFromControls(controlsState as any);
const formData: Record<string, unknown> = {};
Object.keys(controlsState)
.concat(Object.keys(inputFormData))
.concat(Object.keys(cleanedFormData))
.forEach(controlName => {
if (inputFormData[controlName as keyof FormData] === undefined) {
if (cleanedFormData[controlName as keyof FormData] === undefined) {
formData[controlName] = controlFormData[controlName];
} else {
formData[controlName] = inputFormData[controlName as keyof FormData];
formData[controlName] = cleanedFormData[controlName as keyof FormData];
}
});
@@ -34,7 +34,6 @@ import {
Tooltip,
Row,
type OnClickHandler,
type ButtonProps as CoreButtonProps,
} from '@superset-ui/core/components';
import { Icons } from '@superset-ui/core/components/Icons';
import { MenuObjectProps } from 'src/types/bootstrapTypes';
@@ -149,7 +148,7 @@ export interface ButtonProps {
'data-test'?: string;
buttonStyle: 'primary' | 'secondary' | 'dashed' | 'link' | 'tertiary';
loading?: boolean;
icon?: CoreButtonProps['icon'];
icon?: ReactNode;
component?: ReactNode;
}
@@ -307,6 +307,20 @@ export class ThemeController {
return this.currentMode;
}
/**
* Returns the resolved theme mode as 'dark' or 'light'.
* Takes into account SYSTEM mode and returns the actual resolved preference.
*/
public getCurrentModeResolved(): 'dark' | 'light' {
const activeTheme = this.getThemeForMode(this.currentMode);
if (activeTheme) {
const normalizedTheme = this.normalizeTheme(activeTheme);
return isThemeConfigDark(normalizedTheme) ? 'dark' : 'light';
}
return this.currentMode === ThemeMode.DARK ? 'dark' : 'light';
}
/**
* Sets new theme.
* @param theme - The new theme to apply
+11 -3
View File
@@ -53,11 +53,19 @@ export function SupersetThemeProvider({
);
useEffect(() => {
const unsubscribe = themeController.onChange(theme => {
// TODO: Once we migrate to react>=18 is should be possible
// to replace the useState and useEffect with a singular
// useSyncExternalStore, simplifying quite a bit
const updateState = (theme: Theme) => {
setCurrentTheme(theme);
setCurrentThemeMode(themeController.getCurrentMode());
});
document.documentElement.setAttribute(
'data-theme-mode',
themeController.getCurrentModeResolved(),
);
};
const unsubscribe = themeController.onChange(updateState);
updateState(themeController.getTheme());
return unsubscribe;
}, [themeController]);
@@ -1798,3 +1798,38 @@ test('ThemeController invalid initialMode falls back to SYSTEM', () => {
// falling through to the default SYSTEM mode
expect(controller.getCurrentMode()).toBe(ThemeMode.SYSTEM);
});
test('getCurrentModeResolved returns light for light theme', () => {
mockGetBootstrapData.mockReturnValue(
createMockBootstrapData({
default: { token: { colorBgBase: '#ffffff' } },
dark: {
token: { colorBgBase: '#000000' },
algorithm: ThemeAlgorithm.DARK,
},
}),
);
const controller = createController();
expect(controller.getCurrentModeResolved()).toBe('light');
controller.setThemeMode(ThemeMode.DARK);
expect(controller.getCurrentModeResolved()).toBe('dark');
});
test('getResolvedThemeMode returns dark when default theme is dark but mode is DEFAULT', () => {
// Setup: default theme is dark (has dark algorithm)
// This simulates single-theme deployments where THEME_DARK=None but default is dark
mockGetBootstrapData.mockReturnValue(
createMockBootstrapData({
default: {
token: { colorBgBase: '#000000' }, // dark background
algorithm: antdThemeImport.darkAlgorithm,
},
dark: {}, // empty - no separate dark theme
}),
);
const controller = createController();
expect(controller.getCurrentMode()).toBe(ThemeMode.DEFAULT);
expect(controller.getCurrentModeResolved()).toBe('dark');
});
@@ -75,6 +75,7 @@ describe('SupersetThemeProvider', () => {
mockThemeController = {
getTheme: jest.fn().mockReturnValue(mockTheme),
getCurrentMode: jest.fn().mockReturnValue(ThemeMode.DEFAULT),
getCurrentModeResolved: jest.fn().mockReturnValue('dark'),
setTheme: jest.fn(),
setThemeMode: jest.fn(),
resetTheme: jest.fn(),
@@ -276,4 +277,59 @@ describe('SupersetThemeProvider', () => {
);
});
});
afterEach(() => {
document.documentElement.removeAttribute('data-theme-mode');
});
test('should set data-theme-mode="light" on mount when resolved mode is light', () => {
mockThemeController.getCurrentModeResolved.mockReturnValue('light');
render(
<SupersetThemeProvider themeController={mockThemeController}>
<div>Content</div>
</SupersetThemeProvider>,
);
expect(document.documentElement.getAttribute('data-theme-mode')).toBe(
'light',
);
});
test('should set data-theme-mode="dark" on mount when resolved mode is dark', () => {
mockThemeController.getCurrentModeResolved.mockReturnValue('dark');
render(
<SupersetThemeProvider themeController={mockThemeController}>
<div>Content</div>
</SupersetThemeProvider>,
);
expect(document.documentElement.getAttribute('data-theme-mode')).toBe(
'dark',
);
});
test('should update data-theme-mode when theme changes', () => {
mockThemeController.getCurrentModeResolved.mockReturnValue('light');
render(
<SupersetThemeProvider themeController={mockThemeController}>
<div>Content</div>
</SupersetThemeProvider>,
);
expect(document.documentElement.getAttribute('data-theme-mode')).toBe(
'light',
);
act(() => {
mockThemeController.getCurrentModeResolved.mockReturnValue('dark');
mockOnChangeCallback(mockDarkTheme);
});
expect(document.documentElement.getAttribute('data-theme-mode')).toBe(
'dark',
);
});
});
@@ -18,7 +18,6 @@
*/
import { getExtensionsRegistry } from '@superset-ui/core';
import type { ComponentType, ReactNode } from 'react';
import { Provider as ReduxProvider } from 'react-redux';
import { QueryParamProvider } from 'use-query-params';
import { ReactRouter5Adapter } from 'use-query-params/adapters/react-router-5';
@@ -40,7 +39,7 @@ export const RootContextProviders: React.FC<{ children?: React.ReactNode }> = ({
}) => {
const RootContextProviderExtension = extensionsRegistry.get(
'root.context.provider',
) as ComponentType<{ children?: ReactNode }> | undefined;
);
return (
<SupersetThemeProvider themeController={themeController}>
+45 -14
View File
@@ -15,6 +15,7 @@
# specific language governing permissions and limitations
# under the License.
import logging
from collections.abc import Sequence
from datetime import datetime, timedelta
from typing import Any, Optional, Union
from uuid import UUID
@@ -196,7 +197,7 @@ class BaseReportState:
db.session.commit() # pylint: disable=consider-using-transaction
except StaleDataError as ex:
# Report schedule was modified or deleted by another process
db.session.rollback()
db.session.rollback() # pylint: disable=consider-using-transaction
logger.warning(
"Report schedule (execution %s) was modified or deleted "
"during execution. This can occur when a report is deleted "
@@ -280,6 +281,7 @@ class BaseReportState:
)
urls = self._get_tabs_urls(
anchor_list,
dashboard_state=dashboard_state,
native_filter_params=native_filter_params,
user_friendly=user_friendly,
)
@@ -291,12 +293,9 @@ class BaseReportState:
# overwriting — dashboard_state may already have urlParams
# (e.g. standalone=true) that must be preserved.
state: DashboardPermalinkState = {**dashboard_state}
existing_params: list[tuple[str, str]] = state.get("urlParams") or []
merged_params: list[list[str]] = [
list(p) for p in existing_params if p[0] != "native_filters"
]
merged_params.append(["native_filters", native_filter_params or ""])
state["urlParams"] = merged_params # type: ignore[typeddict-item]
state["urlParams"] = self._merge_native_filters_into_url_params(
state.get("urlParams"), native_filter_params
)
return [
self._get_tab_url(
state,
@@ -310,12 +309,17 @@ class BaseReportState:
if filter_warnings:
self._filter_warnings.extend(filter_warnings)
if native_filter_params and native_filter_params != "()":
# Preserve any urlParams from extra.dashboard (e.g. standalone=true)
# set via API even when ALERT_REPORT_TABS is off — same merge
# semantics as the protected branch above.
fallback_state = self._report_schedule.extra.get("dashboard") or {}
return [
self._get_tab_url(
{
"urlParams": [
["native_filters", native_filter_params] # type: ignore
],
"urlParams": self._merge_native_filters_into_url_params(
fallback_state.get("urlParams"),
native_filter_params,
)
},
user_friendly=user_friendly,
)
@@ -353,24 +357,51 @@ class BaseReportState:
user_friendly=user_friendly,
)
@staticmethod
def _merge_native_filters_into_url_params(
existing: Optional[Sequence[Sequence[str]]],
native_filter_params: Optional[str],
) -> list[Sequence[str]]:
"""
Merge the report's ``native_filters`` into a permalink's existing
``urlParams``, deduping any prior ``native_filters`` entry so the
report's value wins. All other params (e.g. ``standalone=true``)
survive in their original order.
"""
merged: list[Sequence[str]] = [
list(p) for p in (existing or []) if p[0] != "native_filters"
]
merged.append(["native_filters", native_filter_params or ""])
return merged
def _get_tabs_urls(
self,
tab_anchors: list[str],
dashboard_state: Optional[DashboardPermalinkState] = None,
native_filter_params: Optional[str] = None,
user_friendly: bool = False,
) -> list[str]:
"""
Get multple tabs urls
Get multiple tabs urls.
Each per-tab permalink merges the report's ``native_filters`` into
the original ``dashboard_state.urlParams`` (deduping any prior
``native_filters`` entry), so params like ``standalone=true`` are
preserved — matching the precedence rules of the single-tab branch
in :meth:`get_dashboard_urls`.
"""
base_state: DashboardPermalinkState = dashboard_state or {}
merged_params = self._merge_native_filters_into_url_params(
base_state.get("urlParams"), native_filter_params
)
return [
self._get_tab_url(
{
"anchor": tab_anchor,
"dataMask": None,
"activeTabs": None,
"urlParams": [
["native_filters", native_filter_params] # type: ignore
],
"urlParams": merged_params,
},
user_friendly=user_friendly,
)
+10 -3
View File
@@ -360,11 +360,18 @@ class QueryObject: # pylint: disable=too-many-instance-attributes
engine = database.db_engine_spec.engine
if needs_transpilation:
clause = transpile_to_dialect(clause, engine)
# source_engine=engine ensures idempotency: this
# method can run more than once (validate() is called
# from both raise_for_access and get_df_payload), so
# the second pass must be able to re-parse the
# dialect-specific output (e.g. BigQuery backticks)
# produced by the first pass.
clause = transpile_to_dialect(
clause, engine, source_engine=engine, identify=True
)
sanitized_clause = sanitize_clause(clause, engine)
if sanitized_clause != clause:
self.extras[param] = sanitized_clause
self.extras[param] = sanitized_clause
except QueryClauseValidationException as ex:
raise QueryObjectValidationError(ex.message) from ex
+4 -14
View File
@@ -55,11 +55,7 @@ from superset.constants import CHANGE_ME_SECRET_KEY
from superset.jinja_context import BaseTemplateProcessor
from superset.key_value.types import JsonKeyValueCodec
from superset.stats_logger import DummyStatsLogger
from superset.superset_typing import (
CacheConfig,
DBConnectionMutator,
EngineContextManager,
)
from superset.superset_typing import CacheConfig
from superset.tasks.types import ExecutorType
from superset.themes.types import Theme
from superset.utils import core as utils
@@ -835,6 +831,7 @@ DEFAULT_FEATURE_FLAGS: dict[str, bool] = {
# FIREWALL (only port 22 is open)
# ----------------------------------------------------------------------
SSH_TUNNEL_MANAGER_CLASS = "superset.extensions.ssh.SSHManager"
SSH_TUNNEL_LOCAL_BIND_ADDRESS = "127.0.0.1"
#: Timeout (seconds) for tunnel connection (open_channel timeout)
SSH_TUNNEL_TIMEOUT_SEC = 10.0
@@ -1723,14 +1720,7 @@ def engine_context_manager( # pylint: disable=unused-argument
yield None
ENGINE_CONTEXT_MANAGER: EngineContextManager = engine_context_manager
# The class used to manage SQLAlchemy engine creation, including SSH tunnels
# and connection details. Deployments that need custom behavior (e.g. bastion
# routing, audit logging, host-key policy, custom credential handling) can
# subclass `superset.engines.manager.EngineManager` and point this setting at
# the subclass.
ENGINE_MANAGER_CLASS = "superset.engines.manager.EngineManager"
ENGINE_CONTEXT_MANAGER = engine_context_manager
# A callable that allows altering the database connection URL and params
# on the fly, at runtime. This allows for things like impersonation or
@@ -1747,7 +1737,7 @@ ENGINE_MANAGER_CLASS = "superset.engines.manager.EngineManager"
#
# Note that the returned uri and params are passed directly to sqlalchemy's
# as such `create_engine(url, **params)`
DB_CONNECTION_MUTATOR: DBConnectionMutator | None = None
DB_CONNECTION_MUTATOR = None
# A callable that is invoked for every invocation of DB Engine Specs
+5 -1
View File
@@ -14,6 +14,7 @@
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.
from collections.abc import Sequence
from typing import Any, Optional, TypedDict
@@ -21,7 +22,10 @@ class DashboardPermalinkState(TypedDict, total=False):
dataMask: Optional[dict[str, Any]]
activeTabs: Optional[list[str]]
anchor: Optional[str]
urlParams: Optional[list[tuple[str, str]]]
# urlParams items are stored/transmitted as JSON arrays, so they
# arrive at runtime as ``list[str]``; ``Sequence[str]`` keeps the
# annotation permissive of both list and tuple shapes.
urlParams: Optional[list[Sequence[str]]]
chartStates: Optional[dict[str, Any]]
-203
View File
@@ -1,203 +0,0 @@
# Licensed to the Apache Software Foundation (ASF) under one
# or more contributor license agreements. See the NOTICE file
# distributed with this work for additional information
# regarding copyright ownership. The ASF licenses this file
# to you under the Apache License, Version 2.0 (the
# "License"); you may not use this file except in compliance
# with the License. You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing,
# software distributed under the License is distributed on an
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.
from contextlib import contextmanager
from datetime import timedelta
from io import StringIO
from typing import Any, Iterator, TYPE_CHECKING
import sshtunnel
from paramiko import RSAKey
from sqlalchemy import create_engine, pool
from sqlalchemy.engine import Engine
from sqlalchemy.engine.url import URL
from sshtunnel import SSHTunnelForwarder
from superset.commands.database.ssh_tunnel.exceptions import SSHTunnelDatabasePortError
from superset.databases.utils import make_url_safe
from superset.superset_typing import DBConnectionMutator, EngineContextManager
from superset.utils.core import get_query_source_from_request, get_user_id, QuerySource
if TYPE_CHECKING:
from superset.databases.ssh_tunnel.models import SSHTunnel
from superset.models.core import Database
class EngineManager:
"""Centralized SQLAlchemy engine creation for Superset."""
def __init__(
self,
engine_context_manager: EngineContextManager,
db_connection_mutator: DBConnectionMutator | None = None,
local_bind_address: str = "127.0.0.1",
tunnel_timeout: timedelta = timedelta(seconds=30),
ssh_timeout: timedelta = timedelta(seconds=1),
) -> None:
self.engine_context_manager = engine_context_manager
self.db_connection_mutator = db_connection_mutator
self.local_bind_address = local_bind_address
sshtunnel.TUNNEL_TIMEOUT = tunnel_timeout.total_seconds()
sshtunnel.SSH_TIMEOUT = ssh_timeout.total_seconds()
@contextmanager
def get_engine(
self,
database: "Database",
catalog: str | None,
schema: str | None,
source: QuerySource | None,
) -> Iterator[Engine]:
"""Context manager to get a SQLAlchemy engine."""
from superset.utils.oauth2 import check_for_oauth2
with self.engine_context_manager(database, catalog, schema):
with check_for_oauth2(database):
uri, kwargs = self._get_engine_args(
database,
catalog,
schema,
source,
get_user_id(),
)
if database.ssh_tunnel:
tunnel = self._create_tunnel(database.ssh_tunnel, uri)
try:
uri = uri.set(
host=tunnel.local_bind_address[0],
port=tunnel.local_bind_port,
)
yield self._create_engine(database, uri, kwargs)
finally:
tunnel.stop()
else:
yield self._create_engine(database, uri, kwargs)
def _get_engine_args(
self,
database: "Database",
catalog: str | None,
schema: str | None,
source: QuerySource | None,
user_id: int | None,
) -> tuple[URL, dict[str, Any]]:
"""Build SQLAlchemy URI and kwargs before engine creation."""
from superset import is_feature_enabled
from superset.extensions import security_manager
uri = make_url_safe(database.sqlalchemy_uri_decrypted)
extra = database.get_extra(source)
kwargs = dict(extra.get("engine_params", {}))
kwargs["poolclass"] = pool.NullPool
connect_args = kwargs.setdefault("connect_args", {})
uri, connect_args = database.db_engine_spec.adjust_engine_params(
uri,
connect_args,
catalog,
schema,
)
username = database.get_effective_user(uri)
if username and is_feature_enabled("IMPERSONATE_WITH_EMAIL_PREFIX"):
user = security_manager.find_user(username=username)
if user and user.email and "@" in user.email:
username = user.email.split("@")[0]
if database.impersonate_user:
oauth2_config = database.get_oauth2_config()
from superset.utils.oauth2 import get_oauth2_access_token
access_token = (
get_oauth2_access_token(
oauth2_config,
database.id,
user_id,
database.db_engine_spec,
)
if oauth2_config and user_id
else None
)
uri, kwargs = database.db_engine_spec.impersonate_user(
database,
username,
access_token,
uri,
kwargs,
)
database.update_params_from_encrypted_extra(kwargs)
if self.db_connection_mutator:
source = source or get_query_source_from_request()
uri, kwargs = self.db_connection_mutator(
uri,
kwargs,
username,
security_manager,
source,
)
database.db_engine_spec.validate_database_uri(uri)
return uri, kwargs
def _create_engine(
self,
database: "Database",
uri: URL,
kwargs: dict[str, Any],
) -> Engine:
try:
return create_engine(uri, **kwargs)
except Exception as ex:
raise database.db_engine_spec.get_dbapi_mapped_exception(ex) from ex
def _create_tunnel(self, ssh_tunnel: "SSHTunnel", uri: URL) -> SSHTunnelForwarder:
kwargs = self._get_tunnel_kwargs(ssh_tunnel, uri)
tunnel = sshtunnel.open_tunnel(**kwargs)
tunnel.start()
return tunnel
def _get_tunnel_kwargs(self, ssh_tunnel: "SSHTunnel", uri: URL) -> dict[str, Any]:
from superset.utils.ssh_tunnel import get_default_port
backend = uri.get_backend_name()
port = uri.port or get_default_port(backend)
if not port:
raise SSHTunnelDatabasePortError()
kwargs = {
"ssh_address_or_host": (ssh_tunnel.server_address, ssh_tunnel.server_port),
"ssh_username": ssh_tunnel.username,
"remote_bind_address": (uri.host, port),
"local_bind_address": (self.local_bind_address,),
}
if ssh_tunnel.password:
kwargs["ssh_password"] = ssh_tunnel.password
elif ssh_tunnel.private_key:
private_key_file = StringIO(ssh_tunnel.private_key)
private_key = RSAKey.from_private_key(
private_key_file,
ssh_tunnel.private_key_password,
)
kwargs["ssh_pkey"] = private_key
return kwargs
+2 -2
View File
@@ -42,7 +42,7 @@ from werkzeug.local import LocalProxy
from superset.async_events.async_query_manager import AsyncQueryManager
from superset.async_events.async_query_manager_factory import AsyncQueryManagerFactory
from superset.extensions.engine_manager import EngineManagerExtension
from superset.extensions.ssh import SSHManagerFactory
from superset.extensions.stats_logger import BaseStatsLoggerManager
from superset.security.manager import SupersetSecurityManager
from superset.utils.cache_manager import CacheManager
@@ -146,7 +146,6 @@ cache_manager = CacheManager()
celery_app = celery.Celery()
csrf = CSRFProtect()
db = get_sqla_class()()
engine_manager_extension = EngineManagerExtension()
_event_logger: dict[str, Any] = {}
encrypted_field_factory = EncryptedFieldFactory()
event_logger = LocalProxy(lambda: _event_logger.get("event_logger"))
@@ -157,5 +156,6 @@ migrate = Migrate()
profiling = ProfilingExtension()
results_backend_manager = ResultsBackendManager()
security_manager: SupersetSecurityManager = LocalProxy(lambda: appbuilder.sm)
ssh_manager_factory = SSHManagerFactory()
stats_logger_manager = BaseStatsLoggerManager()
talisman = Talisman()
-77
View File
@@ -1,77 +0,0 @@
# Licensed to the Apache Software Foundation (ASF) under one
# or more contributor license agreements. See the NOTICE file
# distributed with this work for additional information
# regarding copyright ownership. The ASF licenses this file
# to you under the Apache License, Version 2.0 (the
# "License"); you may not use this file except in compliance
# with the License. You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing,
# software distributed under the License is distributed on an
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.
import logging
from datetime import timedelta
from flask import Flask
from superset.engines.manager import EngineManager
from superset.utils.class_utils import load_class_from_name
logger = logging.getLogger(__name__)
class EngineManagerExtension:
"""
Flask extension for managing SQLAlchemy engines in Superset.
"""
def __init__(self) -> None:
self.engine_manager: EngineManager | None = None
def init_app(self, app: Flask) -> None:
"""
Initialize the EngineManager with Flask app configuration.
"""
engine_context_manager = app.config["ENGINE_CONTEXT_MANAGER"]
db_connection_mutator = app.config["DB_CONNECTION_MUTATOR"]
local_bind_address = app.config["SSH_TUNNEL_LOCAL_BIND_ADDRESS"]
tunnel_timeout = timedelta(seconds=app.config["SSH_TUNNEL_TIMEOUT_SEC"])
ssh_timeout = timedelta(seconds=app.config["SSH_TUNNEL_PACKET_TIMEOUT_SEC"])
engine_manager_class: type[EngineManager] = load_class_from_name(
app.config["ENGINE_MANAGER_CLASS"]
)
self.engine_manager = engine_manager_class(
engine_context_manager,
db_connection_mutator,
local_bind_address,
tunnel_timeout,
ssh_timeout,
)
logger.info(
"Initialized EngineManager with tunnel_timeout=%s, ssh_timeout=%s",
tunnel_timeout.total_seconds(),
ssh_timeout.total_seconds(),
)
@property
def manager(self) -> EngineManager:
"""
Get the EngineManager instance.
Raises:
RuntimeError: If the extension hasn't been initialized with an app.
"""
if self.engine_manager is None:
raise RuntimeError(
"EngineManager extension not initialized. "
"Call init_app() with a Flask app first."
)
return self.engine_manager
+94
View File
@@ -0,0 +1,94 @@
# Licensed to the Apache Software Foundation (ASF) under one
# or more contributor license agreements. See the NOTICE file
# distributed with this work for additional information
# regarding copyright ownership. The ASF licenses this file
# to you under the Apache License, Version 2.0 (the
# "License"); you may not use this file except in compliance
# with the License. You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing,
# software distributed under the License is distributed on an
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.
import logging
from io import StringIO
from typing import TYPE_CHECKING
import sshtunnel
from flask import Flask
from paramiko import RSAKey
from superset.commands.database.ssh_tunnel.exceptions import SSHTunnelDatabasePortError
from superset.databases.utils import make_url_safe
from superset.utils.class_utils import load_class_from_name
if TYPE_CHECKING:
from superset.databases.ssh_tunnel.models import SSHTunnel
class SSHManager:
def __init__(self, app: Flask) -> None:
super().__init__()
self.local_bind_address = app.config["SSH_TUNNEL_LOCAL_BIND_ADDRESS"]
sshtunnel.TUNNEL_TIMEOUT = app.config["SSH_TUNNEL_TIMEOUT_SEC"]
sshtunnel.SSH_TIMEOUT = app.config["SSH_TUNNEL_PACKET_TIMEOUT_SEC"]
def build_sqla_url(
self, sqlalchemy_url: str, server: sshtunnel.SSHTunnelForwarder
) -> str:
# override any ssh tunnel configuration object
url = make_url_safe(sqlalchemy_url)
return url.set(
host=server.local_bind_address[0],
port=server.local_bind_port,
)
def create_tunnel(
self,
ssh_tunnel: "SSHTunnel",
sqlalchemy_database_uri: str,
) -> sshtunnel.SSHTunnelForwarder:
from superset.utils.ssh_tunnel import get_default_port
url = make_url_safe(sqlalchemy_database_uri)
backend = url.get_backend_name()
port = url.port or get_default_port(backend)
if not port:
raise SSHTunnelDatabasePortError()
params = {
"ssh_address_or_host": (ssh_tunnel.server_address, ssh_tunnel.server_port),
"ssh_username": ssh_tunnel.username,
"remote_bind_address": (url.host, port),
"local_bind_address": (self.local_bind_address,),
"debug_level": logging.getLogger("flask_appbuilder").level,
}
if ssh_tunnel.password:
params["ssh_password"] = ssh_tunnel.password
elif ssh_tunnel.private_key:
private_key_file = StringIO(ssh_tunnel.private_key)
private_key = RSAKey.from_private_key(
private_key_file, ssh_tunnel.private_key_password
)
params["ssh_pkey"] = private_key
return sshtunnel.open_tunnel(**params)
class SSHManagerFactory:
def __init__(self) -> None:
self._ssh_manager = None
def init_app(self, app: Flask) -> None:
self._ssh_manager = load_class_from_name(
app.config["SSH_TUNNEL_MANAGER_CLASS"]
)(app)
@property
def instance(self) -> SSHManager:
return self._ssh_manager # type: ignore
+4 -4
View File
@@ -49,13 +49,13 @@ from superset.extensions import (
csrf,
db,
encrypted_field_factory,
engine_manager_extension,
feature_flag_manager,
machine_auth_provider_factory,
manifest_processor,
migrate,
profiling,
results_backend_manager,
ssh_manager_factory,
stats_logger_manager,
talisman,
)
@@ -616,8 +616,8 @@ class SupersetAppInitializer: # pylint: disable=too-many-public-methods
self.configure_url_map_converters()
self.configure_data_sources()
self.configure_auth_provider()
self.configure_engine_manager()
self.configure_async_queries()
self.configure_ssh_manager()
self.configure_stats_manager()
self.configure_task_manager()
@@ -793,8 +793,8 @@ class SupersetAppInitializer: # pylint: disable=too-many-public-methods
def configure_auth_provider(self) -> None:
machine_auth_provider_factory.init_app(self.superset_app)
def configure_engine_manager(self) -> None:
engine_manager_extension.init_app(self.superset_app)
def configure_ssh_manager(self) -> None:
ssh_manager_factory.init_app(self.superset_app)
def configure_stats_manager(self) -> None:
stats_logger_manager.init_app(self.superset_app)
+14 -5
View File
@@ -393,17 +393,26 @@ Used by: `get_chart_info`, `get_chart_preview`, `get_chart_data`, `generate_char
### 11. Compile Check for Chart Creation
When creating or saving charts, run a compile check to verify the query executes:
When creating, saving, or previewing charts, run schema validation (Tier 1)
and optionally a compile check (Tier 2) before persisting or caching.
``validate_and_compile`` glues both together; tools with tight SLAs
(``generate_explore_link``, ``update_chart_preview``) opt out of Tier 2.
```python
from superset.mcp_service.chart.tool.generate_chart import _compile_chart
from superset.mcp_service.chart.compile import validate_and_compile
compile_result = _compile_chart(form_data, dataset.id)
if not compile_result.success:
# Delete broken chart, return error
result = validate_and_compile(
config, form_data, dataset, run_compile_check=True
)
if not result.success:
# ``result.error_obj`` is a ``ChartGenerationError`` with fuzzy-match
# suggestions ("did you mean sum_boys?") so the LLM can self-correct.
...
```
The lower-level ``_compile_chart(form_data, dataset_id)`` is still exported
for callers that have already done their own schema validation.
### 12. Flexible Input Parsing
`ModelListCore` handles JSON string vs. native object parsing automatically via utilities in `superset.mcp_service.utils.schema_utils`:
+362
View File
@@ -0,0 +1,362 @@
# Licensed to the Apache Software Foundation (ASF) under one
# or more contributor license agreements. See the NOTICE file
# distributed with this work for additional information
# regarding copyright ownership. The ASF licenses this file
# to you under the Apache License, Version 2.0 (the
# "License"); you may not use this file except in compliance
# with the License. You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing,
# software distributed under the License is distributed on an
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.
"""
Shared compile/validation helpers for MCP chart-generating tools.
Two tiers are exposed:
* **Tier 1 schema validation** (``DatasetValidator.validate_against_dataset``):
cheap, no SQL execution, catches references to columns or metrics that do
not exist in the dataset and returns fuzzy-match suggestions.
* **Tier 2 compile check** (``_compile_chart``): runs a small (``row_limit=2``)
``ChartDataCommand`` against the underlying database to surface anything Tier
1 cannot catch (incompatible aggregates, virtual-dataset SQL bugs, etc.).
``validate_and_compile`` glues both together so each MCP tool can opt into the
tier(s) appropriate for its SLA.
"""
from __future__ import annotations
import logging
from dataclasses import dataclass, field
from typing import Any, Dict, List, Literal
from superset.commands.exceptions import CommandException
from superset.mcp_service.chart.validation.dataset_validator import DatasetValidator
from superset.mcp_service.common.error_schemas import (
ChartGenerationError,
ColumnSuggestion,
DatasetContext,
)
logger = logging.getLogger(__name__)
@dataclass
class CompileResult:
"""Result of a chart validate-and-compile check.
``error_obj`` carries the structured ``ChartGenerationError`` (with
suggestions, dataset context, etc.) that callers should embed in their
response envelope so LLM clients can self-correct. ``error`` retains the
plain-string form for backwards compatibility with existing call sites.
"""
success: bool
error: str | None = None
error_code: str | None = None
tier: Literal["validation", "compile"] | None = None
error_obj: ChartGenerationError | None = None
warnings: List[str] = field(default_factory=list)
row_count: int | None = None
def build_dataset_context_from_orm(dataset: Any) -> DatasetContext | None:
"""Construct a ``DatasetContext`` from an already-fetched ORM dataset.
Mirrors :py:meth:`DatasetValidator._get_dataset_context` but skips the
``DatasetDAO.find_by_id`` round trip. Callers that have already loaded
the dataset (for permission checks, etc.) should use this instead.
"""
if dataset is None:
return None
columns: List[Dict[str, Any]] = []
for col in getattr(dataset, "columns", []) or []:
columns.append(
{
"name": col.column_name,
"type": str(col.type) if col.type else "UNKNOWN",
"is_temporal": getattr(col, "is_temporal", False),
"is_numeric": getattr(col, "is_numeric", False),
}
)
metrics: List[Dict[str, Any]] = []
for metric in getattr(dataset, "metrics", []) or []:
metrics.append(
{
"name": metric.metric_name,
"expression": metric.expression,
"description": metric.description,
}
)
database = getattr(dataset, "database", None)
# ``DatasetContext.database_name`` is typed as required ``str``; default to
# an empty string when the relationship isn't loaded so we don't blow up
# Pydantic validation. The field is purely informational in error messages.
database_name = getattr(database, "database_name", None) or ""
return DatasetContext(
id=dataset.id,
table_name=dataset.table_name,
schema=dataset.schema,
database_name=database_name,
available_columns=columns,
available_metrics=metrics,
)
def _compile_chart(
form_data: Dict[str, Any],
dataset_id: int,
) -> CompileResult:
"""Execute the chart's query to verify it renders without errors.
Builds a ``QueryContext`` from *form_data* and runs it through
``ChartDataCommand``. A small ``row_limit`` is used so the check is
fast we only need to know the query compiles and returns data, not
fetch the full result set.
Returns a :class:`CompileResult` with ``success=True`` when the
query executes cleanly.
"""
from superset.commands.chart.data.get_data_command import ChartDataCommand
from superset.commands.chart.exceptions import (
ChartDataCacheLoadError,
ChartDataQueryFailedError,
)
from superset.common.query_context_factory import QueryContextFactory
from superset.mcp_service.chart.chart_utils import adhoc_filters_to_query_filters
from superset.mcp_service.chart.preview_utils import _build_query_columns
try:
columns = _build_query_columns(form_data)
query_filters = adhoc_filters_to_query_filters(
form_data.get("adhoc_filters", [])
)
# Big Number charts use singular "metric" instead of "metrics"
metrics = form_data.get("metrics", [])
if not metrics and form_data.get("metric"):
metrics = [form_data["metric"]]
# Big Number with trendline uses granularity_sqla as the time column
if not columns and form_data.get("granularity_sqla"):
columns = [form_data["granularity_sqla"]]
factory = QueryContextFactory()
query_context = factory.create(
datasource={"id": dataset_id, "type": "table"},
queries=[
{
"columns": columns,
"metrics": metrics,
"orderby": form_data.get("orderby", []),
"row_limit": 2,
"filters": query_filters,
"time_range": form_data.get("time_range", "No filter"),
}
],
form_data=form_data,
)
command = ChartDataCommand(query_context)
command.validate()
result = command.run()
warnings: List[str] = []
row_count = 0
for query in result.get("queries", []):
if query.get("error"):
error_str = str(query["error"])
return CompileResult(
success=False,
error=error_str,
error_code="CHART_COMPILE_FAILED",
tier="compile",
error_obj=_build_compile_error(error_str),
)
row_count += len(query.get("data", []))
return CompileResult(success=True, warnings=warnings, row_count=row_count)
except (ChartDataQueryFailedError, ChartDataCacheLoadError) as exc:
return CompileResult(
success=False,
error=str(exc),
error_code="CHART_COMPILE_FAILED",
tier="compile",
error_obj=_build_compile_error(str(exc)),
)
except (CommandException, ValueError, KeyError) as exc:
return CompileResult(
success=False,
error=str(exc),
error_code="CHART_COMPILE_FAILED",
tier="compile",
error_obj=_build_compile_error(str(exc)),
)
def _adhoc_filter_column_valid(
column: str, clause: str, dataset_context: DatasetContext
) -> bool:
"""Return True if *column* is a valid reference for this filter clause.
WHERE filters must reference a physical column; HAVING filters may also
reference a saved metric because Superset resolves metric names there.
"""
if clause == "HAVING":
return DatasetValidator._column_exists(column, dataset_context)
return any(
col["name"].lower() == column.lower()
for col in dataset_context.available_columns
)
def _validate_adhoc_filter_columns(
form_data: Dict[str, Any], dataset_context: DatasetContext
) -> ChartGenerationError | None:
"""Tier-1 check for adhoc-filter column references stored in ``form_data``.
``DatasetValidator._extract_column_references`` walks the typed
``ChartConfig`` and only sees ``config.filters``. Tools like
``update_chart_preview`` and ``update_chart`` (preview path) also merge
*previously cached* ``adhoc_filters`` into ``form_data`` that aren't
represented on the new config those would otherwise bypass validation
and surface only when Explore tries to run the query.
"""
adhoc_filters = form_data.get("adhoc_filters") or []
invalid: List[str] = []
for f in adhoc_filters:
if not isinstance(f, dict):
continue
# SIMPLE filters expose the column via "subject"; SQL-expression
# filters carry a free-form ``sqlExpression`` we can't safely parse,
# so skip those.
if f.get("expressionType") and f.get("expressionType") != "SIMPLE":
continue
column = f.get("subject") or f.get("col")
if not column or not isinstance(column, str):
continue
clause = f.get("clause", "WHERE").upper()
if not _adhoc_filter_column_valid(column, clause, dataset_context):
invalid.append(column)
if not invalid:
return None
suggestions: List[str] = []
for column in invalid:
for suggestion in DatasetValidator._get_column_suggestions(
column, dataset_context
):
name = (
suggestion.name
if isinstance(suggestion, ColumnSuggestion)
else str(suggestion)
)
if name and name not in suggestions:
suggestions.append(name)
bad = ", ".join(sorted(set(invalid)))
return ChartGenerationError(
error_type="invalid_column",
message=(f"Filter references column(s) not in dataset: {bad}"),
details=(
"Adhoc filter columns must exist on the dataset. "
"If these filters were preserved from a previous chart preview, "
"remove them or pass an explicit ``filters`` list on the new config."
),
suggestions=suggestions,
error_code="CHART_VALIDATION_FAILED",
)
def _build_compile_error(message: str) -> ChartGenerationError:
"""Wrap a raw compile-failure string in the structured response envelope."""
return ChartGenerationError(
error_type="compile_error",
message="Chart query failed to execute. The chart was not saved.",
details=message or "",
suggestions=[
"Check that all columns exist in the dataset",
"Verify aggregate functions are compatible with column types",
"Ensure filters reference valid columns",
"Try simplifying the chart configuration",
],
error_code="CHART_COMPILE_FAILED",
)
def validate_and_compile(
config: Any,
form_data: Dict[str, Any],
dataset: Any,
*,
run_compile_check: bool = True,
) -> CompileResult:
"""Run schema validation (Tier 1) and optionally a compile check (Tier 2).
``dataset`` must be an already-fetched ORM dataset; this avoids a second
``DatasetDAO.find_by_id`` round trip inside the validator.
``run_compile_check`` lets fast-path tools (``generate_explore_link``,
``update_chart_preview``) skip the live DB query while still rejecting
obviously bad column references with fuzzy-match suggestions.
Returns a :class:`CompileResult`. On failure, ``error_obj`` carries the
structured :class:`ChartGenerationError` (with ``suggestions``) that the
caller should embed in its response envelope so LLM clients can
self-correct.
"""
if dataset is None:
return CompileResult(
success=False,
error="Dataset not provided to validate_and_compile",
error_code="DATASET_NOT_FOUND",
tier="validation",
)
dataset_context = build_dataset_context_from_orm(dataset)
is_valid, error = DatasetValidator.validate_against_dataset(
config, dataset.id, dataset_context=dataset_context
)
if not is_valid:
details = ""
if error is not None:
details = error.details or error.message
if error.error_code is None:
error.error_code = "CHART_VALIDATION_FAILED"
return CompileResult(
success=False,
error=details,
error_code="CHART_VALIDATION_FAILED",
tier="validation",
error_obj=error,
)
# Validate adhoc-filter columns living only in form_data (e.g. filters
# preserved from a previously cached preview). The typed config-level
# validator above doesn't see those.
if dataset_context is not None:
filter_error = _validate_adhoc_filter_columns(form_data, dataset_context)
if filter_error is not None:
return CompileResult(
success=False,
error=filter_error.details or filter_error.message,
error_code="CHART_VALIDATION_FAILED",
tier="validation",
error_obj=filter_error,
)
if not run_compile_check:
return CompileResult(success=True)
return _compile_chart(form_data, dataset.id)
@@ -20,8 +20,7 @@ MCP tool: generate_chart (simplified schema)
import logging
import time
from dataclasses import dataclass, field
from typing import Any, Dict, List
from typing import Any
from fastmcp import Context
from sqlalchemy.exc import SQLAlchemyError
@@ -39,6 +38,11 @@ from superset.mcp_service.chart.chart_utils import (
map_config_to_form_data,
validate_chart_dataset,
)
from superset.mcp_service.chart.compile import (
_compile_chart,
CompileResult,
validate_and_compile,
)
from superset.mcp_service.chart.schemas import (
AccessibilityMetadata,
CHART_FORM_DATA_EXCLUDED_FIELD_NAMES,
@@ -74,86 +78,7 @@ def _sanitize_generate_chart_form_data_for_llm_context(
)
@dataclass
class CompileResult:
"""Result of a chart compile check (test query execution)."""
success: bool
error: str | None = None
warnings: List[str] = field(default_factory=list)
row_count: int | None = None
def _compile_chart(
form_data: Dict[str, Any],
dataset_id: int,
) -> CompileResult:
"""Execute the chart's query to verify it renders without errors.
Builds a ``QueryContext`` from *form_data* and runs it through
``ChartDataCommand``. A small ``row_limit`` is used so the check is
fast we only need to know the query compiles and returns data, not
fetch the full result set.
Returns a :class:`CompileResult` with ``success=True`` when the
query executes cleanly.
"""
from superset.commands.chart.data.get_data_command import ChartDataCommand
from superset.commands.chart.exceptions import (
ChartDataCacheLoadError,
ChartDataQueryFailedError,
)
from superset.common.query_context_factory import QueryContextFactory
from superset.mcp_service.chart.chart_utils import adhoc_filters_to_query_filters
from superset.mcp_service.chart.preview_utils import _build_query_columns
try:
columns = _build_query_columns(form_data)
query_filters = adhoc_filters_to_query_filters(
form_data.get("adhoc_filters", [])
)
# Big Number charts use singular "metric" instead of "metrics"
metrics = form_data.get("metrics", [])
if not metrics and form_data.get("metric"):
metrics = [form_data["metric"]]
# Big Number with trendline uses granularity_sqla as the time column
if not columns and form_data.get("granularity_sqla"):
columns = [form_data["granularity_sqla"]]
factory = QueryContextFactory()
query_context = factory.create(
datasource={"id": dataset_id, "type": "table"},
queries=[
{
"columns": columns,
"metrics": metrics,
"orderby": form_data.get("orderby", []),
"row_limit": 2,
"filters": query_filters,
"time_range": form_data.get("time_range", "No filter"),
}
],
form_data=form_data,
)
command = ChartDataCommand(query_context)
command.validate()
result = command.run()
warnings: List[str] = []
row_count = 0
for query in result.get("queries", []):
if query.get("error"):
return CompileResult(success=False, error=str(query["error"]))
row_count += len(query.get("data", []))
return CompileResult(success=True, warnings=warnings, row_count=row_count)
except (ChartDataQueryFailedError, ChartDataCacheLoadError) as exc:
return CompileResult(success=False, error=str(exc))
except (CommandException, ValueError, KeyError) as exc:
return CompileResult(success=False, error=str(exc))
__all__ = ["CompileResult", "_compile_chart", "validate_and_compile", "generate_chart"]
@tool(
@@ -199,8 +199,12 @@ async def get_chart_data( # noqa: C901
if not chart:
await ctx.warning("Chart not found: identifier=%s" % (request.identifier,))
safe_id = str(request.identifier)[:200]
return ChartError(
error=f"No chart found with identifier: {request.identifier}",
error=(
f"No chart found with identifier: {safe_id}."
" Use list_charts to get valid chart IDs."
),
error_type="NotFound",
)
@@ -1192,8 +1192,22 @@ async def _get_chart_preview_internal( # noqa: C901
if not chart:
await ctx.warning("Chart not found: identifier=%s" % (request.identifier,))
safe_id = str(request.identifier)[:200]
is_form_data_key = (
isinstance(request.identifier, str)
and len(request.identifier) > 8
and not request.identifier.isdigit()
)
if is_form_data_key:
recovery = (
"If using a form_data_key, it may have expired — "
"use generate_explore_link to get a fresh key, "
"or use list_charts to find a saved chart by ID."
)
else:
recovery = "Use list_charts to get valid chart IDs."
return ChartError(
error=f"No chart found with identifier: {request.identifier}",
error=f"No chart found with identifier: {safe_id}. {recovery}",
error_type="NotFound",
)
@@ -40,6 +40,7 @@ from superset.mcp_service.chart.chart_utils import (
generate_chart_name,
map_config_to_form_data,
)
from superset.mcp_service.chart.compile import validate_and_compile
from superset.mcp_service.chart.schemas import (
AccessibilityMetadata,
GenerateChartResponse,
@@ -162,6 +163,70 @@ def _build_preview_form_data(
return merged
def _validate_update_against_dataset(
parsed_config: Any,
form_data: dict[str, Any],
chart: Any,
) -> GenerateChartResponse | None:
"""Run Tier 1 (schema) + Tier 2 (compile) validation against the chart's
dataset. Returns ``None`` on success, or a :class:`GenerateChartResponse`
error envelope on failure that callers should return as-is.
"""
from superset.daos.dataset import DatasetDAO
dataset = getattr(chart, "datasource", None)
if dataset is None and getattr(chart, "datasource_id", None) is not None:
dataset = DatasetDAO.find_by_id(chart.datasource_id)
if dataset is None:
return GenerateChartResponse.model_validate(
{
"chart": None,
"error": {
"error_type": "DatasetNotAccessible",
"message": "Chart's dataset is not accessible",
"details": (
f"Dataset {getattr(chart, 'datasource_id', None)} "
"is missing or inaccessible."
),
},
"success": False,
"schema_version": "2.0",
"api_version": "v1",
}
)
compile_result = validate_and_compile(
parsed_config, form_data, dataset, run_compile_check=True
)
if compile_result.success:
return None
logger.warning(
"update_chart validation failed for chart %s: %s",
getattr(chart, "id", None),
compile_result.error,
)
if compile_result.error_obj is not None:
error_payload = compile_result.error_obj.model_dump()
else:
error_payload = {
"error_type": "validation_error",
"message": "Chart update validation failed",
"details": compile_result.error or "",
"error_code": compile_result.error_code,
"suggestions": [],
}
return GenerateChartResponse.model_validate(
{
"chart": None,
"error": error_payload,
"success": False,
"schema_version": "2.0",
"api_version": "v1",
}
)
def _create_preview_url(
chart: Any, form_data: dict[str, Any]
) -> tuple[str, str | None, list[str]]:
@@ -272,17 +337,18 @@ async def update_chart( # noqa: C901
chart = find_chart_by_identifier(request.identifier)
if not chart:
safe_id = str(request.identifier)[:200]
not_found_msg = (
f"No chart found with identifier: {safe_id}."
" Use list_charts to get valid chart IDs."
)
return GenerateChartResponse.model_validate(
{
"chart": None,
"error": {
"error_type": "NotFound",
"message": (
f"No chart found with identifier: {request.identifier}"
),
"details": (
f"No chart found with identifier: {request.identifier}"
),
"message": not_found_msg,
"details": not_found_msg,
},
"success": False,
"schema_version": "2.0",
@@ -334,6 +400,18 @@ async def update_chart( # noqa: C901
if "params" in payload_or_error:
new_form_data = json.loads(payload_or_error["params"])
# Validate before persisting — catches bad column refs and runtime
# SQL errors so we don't commit a chart that can't be queried.
# Renames (no parsed_config) skip validation since form_data is
# untouched.
if parsed_config is not None and new_form_data is not None:
with event_logger.log_context(action="mcp.update_chart.validation"):
validation_error = _validate_update_against_dataset(
parsed_config, new_form_data, chart
)
if validation_error is not None:
return validation_error
with event_logger.log_context(action="mcp.update_chart.db_write"):
command = UpdateChartCommand(chart.id, payload_or_error)
updated_chart = command.run()
@@ -346,6 +424,15 @@ async def update_chart( # noqa: C901
if isinstance(preview_or_error, GenerateChartResponse):
return preview_or_error
# Validate before caching the form_data — same rationale as above.
if parsed_config is not None:
with event_logger.log_context(action="mcp.update_chart.validation"):
validation_error = _validate_update_against_dataset(
parsed_config, preview_or_error, chart
)
if validation_error is not None:
return validation_error
with event_logger.log_context(action="mcp.update_chart.preview_link"):
explore_url, form_data_key, warnings = _create_preview_url(
chart, preview_or_error
@@ -30,6 +30,7 @@ from superset_core.mcp.decorators import tool, ToolAnnotations
from superset.commands.exceptions import CommandException
from superset.exceptions import OAuth2Error, OAuth2RedirectError, SupersetException
from superset.extensions import event_logger
from superset.mcp_service.auth import has_dataset_access
from superset.mcp_service.chart.chart_helpers import extract_form_data_key_from_url
from superset.mcp_service.chart.chart_utils import (
analyze_chart_capabilities,
@@ -38,6 +39,7 @@ from superset.mcp_service.chart.chart_utils import (
generate_explore_link,
map_config_to_form_data,
)
from superset.mcp_service.chart.compile import validate_and_compile
from superset.mcp_service.chart.schemas import (
AccessibilityMetadata,
PerformanceMetadata,
@@ -88,7 +90,7 @@ def _get_previous_form_data(form_data_key: str) -> dict[str, Any] | None:
destructiveHint=True,
),
)
def update_chart_preview(
def update_chart_preview( # noqa: C901
request: UpdateChartPreviewRequest, ctx: Context
) -> Dict[str, Any]:
"""Update cached chart preview without saving.
@@ -133,6 +135,61 @@ def update_chart_preview(
if old_adhoc_filters:
new_form_data["adhoc_filters"] = old_adhoc_filters
# Tier-1 schema validation against the dataset (no DB roundtrip).
# Runs AFTER the filter merge so filter columns are also validated.
from superset.daos.dataset import DatasetDAO
if isinstance(request.dataset_id, int) or (
isinstance(request.dataset_id, str) and request.dataset_id.isdigit()
):
dataset = DatasetDAO.find_by_id(int(request.dataset_id))
else:
dataset = DatasetDAO.find_by_id(request.dataset_id, id_column="uuid")
if dataset is None or not has_dataset_access(dataset):
return {
"chart": None,
"error": {
"error_type": "DatasetNotAccessible",
"message": (
f"Dataset not found: {request.dataset_id}. "
"Use list_datasets to find valid dataset IDs."
),
"details": (
f"Dataset {request.dataset_id} is missing or inaccessible."
),
},
"success": False,
"schema_version": "2.0",
"api_version": "v1",
}
compile_result = validate_and_compile(
config, new_form_data, dataset, run_compile_check=False
)
if not compile_result.success:
logger.warning(
"update_chart_preview validation failed: %s",
compile_result.error,
)
if compile_result.error_obj is not None:
error_payload = compile_result.error_obj.model_dump()
else:
error_payload = {
"error_type": "validation_error",
"message": "Chart preview validation failed",
"details": compile_result.error or "",
"error_code": compile_result.error_code,
"suggestions": [],
}
return {
"chart": None,
"error": error_payload,
"success": False,
"schema_version": "2.0",
"api_version": "v1",
}
# Generate new explore link with updated form_data
explore_url = generate_explore_link(request.dataset_id, new_form_data)
@@ -25,7 +25,12 @@ import logging
from typing import Any, Dict, List, Tuple
from superset.mcp_service.chart.schemas import (
BigNumberChartConfig,
ColumnRef,
HandlebarsChartConfig,
MixedTimeseriesChartConfig,
PieChartConfig,
PivotTableChartConfig,
TableChartConfig,
XYChartConfig,
)
@@ -53,7 +58,7 @@ class DatasetValidator:
@staticmethod
def validate_against_dataset(
config: TableChartConfig | XYChartConfig,
config: Any,
dataset_id: int | str,
dataset_context: DatasetContext | None = None,
) -> Tuple[bool, ChartGenerationError | None]:
@@ -96,13 +101,16 @@ class DatasetValidator:
if column_error:
return False, column_error
# Validate aggregation compatibility
if isinstance(config, (TableChartConfig, XYChartConfig)):
aggregation_errors = DatasetValidator._validate_aggregations(
column_refs, dataset_context
)
if aggregation_errors:
return False, aggregation_errors[0]
# Validate aggregation compatibility for every config that produced
# column refs. ``_validate_aggregations`` is config-agnostic — gating
# it to Table/XY would let pie / pivot table / mixed timeseries /
# handlebars / big number slip through ``SUM(non_numeric)`` patterns
# for the fast-path tools that skip Tier 2.
aggregation_errors = DatasetValidator._validate_aggregations(
column_refs, dataset_context
)
if aggregation_errors:
return False, aggregation_errors[0]
return True, None
@@ -110,14 +118,41 @@ class DatasetValidator:
def _validate_columns_exist(
column_refs: List[ColumnRef], dataset_context: DatasetContext
) -> ChartGenerationError | None:
"""Validate that non-saved-metric column refs exist in the dataset."""
invalid_columns = []
"""Validate that non-saved-metric column refs exist in the dataset.
A ``ColumnRef`` with ``saved_metric=False`` must match an entry in
``available_columns``. Saved-metric *names* don't satisfy this check —
otherwise ``{name: "sum_boys", aggregate: "SUM"}`` (no
``saved_metric=true``) would slip through and downstream code would
emit ``SUM(sum_boys)`` as an ad-hoc SIMPLE metric, producing the
broken-SQL pattern this validator is meant to prevent.
"""
column_names_lower = {
col["name"].lower() for col in dataset_context.available_columns
}
metric_names_lower = {
metric["name"].lower() for metric in dataset_context.available_metrics
}
invalid_columns: List[ColumnRef] = []
saved_metric_typo: List[ColumnRef] = []
for col_ref in column_refs:
if col_ref.saved_metric:
continue
if not DatasetValidator._column_exists(col_ref.name, dataset_context):
name_lower = col_ref.name.lower()
if name_lower in column_names_lower:
continue
if name_lower in metric_names_lower:
# Name matches a saved metric but the ref didn't opt into
# saved-metric resolution. Surface a tailored hint so the
# caller (typically an LLM) can flip ``saved_metric=true``.
saved_metric_typo.append(col_ref)
else:
invalid_columns.append(col_ref)
if saved_metric_typo:
return DatasetValidator._build_saved_metric_hint_error(saved_metric_typo)
if not invalid_columns:
return None
@@ -132,6 +167,36 @@ class DatasetValidator:
invalid_columns, suggestions_map, dataset_context
)
@staticmethod
def _build_saved_metric_hint_error(
refs: List[ColumnRef],
) -> ChartGenerationError:
"""Error response when a non-saved-metric ref names a saved metric."""
names = [r.name for r in refs]
names_str = ", ".join(f"'{n}'" for n in names)
first = names[0]
return ChartGenerationError(
error_type="saved_metric_not_marked",
message=(
f"{names_str} matches a saved metric but the ref doesn't "
f"have saved_metric=true"
),
details=(
f"The dataset has a saved metric named {names_str}. To use "
f"it, set 'saved_metric': true on the column ref instead of "
f"providing an 'aggregate'. With the current shape, the "
f"chart would emit ad-hoc SQL like SUM({first}) — which is "
f"invalid because {first} is a metric expression, not a "
f"column."
),
suggestions=[
f'Did you mean: {{"name": "{first}", "saved_metric": true}}?',
"Use saved_metric=true to reference a saved dataset metric",
"Or pick a real column name and apply an aggregate to it",
],
error_code="SAVED_METRIC_NOT_MARKED",
)
@staticmethod
def _get_dataset_context(dataset_id: int | str) -> DatasetContext | None:
"""Get dataset context with column information."""
@@ -195,11 +260,16 @@ class DatasetValidator:
return None
@staticmethod
def _extract_column_references(
config: TableChartConfig | XYChartConfig,
) -> List[ColumnRef]:
"""Extract all column references from configuration."""
refs = []
def _extract_column_references(config: Any) -> List[ColumnRef]: # noqa: C901
"""Extract all column references from a chart configuration.
Covers every supported ``ChartConfig`` variant so fast-path tools
(``generate_explore_link``, ``update_chart_preview``) that only run
Tier-1 validation still catch bad column refs in pie / pivot table /
mixed timeseries / handlebars / big number charts not just XY and
table.
"""
refs: List[ColumnRef] = []
if isinstance(config, TableChartConfig):
refs.extend(config.columns)
@@ -209,10 +279,37 @@ class DatasetValidator:
refs.extend(config.y)
if config.group_by:
refs.extend(config.group_by)
elif isinstance(config, PieChartConfig):
refs.append(config.dimension)
refs.append(config.metric)
elif isinstance(config, PivotTableChartConfig):
refs.extend(config.rows)
if config.columns:
refs.extend(config.columns)
refs.extend(config.metrics)
elif isinstance(config, MixedTimeseriesChartConfig):
refs.append(config.x)
refs.extend(config.y)
if config.group_by:
refs.extend(config.group_by)
refs.extend(config.y_secondary)
if config.group_by_secondary:
refs.extend(config.group_by_secondary)
elif isinstance(config, HandlebarsChartConfig):
if config.columns:
refs.extend(config.columns)
if config.groupby:
refs.extend(config.groupby)
if config.metrics:
refs.extend(config.metrics)
elif isinstance(config, BigNumberChartConfig):
refs.append(config.metric)
if config.temporal_column:
refs.append(ColumnRef(name=config.temporal_column))
# Add filter columns
if hasattr(config, "filters") and config.filters:
for filter_config in config.filters:
# Filter columns (shared by every config type that defines ``filters``).
if filters := getattr(config, "filters", None):
for filter_config in filters:
refs.append(ColumnRef(name=filter_config.column))
return refs
@@ -379,20 +476,28 @@ class DatasetValidator:
# Find close matches
column_lower = column_name.lower()
candidate_lookup = [name[0].lower() for name in all_names]
close_matches = difflib.get_close_matches(
column_lower,
[name[0].lower() for name in all_names],
candidate_lookup,
n=max_suggestions,
cutoff=0.6,
)
# Build suggestions with proper case and type info
# Build suggestions with proper case and type info. ``ColumnSuggestion``
# requires ``similarity_score`` and does not have a ``data_type`` field;
# we score via difflib ratio and store the candidate kind in ``type``.
suggestions = []
for match in close_matches:
for name, col_type, data_type in all_names:
for name, col_type, _data_type in all_names:
if name.lower() == match:
score = difflib.SequenceMatcher(None, column_lower, match).ratio()
suggestions.append(
ColumnSuggestion(name=name, type=col_type, data_type=data_type)
ColumnSuggestion(
name=name,
type=col_type,
similarity_score=round(score, 3),
)
)
break
@@ -503,8 +608,12 @@ class DatasetValidator:
break
if col_info:
# Check numeric aggregates on non-numeric columns
numeric_aggs = ["SUM", "AVG", "MIN", "MAX", "STDDEV", "VAR", "MEDIAN"]
# Check numeric aggregates on non-numeric columns.
# MIN and MAX are intentionally excluded: they work on dates
# and text in most SQL engines, so restricting them here would
# produce false-positive errors. Leave those to the Tier-2
# compile check.
numeric_aggs = ["SUM", "AVG", "STDDEV", "VAR", "MEDIAN"]
if (
col_ref.aggregate in numeric_aggs
and not col_info.get("is_numeric", False)
@@ -334,7 +334,10 @@ def _find_and_authorize_dashboard(
dashboard=None,
dashboard_url=None,
position=None,
error=f"Dashboard with ID {dashboard_id} not found",
error=(
f"Dashboard with ID {dashboard_id} not found."
" Use list_dashboards to get valid dashboard IDs."
),
)
try:
@@ -392,7 +395,10 @@ def add_chart_to_existing_dashboard(
dashboard=None,
dashboard_url=None,
position=None,
error=f"Chart with ID {request.chart_id} not found",
error=(
f"Chart with ID {request.chart_id} not found."
" Use list_charts to get valid chart IDs."
),
)
# Validate dataset access for the chart.
@@ -230,7 +230,10 @@ def generate_dashboard( # noqa: C901
return GenerateDashboardResponse(
dashboard=None,
dashboard_url=None,
error=f"Charts not found: {list(missing_chart_ids)}",
error=(
f"Charts not found: {list(missing_chart_ids)}."
" Use list_charts to get valid chart IDs."
),
)
# Validate dataset access for each chart.
@@ -183,7 +183,10 @@ async def query_dataset( # noqa: C901
if dataset is None:
await ctx.error("Dataset not found: identifier=%s" % (request.dataset_id,))
return DatasetError.create(
error=f"No dataset found with identifier: {request.dataset_id}",
error=(
f"No dataset found with identifier: {request.dataset_id}."
" Use list_datasets to get valid dataset IDs."
),
error_type="NotFound",
)
@@ -22,21 +22,26 @@ This tool generates a URL to the Superset explore interface with the specified
chart configuration.
"""
import logging
from typing import Any, Dict
from fastmcp import Context
from superset_core.mcp.decorators import tool, ToolAnnotations
from superset.extensions import event_logger
from superset.mcp_service.auth import has_dataset_access
from superset.mcp_service.chart.chart_helpers import extract_form_data_key_from_url
from superset.mcp_service.chart.chart_utils import (
generate_explore_link as generate_url,
map_config_to_form_data,
)
from superset.mcp_service.chart.compile import validate_and_compile
from superset.mcp_service.chart.schemas import (
GenerateExploreLinkRequest,
)
logger = logging.getLogger(__name__)
@tool(
tags=["explore"],
@@ -131,6 +136,24 @@ async def generate_explore_link(
),
}
if not has_dataset_access(dataset):
logger.warning(
"User attempted to access dataset %s without permission",
request.dataset_id,
)
await ctx.warning(
"Dataset access denied: dataset_id=%s" % (request.dataset_id,)
)
return {
"url": "",
"form_data": {},
"form_data_key": None,
"error": (
f"Dataset not found: {request.dataset_id}. "
"Use list_datasets to find valid dataset IDs."
),
}
await ctx.report_progress(2, 4, "Converting configuration to form data")
with event_logger.log_context(action="mcp.generate_explore_link.form_data"):
# Normalize column names to match canonical dataset column names
@@ -165,6 +188,38 @@ async def generate_explore_link(
)
)
# Tier-1 schema validation against the dataset (no DB roundtrip).
# Catches references to non-existent columns/metrics with fuzzy
# suggestions so the LLM can self-correct ("did you mean sum_boys?").
with event_logger.log_context(action="mcp.generate_explore_link.validation"):
compile_result = validate_and_compile(
normalized_config,
form_data,
dataset,
run_compile_check=False,
)
if not compile_result.success:
await ctx.warning(
"Explore link validation failed: error=%s" % (compile_result.error,)
)
error_payload: Dict[str, Any]
if compile_result.error_obj is not None:
error_payload = compile_result.error_obj.model_dump()
else:
error_payload = {
"error_type": "validation_error",
"message": "Explore link validation failed",
"details": compile_result.error or "",
"error_code": compile_result.error_code,
"suggestions": [],
}
return {
"url": "",
"form_data": form_data,
"form_data_key": None,
"error": error_payload,
}
await ctx.report_progress(3, 4, "Generating explore URL")
with event_logger.log_context(
action="mcp.generate_explore_link.url_generation"
@@ -100,7 +100,10 @@ async def execute_sql(request: ExecuteSqlRequest, ctx: Context) -> ExecuteSqlRes
)
return ExecuteSqlResponse(
success=False,
error=f"Database with ID {request.database_id} not found",
error=(
f"Database with ID {request.database_id} not found."
" Use list_databases to get valid database IDs."
),
error_type=SupersetErrorType.DATABASE_NOT_FOUND_ERROR.value,
)
@@ -103,7 +103,8 @@ def open_sql_lab_with_context(
database = DatabaseDAO.find_by_id(request.database_connection_id)
if not database:
error_message = (
f"Database with ID {request.database_connection_id} not found"
f"Database with ID {request.database_connection_id} not found."
" Use list_databases to get valid database IDs."
)
return _sanitize_sql_lab_response_for_llm_context(
SqlLabResponse(
+120 -29
View File
@@ -25,22 +25,24 @@ import builtins
import logging
import textwrap
from ast import literal_eval
from contextlib import closing, contextmanager, suppress
from contextlib import closing, contextmanager, nullcontext, suppress
from copy import deepcopy
from datetime import datetime
from functools import lru_cache
from inspect import signature
from typing import Any, Callable, cast, Iterator, Optional, TYPE_CHECKING
from typing import Any, Callable, cast, Optional, TYPE_CHECKING
import numpy
import pandas as pd
import sqlalchemy as sqla
import sshtunnel
from flask import current_app as app, g, has_app_context
from flask_appbuilder import Model
from marshmallow.exceptions import ValidationError
from sqlalchemy import (
Boolean,
Column,
create_engine,
DateTime,
ForeignKey,
Integer,
@@ -55,6 +57,7 @@ from sqlalchemy.engine.url import URL
from sqlalchemy.exc import NoSuchModuleError
from sqlalchemy.ext.hybrid import hybrid_property
from sqlalchemy.orm import relationship
from sqlalchemy.pool import NullPool
from sqlalchemy.schema import UniqueConstraint
from sqlalchemy.sql import ColumnElement, expression, Select
from superset_core.common.models import Database as CoreDatabase
@@ -69,6 +72,7 @@ from superset.extensions import (
encrypted_field_factory,
event_logger,
security_manager,
ssh_manager_factory,
)
from superset.models.helpers import AuditMixinNullable, ImportExportMixin, UUIDMixin
from superset.result_set import SupersetResultSet
@@ -80,9 +84,10 @@ from superset.superset_typing import (
)
from superset.utils import cache as cache_util, core as utils, json
from superset.utils.backports import StrEnum
from superset.utils.core import get_username
from superset.utils.core import get_query_source_from_request, get_username
from superset.utils.oauth2 import (
check_for_oauth2,
get_oauth2_access_token,
OAuth2ClientConfigSchema,
)
@@ -419,46 +424,130 @@ class Database(CoreDatabase, AuditMixinNullable, ImportExportMixin): # pylint:
)
@contextmanager
def get_sqla_engine(
def get_sqla_engine( # pylint: disable=too-many-arguments
self,
catalog: str | None = None,
schema: str | None = None,
nullpool: bool = True,
source: utils.QuerySource | None = None,
nullpool: bool | None = None,
) -> Iterator[Engine]:
) -> Engine:
"""
Context manager for a SQLAlchemy engine.
This method will return a context manager for a SQLAlchemy engine. The engine
manager handles engine creation, SSH tunnels, and connection details in a
centralized place.
The ``nullpool`` argument is deprecated and ignored the engine manager
always uses ``NullPool``. It is kept temporarily for backwards compatibility
with external callers and will be removed in a future release.
This method will return a context manager for a SQLAlchemy engine. Using the
context manager (as opposed to the engine directly) is important because we need
to potentially establish SSH tunnels before the connection is created, and clean
them up once the engine is no longer used.
"""
if nullpool is not None:
import warnings
warnings.warn(
"The `nullpool` argument to `Database.get_sqla_engine` is "
"deprecated and ignored; the engine manager always uses NullPool.",
DeprecationWarning,
stacklevel=2,
sqlalchemy_uri = self.sqlalchemy_uri_decrypted
ssh_context_manager = (
ssh_manager_factory.instance.create_tunnel(
ssh_tunnel=self.ssh_tunnel,
sqlalchemy_database_uri=sqlalchemy_uri,
)
if self.ssh_tunnel
else nullcontext()
)
# Import here to avoid circular imports
from superset.extensions import engine_manager_extension
with ssh_context_manager as ssh_context:
if ssh_context:
logger.info(
"[SSH] Successfully created tunnel w/ %s tunnel_timeout + %s "
"ssh_timeout at %s",
sshtunnel.TUNNEL_TIMEOUT,
sshtunnel.SSH_TIMEOUT,
ssh_context.local_bind_address,
)
sqlalchemy_uri = ssh_manager_factory.instance.build_sqla_url(
sqlalchemy_uri,
ssh_context,
)
# Use the engine manager to get the engine
engine_manager = engine_manager_extension.manager
with engine_manager.get_engine(
database=self,
engine_context_manager = app.config["ENGINE_CONTEXT_MANAGER"]
with engine_context_manager(self, catalog, schema):
with check_for_oauth2(self):
yield self._get_sqla_engine(
catalog=catalog,
schema=schema,
nullpool=nullpool,
source=source,
sqlalchemy_uri=sqlalchemy_uri,
)
def _get_sqla_engine( # pylint: disable=too-many-locals # noqa: C901
self,
catalog: str | None = None,
schema: str | None = None,
nullpool: bool = True,
source: utils.QuerySource | None = None,
sqlalchemy_uri: str | None = None,
) -> Engine:
sqlalchemy_url = make_url_safe(
sqlalchemy_uri if sqlalchemy_uri else self.sqlalchemy_uri_decrypted
)
self.db_engine_spec.validate_database_uri(sqlalchemy_url)
extra = self.get_extra(source)
engine_kwargs = extra.get("engine_params", {})
if nullpool:
engine_kwargs["poolclass"] = NullPool
connect_args = engine_kwargs.setdefault("connect_args", {})
# modify URL/args for a specific catalog/schema
sqlalchemy_url, connect_args = self.db_engine_spec.adjust_engine_params(
uri=sqlalchemy_url,
connect_args=connect_args,
catalog=catalog,
schema=schema,
source=source,
) as engine:
yield engine
)
effective_username = self.get_effective_user(sqlalchemy_url)
if effective_username and is_feature_enabled("IMPERSONATE_WITH_EMAIL_PREFIX"):
user = security_manager.find_user(username=effective_username)
if user and user.email:
effective_username = user.email.split("@")[0]
oauth2_config = self.get_oauth2_config()
access_token = (
get_oauth2_access_token(
oauth2_config,
self.id,
g.user.id,
self.db_engine_spec,
)
if oauth2_config and hasattr(g, "user") and hasattr(g.user, "id")
else None
)
masked_url = self.get_password_masked_url(sqlalchemy_url)
logger.debug("Database._get_sqla_engine(). Masked URL: %s", str(masked_url))
if self.impersonate_user:
sqlalchemy_url, engine_kwargs = self.db_engine_spec.impersonate_user(
self,
effective_username,
access_token,
sqlalchemy_url,
engine_kwargs,
)
self.update_params_from_encrypted_extra(engine_kwargs)
if DB_CONNECTION_MUTATOR := app.config["DB_CONNECTION_MUTATOR"]: # noqa: N806
source = source or get_query_source_from_request()
sqlalchemy_url, engine_kwargs = DB_CONNECTION_MUTATOR(
sqlalchemy_url,
engine_kwargs,
effective_username,
security_manager,
source,
)
try:
return create_engine(sqlalchemy_url, **engine_kwargs)
except Exception as ex:
raise self.db_engine_spec.get_dbapi_mapped_exception(ex) from ex
def add_database_to_signature(
self,
@@ -483,11 +572,13 @@ class Database(CoreDatabase, AuditMixinNullable, ImportExportMixin): # pylint:
self,
catalog: str | None = None,
schema: str | None = None,
nullpool: bool = True,
source: utils.QuerySource | None = None,
) -> Connection:
with self.get_sqla_engine(
catalog=catalog,
schema=schema,
nullpool=nullpool,
source=source,
) as engine:
with check_for_oauth2(self):
+8
View File
@@ -3084,6 +3084,14 @@ class ExploreMixin: # pylint: disable=too-many-public-methods
sqla_col = self.convert_tbl_column_to_sqla_col(
tbl_column=col_obj, template_processor=template_processor
)
# Parenthesize expression-based columns to prevent operator
# precedence issues (e.g. OR in a calculated column breaking
# surrounding AND filters). Same pattern as extras.where
# wrapping added in PR #38183.
if sqla_col is not None and (
(col_obj and col_obj.expression) or is_adhoc_column(flt_col)
):
sqla_col = Grouping(sqla_col)
col_type = col_obj.type if col_obj else None
col_spec = db_engine_spec.get_column_spec(native_type=col_type)
is_list_target = op in (
+3
View File
@@ -1568,6 +1568,7 @@ def transpile_to_dialect(
sql: str,
target_engine: str,
source_engine: str | None = None,
identify: bool = False,
) -> str:
"""
Transpile SQL from one database dialect to another using SQLGlot.
@@ -1576,6 +1577,7 @@ def transpile_to_dialect(
sql: The SQL query to transpile
target_engine: The target database engine (e.g., "mysql", "postgresql")
source_engine: The source database engine. If None, uses generic SQL dialect.
identify: If True, quote all identifiers per the target dialect.
Returns:
The transpiled SQL string
@@ -1598,6 +1600,7 @@ def transpile_to_dialect(
copy=True,
comments=False,
pretty=False,
identify=identify,
)
except ParseError as ex:
raise QueryClauseValidationException(f"Cannot parse SQL clause: {sql}") from ex
+2 -28
View File
@@ -18,28 +18,14 @@ from __future__ import annotations
from collections.abc import Hashable, Sequence
from datetime import datetime
from typing import (
Any,
Callable,
ContextManager,
Literal,
TYPE_CHECKING,
TypeAlias,
TypedDict,
)
from typing import Any, Literal, TYPE_CHECKING, TypeAlias, TypedDict
from sqlalchemy.engine.url import URL
from sqlalchemy.sql.type_api import TypeEngine
from typing_extensions import NotRequired
from werkzeug.wrappers import Response
if TYPE_CHECKING:
from superset.models.core import Database
from superset.utils.core import (
GenericDataType,
QueryObjectFilterClause,
QuerySource,
)
from superset.utils.core import GenericDataType, QueryObjectFilterClause
SQLType: TypeAlias = TypeEngine | type[TypeEngine]
@@ -84,18 +70,6 @@ class DatasetMetricData(TypedDict, total=False):
verbose_name: str | None
# Type alias for database connection mutator function
DBConnectionMutator: TypeAlias = Callable[
[URL, dict[str, Any], str | None, Any, "QuerySource | None"],
tuple[URL, dict[str, Any]],
]
# Type alias for engine context manager
EngineContextManager: TypeAlias = Callable[
["Database", str | None, str | None], ContextManager[None]
]
class LegacyMetric(TypedDict):
label: str | None
+5
View File
@@ -183,6 +183,11 @@ def load_explore_json_into_cache( # pylint: disable=too-many-locals
except Exception as ex:
if isinstance(ex, SupersetVizException):
errors = ex.errors
# Extract SIP-40 style errors when available
elif isinstance(ex, SupersetErrorException):
errors = [dataclasses.asdict(ex.error)] # type: ignore
elif isinstance(ex, SupersetErrorsException):
errors = [dataclasses.asdict(error) for error in ex.errors] # type: ignore
else:
error = ex.message if hasattr(ex, "message") else str(ex)
errors = [error] # type: ignore
File diff suppressed because it is too large Load Diff
+10 -1
View File
@@ -68,6 +68,7 @@ from superset.daos.datasource import DatasourceDAO
from superset.dashboards.permalink.exceptions import DashboardPermalinkGetFailedError
from superset.exceptions import (
CacheLoadError,
SupersetErrorException,
SupersetException,
SupersetSecurityException,
)
@@ -266,6 +267,10 @@ class Superset(BaseSupersetView):
)
return self.generate_json(viz_obj, response_type)
except SupersetErrorException:
# Let structured Superset errors (e.g. OAuth2RedirectError) propagate
# so the global Flask error handler serializes them.
raise
except SupersetException as ex:
return json_error_response(utils.error_msg_from_exception(ex), 400)
@@ -290,7 +295,7 @@ class Superset(BaseSupersetView):
@etag_cache()
@check_resource_permissions(check_datasource_perms)
@deprecated(eol_version="5.0.0")
def explore_json(
def explore_json( # noqa: C901
self, datasource_type: str | None = None, datasource_id: int | None = None
) -> FlaskResponse:
"""Serves all request that GET or POST form_data
@@ -377,6 +382,10 @@ class Superset(BaseSupersetView):
)
return self.generate_json(viz_obj, response_type)
except SupersetErrorException:
# Let structured Superset errors (e.g. OAuth2RedirectError) propagate
# so the global Flask error handler serializes them.
raise
except SupersetException as ex:
return json_error_response(utils.error_msg_from_exception(ex), 400)
+5
View File
@@ -51,6 +51,7 @@ from superset.exceptions import (
NullValueException,
QueryObjectValidationError,
SpatialException,
SupersetErrorException,
SupersetSecurityException,
)
from superset.extensions import cache_manager, security_manager
@@ -612,6 +613,10 @@ class BaseViz: # pylint: disable=too-many-public-methods
)
self.errors.append(error)
self.status = QueryStatus.FAILED
except SupersetErrorException:
# Let structured Superset errors (e.g. OAuth2RedirectError) propagate
# so the global Flask error handler serializes them.
raise
except Exception as ex: # pylint: disable=broad-except
logger.exception(ex)
+1
View File
@@ -170,6 +170,7 @@ def example_db_provider() -> Callable[[], Database]:
return self._db
def _load_lazy_data_to_decouple_from_session(self) -> None:
self._db._get_sqla_engine() # type: ignore
self._db.backend # type: ignore # noqa: B018
def remove(self) -> None:
+54 -1
View File
@@ -41,7 +41,7 @@ from superset.common.db_query_status import QueryStatus
from superset.connectors.sqla.models import SqlaTable
from superset.db_engine_specs.base import BaseEngineSpec
from superset.db_engine_specs.mssql import MssqlEngineSpec
from superset.exceptions import SupersetException
from superset.exceptions import OAuth2RedirectError, SupersetException
from superset.extensions import cache_manager
from superset.models import core as models
from superset.models.dashboard import Dashboard
@@ -458,6 +458,59 @@ class TestCore(SupersetTestCase):
assert rv.status_code == 404
assert data["error"] == "Cached data not found"
@pytest.mark.usefixtures("load_world_bank_dashboard_with_slices")
@mock.patch("superset.viz.BaseViz.get_df")
def test_explore_json_propagates_oauth2_redirect_error(
self, mock_get_df: mock.Mock
) -> None:
"""
SupersetErrorException exceptions bubble up properly.
"""
mock_get_df.side_effect = OAuth2RedirectError(
url="https://accounts.example.com/o/oauth2/v2/auth?...",
tab_id="tab-123",
redirect_uri="https://superset.example.com/oauth2/redirect",
)
self.login(ADMIN_USERNAME)
slc = self.get_slice("Life Expectancy VS Rural %")
rv = self.client.post(
f"/superset/explore_json/{slc.datasource_type}/{slc.datasource_id}/",
data={"form_data": json.dumps(slc.form_data)},
)
data = json.loads(rv.data.decode("utf-8"))
assert "errors" in data, data
assert data["errors"][0]["error_type"] == "OAUTH2_REDIRECT"
assert data["errors"][0]["extra"] == {
"url": "https://accounts.example.com/o/oauth2/v2/auth?...",
"tab_id": "tab-123",
"redirect_uri": "https://superset.example.com/oauth2/redirect",
}
@pytest.mark.usefixtures("load_world_bank_dashboard_with_slices")
@mock.patch("superset.viz.BaseViz.get_df")
def test_explore_json_generic_exception_still_returns_viz_get_df_error(
self, mock_get_df: mock.Mock
) -> None:
"""
Non-Superset exceptions raised by ``get_df`` are reported as the
generic ``VIZ_GET_DF_ERROR``.
"""
mock_get_df.side_effect = RuntimeError("boom")
self.login(ADMIN_USERNAME)
slc = self.get_slice("Life Expectancy VS Rural %")
rv = self.client.post(
f"/superset/explore_json/{slc.datasource_type}/{slc.datasource_id}/",
data={"form_data": json.dumps(slc.form_data)},
)
data = json.loads(rv.data.decode("utf-8"))
assert "errors" in data, data
assert data["errors"][0]["error_type"] == "VIZ_GET_DF_ERROR"
assert data["errors"][0]["message"] == "boom"
def test_results_default_deserialization(self):
use_new_deserialization = False
data = [("a", 4, 4.0, "2019-08-18T16:39:16.660000")]
@@ -897,7 +897,7 @@ class TestImportDatabasesCommand(SupersetTestCase):
class TestTestConnectionDatabaseCommand(SupersetTestCase):
@patch("superset.models.core.Database.get_sqla_engine")
@patch("superset.models.core.Database._get_sqla_engine")
@patch("superset.commands.database.test_connection.event_logger.log_with_context")
@patch("superset.utils.core.g")
def test_connection_db_exception(
@@ -906,23 +906,19 @@ class TestTestConnectionDatabaseCommand(SupersetTestCase):
"""Test to make sure event_logger is called when an exception is raised"""
database = get_example_database()
mock_g.user = security_manager.find_user("admin")
mock_get_sqla_engine.return_value.__enter__.side_effect = Exception(
"An error has occurred!"
)
mock_get_sqla_engine.side_effect = Exception("An error has occurred!")
db_uri = database.sqlalchemy_uri_decrypted
json_payload = {"sqlalchemy_uri": db_uri}
command_without_db_name = TestConnectionDatabaseCommand(json_payload)
with pytest.raises(DatabaseTestConnectionUnexpectedError) as excinfo:
with pytest.raises(DatabaseTestConnectionUnexpectedError) as excinfo: # noqa: PT012
command_without_db_name.run()
# Exception wraps errors from db_engine_spec.extract_errors()
assert (
excinfo.value.errors[0].error_type
== SupersetErrorType.GENERIC_DB_ENGINE_ERROR
)
assert str(excinfo.value) == (
"Unexpected error occurred, please check your logs for details"
)
mock_event_logger.assert_called()
@patch("superset.models.core.Database.get_sqla_engine")
@patch("superset.models.core.Database._get_sqla_engine")
@patch("superset.commands.database.test_connection.event_logger.log_with_context")
@patch("superset.utils.core.g")
def test_connection_do_ping_exception(
@@ -931,8 +927,9 @@ class TestTestConnectionDatabaseCommand(SupersetTestCase):
"""Test to make sure do_ping exceptions gets captured"""
database = get_example_database()
mock_g.user = security_manager.find_user("admin")
mock_engine = mock_get_sqla_engine.return_value.__enter__.return_value
mock_engine.dialect.do_ping.side_effect = Exception("An error has occurred!")
mock_get_sqla_engine.return_value.dialect.do_ping.side_effect = Exception(
"An error has occurred!"
)
db_uri = database.sqlalchemy_uri_decrypted
json_payload = {"sqlalchemy_uri": db_uri}
command_without_db_name = TestConnectionDatabaseCommand(json_payload)
@@ -970,7 +967,7 @@ class TestTestConnectionDatabaseCommand(SupersetTestCase):
== SupersetErrorType.CONNECTION_DATABASE_TIMEOUT
)
@patch("superset.models.core.Database.get_sqla_engine")
@patch("superset.models.core.Database._get_sqla_engine")
@patch("superset.commands.database.test_connection.event_logger.log_with_context")
@patch("superset.utils.core.g")
def test_connection_superset_security_connection(
@@ -980,20 +977,20 @@ class TestTestConnectionDatabaseCommand(SupersetTestCase):
connection exc is raised"""
database = get_example_database()
mock_g.user = security_manager.find_user("admin")
mock_get_sqla_engine.return_value.__enter__.side_effect = (
SupersetSecurityException(
SupersetError(error_type=500, message="test", level="info")
)
mock_get_sqla_engine.side_effect = SupersetSecurityException(
SupersetError(error_type=500, message="test", level="info")
)
db_uri = database.sqlalchemy_uri_decrypted
json_payload = {"sqlalchemy_uri": db_uri}
command_without_db_name = TestConnectionDatabaseCommand(json_payload)
with pytest.raises(DatabaseSecurityUnsafeError):
with pytest.raises(DatabaseSecurityUnsafeError) as excinfo: # noqa: PT012
command_without_db_name.run()
assert str(excinfo.value) == ("Stopped an unsafe database connection")
mock_event_logger.assert_called()
@patch("superset.models.core.Database.get_sqla_engine")
@patch("superset.models.core.Database._get_sqla_engine")
@patch("superset.commands.database.test_connection.event_logger.log_with_context")
@patch("superset.utils.core.g")
def test_connection_db_api_exc(
@@ -1002,20 +999,19 @@ class TestTestConnectionDatabaseCommand(SupersetTestCase):
"""Test to make sure event_logger is called when DBAPIError is raised"""
database = get_example_database()
mock_g.user = security_manager.find_user("admin")
mock_get_sqla_engine.return_value.__enter__.side_effect = DBAPIError(
mock_get_sqla_engine.side_effect = DBAPIError(
statement="error", params={}, orig={}
)
db_uri = database.sqlalchemy_uri_decrypted
json_payload = {"sqlalchemy_uri": db_uri}
command_without_db_name = TestConnectionDatabaseCommand(json_payload)
with pytest.raises(SupersetErrorsException) as excinfo:
with pytest.raises(SupersetErrorsException) as excinfo: # noqa: PT012
command_without_db_name.run()
# Exception wraps errors from db_engine_spec.extract_errors()
assert (
excinfo.value.errors[0].error_type
== SupersetErrorType.GENERIC_DB_ENGINE_ERROR
)
assert str(excinfo.value) == (
"Connection failed, please check your connection settings"
)
mock_event_logger.assert_called()
@@ -1151,7 +1147,7 @@ class TestTablesDatabaseCommand(SupersetTestCase):
with pytest.raises(DatabaseNotFoundError) as excinfo: # noqa: PT012
command.run()
assert str(excinfo.value) == ("Database not found.")
assert str(excinfo.value) == ("Database not found.")
@patch("superset.daos.database.DatabaseDAO.find_by_id")
@patch("superset.security.manager.SupersetSecurityManager.can_access_database")
@@ -1170,35 +1166,26 @@ class TestTablesDatabaseCommand(SupersetTestCase):
command = TablesDatabaseCommand(database.id, None, "main", False)
with pytest.raises(SupersetException) as excinfo: # noqa: PT012
command.run()
assert str(excinfo.value) == "Test Error"
assert str(excinfo.value) == "Test Error"
@patch("superset.daos.database.DatabaseDAO.find_by_id")
@patch("superset.models.core.Database.get_all_materialized_view_names_in_schema")
@patch("superset.models.core.Database.get_all_view_names_in_schema")
@patch("superset.models.core.Database.get_all_table_names_in_schema")
@patch("superset.security.manager.SupersetSecurityManager.can_access_database")
@patch("superset.utils.core.g")
def test_database_tables_exception(
self,
mock_g,
mock_can_access_database,
mock_get_tables,
mock_get_views,
mock_get_mvs,
mock_find_by_id,
self, mock_g, mock_can_access_database, mock_find_by_id
):
database = get_example_database()
mock_find_by_id.return_value = database
mock_get_tables.return_value = {("table1", "main", None)}
mock_get_views.return_value = set()
mock_get_mvs.return_value = []
mock_can_access_database.side_effect = Exception("Test Error")
mock_g.user = security_manager.find_user("admin")
command = TablesDatabaseCommand(database.id, None, "main", False)
with pytest.raises(DatabaseTablesUnexpectedError) as excinfo: # noqa: PT012
command.run()
assert str(excinfo.value) == "Test Error"
assert (
str(excinfo.value)
== "Unexpected error occurred, please check your logs for details"
)
@patch("superset.daos.database.DatabaseDAO.find_by_id")
@patch("superset.security.manager.SupersetSecurityManager.can_access_database")
+14 -23
View File
@@ -145,7 +145,7 @@ class TestDatabaseModel(SupersetTestCase):
username = make_url(engine.url).username
assert example_user.username != username
@mock.patch("superset.engines.manager.create_engine")
@mock.patch("superset.models.core.create_engine")
@unittest.skipUnless(
SupersetTestCase.is_module_installed("pyhive"), "pyhive not installed"
)
@@ -172,8 +172,7 @@ class TestDatabaseModel(SupersetTestCase):
database_name="test_database", sqlalchemy_uri=uri, extra=extra
)
model.impersonate_user = True
with model.get_sqla_engine():
pass
model._get_sqla_engine()
call_args = mocked_create_engine.call_args
assert str(call_args[0][0]) == "presto://gamma@localhost/"
@@ -186,8 +185,7 @@ class TestDatabaseModel(SupersetTestCase):
}
model.impersonate_user = False
with model.get_sqla_engine():
pass
model._get_sqla_engine()
call_args = mocked_create_engine.call_args
assert str(call_args[0][0]) == "presto://localhost/"
@@ -201,14 +199,13 @@ class TestDatabaseModel(SupersetTestCase):
@unittest.skipUnless(
SupersetTestCase.is_module_installed("mysqlclient"), "mysqlclient not installed"
)
@mock.patch("superset.engines.manager.create_engine")
@mock.patch("superset.models.core.create_engine")
def test_adjust_engine_params_mysql(self, mocked_create_engine):
model = Database(
database_name="test_database1",
sqlalchemy_uri="mysql://user:password@localhost",
)
with model.get_sqla_engine():
pass
model._get_sqla_engine()
call_args = mocked_create_engine.call_args
assert str(call_args[0][0]) == "mysql://user:password@localhost"
@@ -218,14 +215,13 @@ class TestDatabaseModel(SupersetTestCase):
database_name="test_database2",
sqlalchemy_uri="mysql+mysqlconnector://user:password@localhost",
)
with model.get_sqla_engine():
pass
model._get_sqla_engine()
call_args = mocked_create_engine.call_args
assert str(call_args[0][0]) == "mysql+mysqlconnector://user:password@localhost"
assert call_args[1]["connect_args"]["allow_local_infile"] == 0
@mock.patch("superset.engines.manager.create_engine")
@mock.patch("superset.models.core.create_engine")
def test_impersonate_user_trino(self, mocked_create_engine):
principal_user = security_manager.find_user(username="gamma")
@@ -234,8 +230,7 @@ class TestDatabaseModel(SupersetTestCase):
database_name="test_database", sqlalchemy_uri="trino://localhost"
)
model.impersonate_user = True
with model.get_sqla_engine():
pass
model._get_sqla_engine()
call_args = mocked_create_engine.call_args
assert str(call_args[0][0]) == "trino://localhost/"
@@ -247,8 +242,7 @@ class TestDatabaseModel(SupersetTestCase):
)
model.impersonate_user = True
with model.get_sqla_engine():
pass
model._get_sqla_engine()
call_args = mocked_create_engine.call_args
assert (
@@ -257,7 +251,7 @@ class TestDatabaseModel(SupersetTestCase):
)
assert call_args[1]["connect_args"]["user"] == "gamma"
@mock.patch("superset.engines.manager.create_engine")
@mock.patch("superset.models.core.create_engine")
@unittest.skipUnless(
SupersetTestCase.is_module_installed("pyhive"), "pyhive not installed"
)
@@ -287,8 +281,7 @@ class TestDatabaseModel(SupersetTestCase):
database_name="test_database", sqlalchemy_uri=uri, extra=extra
)
model.impersonate_user = True
with model.get_sqla_engine():
pass
model._get_sqla_engine()
call_args = mocked_create_engine.call_args
assert str(call_args[0][0]) == "hive://localhost"
@@ -301,8 +294,7 @@ class TestDatabaseModel(SupersetTestCase):
}
model.impersonate_user = False
with model.get_sqla_engine():
pass
model._get_sqla_engine()
call_args = mocked_create_engine.call_args
assert str(call_args[0][0]) == "hive://localhost"
@@ -384,7 +376,7 @@ class TestDatabaseModel(SupersetTestCase):
df = main_db.get_df("USE superset; SELECT ';';", None, None)
assert df.iat[0, 0] == ";"
@mock.patch("superset.engines.manager.create_engine")
@mock.patch("superset.models.core.create_engine")
def test_get_sqla_engine(self, mocked_create_engine):
model = Database(
database_name="test_database",
@@ -395,8 +387,7 @@ class TestDatabaseModel(SupersetTestCase):
)
mocked_create_engine.side_effect = Exception()
with self.assertRaises(SupersetException): # noqa: PT027
with model.get_sqla_engine():
pass
model._get_sqla_engine()
class TestSqlaTableModel(SupersetTestCase):
@@ -624,6 +624,65 @@ def test_get_tab_urls(
]
@patch("superset.commands.report.execute.CreateDashboardPermalinkCommand")
@with_feature_flags(ALERT_REPORT_TABS=True)
def test_get_dashboard_urls_multitab_preserves_url_params(
mock_permalink_cls,
mocker: MockerFixture,
app,
) -> None:
"""Multi-tab fan-out must preserve dashboard_state.urlParams (e.g. standalone)
and replace any pre-existing native_filters entry with the report's value —
matching the single-tab branch's merge semantics."""
mock_report_schedule: ReportSchedule = mocker.Mock(spec=ReportSchedule)
mock_report_schedule.chart = False
mock_report_schedule.chart_id = None
mock_report_schedule.dashboard_id = 123
mock_report_schedule.type = "report_type"
mock_report_schedule.report_format = "report_format"
mock_report_schedule.owners = [1, 2]
mock_report_schedule.recipients = []
native_filter_rison = "(NATIVE_FILTER-1:(filterType:filter_select))"
# Use list-of-lists (not tuples) — extra_json deserializes urlParams from
# JSON arrays. Includes a stale native_filters entry to exercise the
# dedup-then-append step in the merge.
mock_report_schedule.extra = {
"dashboard": {
"anchor": json.dumps(["TAB-1", "TAB-2"]),
"urlParams": [
["standalone", "true"],
["native_filters", "(STALE_FILTER:(filterType:filter_select))"],
["show_filters", "0"],
],
}
}
mock_report_schedule.get_native_filters_params.return_value = ( # type: ignore[attr-defined]
native_filter_rison,
[],
)
mock_permalink_cls.return_value.run.side_effect = ["key1", "key2"]
class_instance: BaseReportState = BaseReportState(
mock_report_schedule, "January 1, 2021", "execution_id_example"
)
class_instance._report_schedule = mock_report_schedule
class_instance.get_dashboard_urls()
assert mock_permalink_cls.call_count == 2
for idx, expected_anchor in enumerate(["TAB-1", "TAB-2"]):
state = mock_permalink_cls.call_args_list[idx].kwargs["state"]
# Stale native_filters is replaced (not duplicated); other params
# survive in their original order; report's native_filters appended.
assert state["urlParams"] == [
["standalone", "true"],
["show_filters", "0"],
["native_filters", native_filter_rison],
]
# Each per-tab permalink targets exactly that tab.
assert state["anchor"] == expected_anchor
@patch(
"superset.commands.dashboard.permalink.create.CreateDashboardPermalinkCommand.run"
)
@@ -702,6 +761,58 @@ def test_get_dashboard_urls_native_filters_without_tabs(
assert "permalink_key" in result[0]
@patch("superset.commands.report.execute.CreateDashboardPermalinkCommand")
@with_feature_flags(ALERT_REPORT_TABS=False)
def test_get_dashboard_urls_flag_off_preserves_url_params(
mock_permalink_cls,
mocker: MockerFixture,
app,
) -> None:
"""The post-``if``-block fall-through in ``get_dashboard_urls`` must
honor any urlParams set in ``extra.dashboard`` (e.g. via API) same
merge semantics as the protected branch.
Reachability: only when ``dashboard_state`` is falsy OR
``ALERT_REPORT_TABS=False``. The flag-on / no-anchor case lands in
the single-tab merge at L290-306, not here.
"""
mock_report_schedule: ReportSchedule = mocker.Mock(spec=ReportSchedule)
mock_report_schedule.chart = False
mock_report_schedule.chart_id = None
mock_report_schedule.dashboard_id = 123
native_filter_rison = "(NATIVE_FILTER-abc:!(val1))"
mock_report_schedule.extra = {
"dashboard": {
"urlParams": [
["standalone", "true"],
["native_filters", "(STALE_FILTER:!(stale))"],
["show_filters", "0"],
],
}
}
mock_report_schedule.get_native_filters_params.return_value = ( # type: ignore[attr-defined]
native_filter_rison,
[],
)
class_instance: BaseReportState = BaseReportState(
mock_report_schedule, "January 1, 2021", "execution_id_example"
)
class_instance._report_schedule = mock_report_schedule
mock_permalink_cls.return_value.run.return_value = "permalink_key"
class_instance.get_dashboard_urls()
state = mock_permalink_cls.call_args_list[0].kwargs["state"]
# Stale native_filters replaced; existing params survive in order;
# report's native_filters appended.
assert state["urlParams"] == [
["standalone", "true"],
["show_filters", "0"],
["native_filters", native_filter_rison],
]
def create_report_schedule(
mocker: MockerFixture,
custom_width: int | None = None,
-109
View File
@@ -1,109 +0,0 @@
# Licensed to the Apache Software Foundation (ASF) under one
# or more contributor license agreements. See the NOTICE file
# distributed with this work for additional information
# regarding copyright ownership. The ASF licenses this file
# to you under the Apache License, Version 2.0 (the
# "License"); you may not use this file except in compliance
# with the License. You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing,
# software distributed under the License is distributed on an
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.
from contextlib import contextmanager
from unittest.mock import MagicMock, patch
import pytest
from sqlalchemy import pool
from sqlalchemy.engine.url import make_url
from superset.commands.database.ssh_tunnel.exceptions import SSHTunnelDatabasePortError
from superset.engines.manager import EngineManager
@pytest.fixture
def engine_manager() -> EngineManager:
@contextmanager
def dummy_context_manager(
database: MagicMock,
catalog: str | None,
schema: str | None,
):
yield
return EngineManager(engine_context_manager=dummy_context_manager)
@pytest.fixture
def mock_database() -> MagicMock:
database = MagicMock()
database.id = 1
database.sqlalchemy_uri_decrypted = "trino://"
database.get_extra.return_value = {"engine_params": {"poolclass": "queue"}}
database.get_effective_user.return_value = "alice"
database.impersonate_user = False
database.update_params_from_encrypted_extra = MagicMock()
database.db_engine_spec = MagicMock()
database.db_engine_spec.adjust_engine_params.return_value = (
make_url("trino://"),
{"source": "Apache Superset"},
)
database.db_engine_spec.validate_database_uri = MagicMock()
return database
@patch("superset.engines.manager.make_url_safe")
def test_get_engine_args_uses_null_pool(
mock_make_url: MagicMock,
engine_manager: EngineManager,
mock_database: MagicMock,
) -> None:
mock_make_url.return_value = make_url("trino://")
_, kwargs = engine_manager._get_engine_args(mock_database, None, None, None, None)
assert kwargs["poolclass"] is pool.NullPool
@patch("superset.engines.manager.make_url_safe")
def test_get_engine_args_with_impersonation(
mock_make_url: MagicMock,
engine_manager: EngineManager,
mock_database: MagicMock,
) -> None:
mock_make_url.return_value = make_url("trino://")
mock_database.impersonate_user = True
mock_database.get_oauth2_config.return_value = None
mock_database.db_engine_spec.impersonate_user.return_value = (
make_url("trino://"),
{"connect_args": {"user": "alice"}, "poolclass": pool.NullPool},
)
engine_manager._get_engine_args(mock_database, None, None, None, None)
mock_database.db_engine_spec.impersonate_user.assert_called_once()
def test_get_tunnel_kwargs_requires_database_port(
engine_manager: EngineManager,
) -> None:
ssh_tunnel = MagicMock()
ssh_tunnel.server_address = "ssh.example.com"
ssh_tunnel.server_port = 22
ssh_tunnel.username = "ssh_user"
ssh_tunnel.password = None
ssh_tunnel.private_key = None
ssh_tunnel.private_key_password = None
uri = MagicMock()
uri.port = None
uri.get_backend_name.return_value = "unknown"
with patch("superset.utils.ssh_tunnel.get_default_port", return_value=None):
with pytest.raises(SSHTunnelDatabasePortError):
engine_manager._get_tunnel_kwargs(ssh_tunnel, uri)
@@ -14,3 +14,23 @@
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.
from unittest.mock import Mock
import sshtunnel
from superset.extensions.ssh import SSHManagerFactory
def test_ssh_tunnel_timeout_setting() -> None:
app = Mock()
app.config = {
"SSH_TUNNEL_MAX_RETRIES": 2,
"SSH_TUNNEL_LOCAL_BIND_ADDRESS": "test",
"SSH_TUNNEL_TIMEOUT_SEC": 123.0,
"SSH_TUNNEL_PACKET_TIMEOUT_SEC": 321.0,
"SSH_TUNNEL_MANAGER_CLASS": "superset.extensions.ssh.SSHManager",
}
factory = SSHManagerFactory()
factory.init_app(app)
assert sshtunnel.TUNNEL_TIMEOUT == 123.0
assert sshtunnel.SSH_TIMEOUT == 321.0
+1 -1
View File
@@ -124,7 +124,7 @@ class TestSupersetAppInitializer:
patch.object(app_initializer, "configure_data_sources"),
patch.object(app_initializer, "configure_auth_provider"),
patch.object(app_initializer, "configure_async_queries"),
patch.object(app_initializer, "configure_engine_manager"),
patch.object(app_initializer, "configure_ssh_manager"),
patch.object(app_initializer, "configure_stats_manager"),
patch.object(app_initializer, "init_views"),
):
@@ -0,0 +1,445 @@
# Licensed to the Apache Software Foundation (ASF) under one
# or more contributor license agreements. See the NOTICE file
# distributed with this work for additional information
# regarding copyright ownership. The ASF licenses this file
# to you under the Apache License, Version 2.0 (the
# "License"); you may not use this file except in compliance
# with the License. You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing,
# software distributed under the License is distributed on an
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.
"""
Integration-style tests for ``validate_and_compile``.
These tests exercise the real ``DatasetValidator.validate_against_dataset``
path so fast-path tools (``generate_explore_link``, ``update_chart_preview``)
that only use Tier-1 validation are exercised end-to-end.
"""
from unittest.mock import Mock, patch
import pytest
from superset.mcp_service.chart.compile import (
build_dataset_context_from_orm,
CompileResult,
validate_and_compile,
)
from superset.mcp_service.chart.schemas import (
BigNumberChartConfig,
ColumnRef,
FilterConfig,
PieChartConfig,
PivotTableChartConfig,
TableChartConfig,
XYChartConfig,
)
def _orm_dataset(
*,
column_names: list[str] | None = None,
metric_names: list[str] | None = None,
has_database: bool = True,
) -> Mock:
"""Build a Mock dataset that satisfies build_dataset_context_from_orm."""
columns = []
for name in column_names or ["ds", "gender", "name", "num"]:
col = Mock()
col.column_name = name
col.type = "TEXT"
col.is_temporal = name == "ds"
col.is_numeric = name == "num"
columns.append(col)
metrics = []
for name in metric_names or ["sum_boys", "sum_girls"]:
m = Mock()
m.metric_name = name
m.expression = f"SUM({name})"
m.description = None
metrics.append(m)
dataset = Mock()
dataset.id = 3
dataset.table_name = "birth_names"
dataset.schema = None
dataset.columns = columns
dataset.metrics = metrics
if has_database:
db = Mock()
db.database_name = "examples"
dataset.database = db
else:
dataset.database = None
return dataset
class TestBuildDatasetContextFromOrm:
"""Cover the helper that converts ORM dataset → DatasetContext."""
def test_handles_missing_database_relationship(self):
"""``database_name`` defaults to '' when ``dataset.database`` is None
so Pydantic validation doesn't blow up."""
ds = _orm_dataset(has_database=False)
ctx = build_dataset_context_from_orm(ds)
assert ctx is not None
assert ctx.database_name == ""
assert ctx.id == 3
assert {c["name"] for c in ctx.available_columns} == {
"ds",
"gender",
"name",
"num",
}
assert {m["name"] for m in ctx.available_metrics} == {
"sum_boys",
"sum_girls",
}
def test_returns_none_for_none_input(self):
assert build_dataset_context_from_orm(None) is None
class TestValidateAndCompileChartTypeCoverage:
"""Tier-1 validation must catch bad column refs in every supported
chart-config variant not just XY and table."""
def test_xy_bad_metric_column_rejected(self):
ds = _orm_dataset()
config = XYChartConfig(
chart_type="xy",
x=ColumnRef(name="ds"),
y=[ColumnRef(name="num_boys", aggregate="SUM")],
kind="line",
)
result = validate_and_compile(config, {}, ds, run_compile_check=False)
assert not result.success
assert result.tier == "validation"
assert result.error_obj is not None
assert any("sum_boys" in s for s in (result.error_obj.suggestions or []))
def test_pie_bad_metric_column_rejected(self):
ds = _orm_dataset()
config = PieChartConfig(
dimension=ColumnRef(name="gender"),
metric=ColumnRef(name="num_boys", aggregate="SUM"),
)
result = validate_and_compile(config, {}, ds, run_compile_check=False)
assert not result.success, "Pie chart with bad metric column should fail"
assert result.tier == "validation"
assert result.error_obj is not None
assert any("sum_boys" in s for s in (result.error_obj.suggestions or []))
def test_pie_valid_dimension_and_saved_metric_passes(self):
ds = _orm_dataset()
config = PieChartConfig(
dimension=ColumnRef(name="gender"),
metric=ColumnRef(name="sum_boys", saved_metric=True),
)
result = validate_and_compile(config, {}, ds, run_compile_check=False)
assert result.success, result.error
def test_pivot_table_bad_row_rejected(self):
ds = _orm_dataset()
config = PivotTableChartConfig(
rows=[ColumnRef(name="bogus_dim")],
metrics=[ColumnRef(name="sum_boys", saved_metric=True)],
)
result = validate_and_compile(config, {}, ds, run_compile_check=False)
assert not result.success
assert result.error_obj is not None
def test_big_number_bad_temporal_column_rejected(self):
ds = _orm_dataset()
config = BigNumberChartConfig(
chart_type="big_number",
metric=ColumnRef(name="sum_boys", saved_metric=True),
temporal_column="not_a_real_temporal",
show_trendline=True,
)
result = validate_and_compile(config, {}, ds, run_compile_check=False)
assert not result.success, "BigNumber temporal_column must be validated"
assert result.error_obj is not None
assert "not_a_real_temporal" in (result.error_obj.message or "")
def test_pie_with_sum_on_non_numeric_column_rejected(self):
"""Tier-1 aggregation compatibility now runs for non-Table/XY too —
a pie ``metric={"name": "gender", "aggregate": "SUM"}`` would emit
``SUM(gender)`` which the DB rejects, so the validator must catch it
before we hand back an explore URL."""
ds = _orm_dataset()
config = PieChartConfig(
dimension=ColumnRef(name="name"),
metric=ColumnRef(name="gender", aggregate="SUM"),
)
result = validate_and_compile(config, {}, ds, run_compile_check=False)
assert not result.success, "SUM on a TEXT column must reject"
assert result.error_obj is not None
assert result.error_obj.error_code == "INVALID_AGGREGATION"
def test_pivot_table_sum_on_non_numeric_column_rejected(self):
ds = _orm_dataset()
config = PivotTableChartConfig(
rows=[ColumnRef(name="gender")],
metrics=[ColumnRef(name="name", aggregate="SUM")],
)
result = validate_and_compile(config, {}, ds, run_compile_check=False)
assert not result.success
assert result.error_obj is not None
assert result.error_obj.error_code == "INVALID_AGGREGATION"
def test_pivot_table_min_on_non_numeric_column_passes(self):
"""MIN and MAX are not numeric-only (valid on dates/text in SQL).
They are left to the Tier-2 compile check rather than being rejected
by Tier-1 schema validation.
"""
ds = _orm_dataset()
config = PivotTableChartConfig(
rows=[ColumnRef(name="gender")],
metrics=[ColumnRef(name="name", aggregate="MIN")],
)
result = validate_and_compile(config, {}, ds, run_compile_check=False)
assert result.success, (
"MIN on a text column should not be rejected by Tier-1 validation"
)
def test_table_with_invalid_filter_column_rejected(self):
ds = _orm_dataset()
config = TableChartConfig(
chart_type="table",
columns=[ColumnRef(name="gender")],
filters=[FilterConfig(column="bogus", op="=", value="x")],
)
result = validate_and_compile(config, {}, ds, run_compile_check=False)
assert not result.success
assert result.error_obj is not None
class TestSavedMetricNotMarked:
"""A non-saved-metric ColumnRef whose name matches a saved metric is a
common LLM mistake (forgetting to set ``saved_metric=true``). The
validator should surface a tailored hint instead of letting the bad SQL
through."""
def test_table_metric_name_without_saved_metric_flag_rejected(self):
ds = _orm_dataset()
config = TableChartConfig(
chart_type="table",
columns=[
ColumnRef(name="gender"),
# ``sum_boys`` is a saved metric on the dataset, but
# saved_metric=False (default) would render as
# ``SUM(sum_boys)`` ad-hoc SQL — broken.
ColumnRef(name="sum_boys", aggregate="SUM"),
],
)
result = validate_and_compile(config, {}, ds, run_compile_check=False)
assert not result.success, (
"ref.name matches a saved metric but saved_metric=False -> reject"
)
assert result.error_obj is not None
assert result.error_obj.error_code == "SAVED_METRIC_NOT_MARKED"
# Suggestion should point the LLM at the right correction.
suggestions_text = " ".join(result.error_obj.suggestions or [])
assert "saved_metric" in suggestions_text
assert "sum_boys" in suggestions_text
def test_pie_metric_name_without_saved_metric_flag_rejected(self):
ds = _orm_dataset()
config = PieChartConfig(
dimension=ColumnRef(name="gender"),
metric=ColumnRef(name="sum_boys", aggregate="SUM"),
)
result = validate_and_compile(config, {}, ds, run_compile_check=False)
assert not result.success
assert result.error_obj is not None
assert result.error_obj.error_code == "SAVED_METRIC_NOT_MARKED"
def test_explicit_saved_metric_passes(self):
ds = _orm_dataset()
config = PieChartConfig(
dimension=ColumnRef(name="gender"),
metric=ColumnRef(name="sum_boys", saved_metric=True),
)
result = validate_and_compile(config, {}, ds, run_compile_check=False)
assert result.success, result.error
class TestAdhocFiltersFromFormData:
"""Filters merged into form_data (not present on the typed config) must
also be validated. Without this hook, ``update_chart_preview`` could
smuggle bad column refs through preserved adhoc filters."""
def test_unknown_adhoc_filter_subject_rejected(self):
ds = _orm_dataset()
config = TableChartConfig(
chart_type="table", columns=[ColumnRef(name="gender")]
)
form_data = {
"adhoc_filters": [
{
"expressionType": "SIMPLE",
"subject": "removed_column",
"operator": "==",
"comparator": "x",
}
]
}
result = validate_and_compile(config, form_data, ds, run_compile_check=False)
assert not result.success
assert result.error_obj is not None
assert "removed_column" in (result.error_obj.message or "")
def test_known_adhoc_filter_subject_passes(self):
ds = _orm_dataset()
config = TableChartConfig(
chart_type="table", columns=[ColumnRef(name="gender")]
)
form_data = {
"adhoc_filters": [
{
"expressionType": "SIMPLE",
"subject": "gender",
"operator": "==",
"comparator": "boy",
}
]
}
result = validate_and_compile(config, form_data, ds, run_compile_check=False)
assert result.success, result.error
def test_sql_expression_filter_skipped(self):
"""SQL-expression filters carry a free-form ``sqlExpression`` we can't
safely parse, so they should pass Tier-1 untouched."""
ds = _orm_dataset()
config = TableChartConfig(
chart_type="table", columns=[ColumnRef(name="gender")]
)
form_data = {
"adhoc_filters": [
{
"expressionType": "SQL",
"clause": "WHERE",
"sqlExpression": "1 = 1",
}
]
}
result = validate_and_compile(config, form_data, ds, run_compile_check=False)
assert result.success
def test_where_filter_with_metric_name_rejected(self):
"""A saved-metric name used as a WHERE filter subject must be rejected.
WHERE filters need a physical column; metric names are only valid in
HAVING clauses where Superset can resolve them.
"""
ds = _orm_dataset()
config = TableChartConfig(
chart_type="table", columns=[ColumnRef(name="gender")]
)
form_data = {
"adhoc_filters": [
{
"expressionType": "SIMPLE",
"clause": "WHERE",
"subject": "sum_boys", # saved metric, not a physical column
"operator": ">",
"comparator": "0",
}
]
}
result = validate_and_compile(config, form_data, ds, run_compile_check=False)
assert not result.success, (
"A saved-metric name used in a WHERE filter must not pass Tier-1"
)
assert result.error_obj is not None
assert "sum_boys" in (result.error_obj.message or "")
def test_having_filter_with_metric_name_passes(self):
"""A saved-metric name used in a HAVING filter must be accepted.
HAVING filters are aggregate-level conditions; Superset resolves metric
names there so they are valid references.
"""
ds = _orm_dataset()
config = TableChartConfig(
chart_type="table", columns=[ColumnRef(name="gender")]
)
form_data = {
"adhoc_filters": [
{
"expressionType": "SIMPLE",
"clause": "HAVING",
"subject": "sum_boys", # saved metric — valid in HAVING
"operator": ">",
"comparator": "0",
}
]
}
result = validate_and_compile(config, form_data, ds, run_compile_check=False)
assert result.success, (
"A saved-metric name in a HAVING filter should pass Tier-1 validation"
)
class TestValidateAndCompileTier2:
"""When ``run_compile_check=True`` and Tier-1 passes, the helper must
invoke ``_compile_chart`` and surface its outcome."""
@patch("superset.mcp_service.chart.compile._compile_chart")
def test_tier2_runs_when_tier1_passes(self, mock_compile):
mock_compile.return_value = CompileResult(success=True)
ds = _orm_dataset()
config = TableChartConfig(
chart_type="table", columns=[ColumnRef(name="gender")]
)
result = validate_and_compile(
config, {"adhoc_filters": []}, ds, run_compile_check=True
)
assert result.success
mock_compile.assert_called_once()
@patch("superset.mcp_service.chart.compile._compile_chart")
def test_tier2_skipped_on_tier1_failure(self, mock_compile):
ds = _orm_dataset()
config = TableChartConfig(chart_type="table", columns=[ColumnRef(name="bogus")])
result = validate_and_compile(config, {}, ds, run_compile_check=True)
assert not result.success
assert result.tier == "validation"
mock_compile.assert_not_called()
def test_dataset_none_returns_dataset_not_found(self):
result = validate_and_compile(None, {}, None, run_compile_check=True)
assert not result.success
assert result.error_code == "DATASET_NOT_FOUND"
@pytest.mark.parametrize(
"config_factory",
[
lambda: PieChartConfig(
dimension=ColumnRef(name="gender"),
metric=ColumnRef(name="sum_boys", saved_metric=True),
),
lambda: TableChartConfig(
chart_type="table",
columns=[
ColumnRef(name="gender"),
ColumnRef(name="sum_boys", saved_metric=True),
],
),
],
)
def test_valid_configs_pass_tier1(config_factory):
ds = _orm_dataset()
result = validate_and_compile(config_factory(), {}, ds, run_compile_check=False)
assert result.success, result.error
@@ -745,6 +745,9 @@ class TestUpdateChartNameOnly:
class TestUpdateChartPreviewFirst:
"""Integration-style tests for the preview-first default flow."""
@patch.object(
update_chart_module, "_validate_update_against_dataset", return_value=None
)
@patch.object(update_chart_module, "_create_preview_url", new_callable=Mock)
@patch(
"superset.commands.chart.update.UpdateChartCommand",
@@ -764,6 +767,7 @@ class TestUpdateChartPreviewFirst:
mock_check_access,
mock_update_cmd_cls,
mock_create_preview,
unused_validate_mock,
mcp_server,
):
"""Default update flow returns a preview URL and does NOT save."""
@@ -934,6 +938,9 @@ class TestBuildPreviewFormData:
class TestUpdateChartSaveWithConfig:
"""Save-path integration tests for update_chart with a full config payload."""
@patch.object(
update_chart_module, "_validate_update_against_dataset", return_value=None
)
@patch(
"superset.commands.chart.update.UpdateChartCommand",
new_callable=Mock,
@@ -951,6 +958,7 @@ class TestUpdateChartSaveWithConfig:
mock_find_by_id,
mock_check_access,
mock_update_cmd_cls,
unused_validate_mock,
mcp_server,
):
"""generate_preview=False with config persists and returns saved chart."""
@@ -1086,6 +1094,9 @@ class TestUpdateChartErrorPaths:
assert error["error_type"] == "CommandException"
assert "boom" in error["details"]
@patch.object(
update_chart_module, "_validate_update_against_dataset", return_value=None
)
@patch.object(update_chart_module, "_create_preview_url", new_callable=Mock)
@patch(
"superset.mcp_service.auth.check_chart_data_access",
@@ -1100,6 +1111,7 @@ class TestUpdateChartErrorPaths:
mock_find_by_id,
mock_check_access,
mock_create_preview,
unused_validate_mock,
mcp_server,
):
"""If _create_preview_url returns (url, None), form_data_key comes from url."""
@@ -1137,3 +1149,140 @@ class TestUpdateChartErrorPaths:
assert result.structured_content["success"] is True
assert result.structured_content["form_data_key"] == "url_embedded_key"
class TestUpdateChartValidationGate:
"""Tier-1+2 validation prevents bad config from reaching DB or cache."""
@staticmethod
def _mock_chart_with_dataset(chart_id: int = 1) -> Mock:
chart = Mock()
chart.id = chart_id
chart.datasource_id = 10
chart.slice_name = "Existing"
chart.viz_type = "table"
chart.uuid = "abc-123"
chart.params = '{"viz_type": "table", "datasource": "10__table"}'
# validate_and_compile is mocked, so dataset shape doesn't matter.
chart.datasource = Mock()
return chart
@patch.object(update_chart_module, "validate_and_compile")
@patch.object(update_chart_module, "_create_preview_url", new_callable=Mock)
@patch(
"superset.mcp_service.auth.check_chart_data_access",
new_callable=Mock,
)
@patch("superset.daos.chart.ChartDAO.find_by_id", new_callable=Mock)
@patch("superset.db.session")
@pytest.mark.asyncio
async def test_preview_path_validation_failure_skips_cache(
self,
mock_db_session,
mock_find_by_id,
mock_check_access,
mock_create_preview,
mock_validate,
mcp_server,
):
"""Preview path: bad column → structured error, _create_preview_url
must NOT be called."""
from superset.mcp_service.chart.compile import CompileResult
from superset.mcp_service.common.error_schemas import ChartGenerationError
mock_find_by_id.return_value = self._mock_chart_with_dataset()
mock_check_access.return_value = DatasetValidationResult(
is_valid=True, dataset_id=10, dataset_name="ds", warnings=[]
)
mock_validate.return_value = CompileResult(
success=False,
error="Column 'num_boys' does not exist",
error_code="CHART_VALIDATION_FAILED",
tier="validation",
error_obj=ChartGenerationError(
error_type="invalid_column",
message="Column 'num_boys' does not exist",
details="Available: ds, gender, sum_boys",
suggestions=["sum_boys"],
error_code="CHART_VALIDATION_FAILED",
),
)
request = {
"identifier": 1,
"config": {
"chart_type": "xy",
"x": {"name": "ds"},
"y": [{"name": "num_boys", "aggregate": "SUM"}],
"kind": "line",
},
}
async with Client(mcp) as client:
result = await client.call_tool("update_chart", {"request": request})
assert result.structured_content["success"] is False
error = result.structured_content["error"]
assert error["error_code"] == "CHART_VALIDATION_FAILED"
assert "sum_boys" in error["suggestions"]
mock_create_preview.assert_not_called()
@patch.object(update_chart_module, "validate_and_compile")
@patch(
"superset.commands.chart.update.UpdateChartCommand",
new_callable=Mock,
)
@patch(
"superset.mcp_service.auth.check_chart_data_access",
new_callable=Mock,
)
@patch("superset.daos.chart.ChartDAO.find_by_id", new_callable=Mock)
@patch("superset.db.session")
@pytest.mark.asyncio
async def test_persist_path_validation_failure_skips_db_write(
self,
mock_db_session,
mock_find_by_id,
mock_check_access,
mock_update_cmd_cls,
mock_validate,
mcp_server,
):
"""Persist path: validation failure → UpdateChartCommand NOT called."""
from superset.mcp_service.chart.compile import CompileResult
from superset.mcp_service.common.error_schemas import ChartGenerationError
mock_find_by_id.return_value = self._mock_chart_with_dataset(chart_id=42)
mock_check_access.return_value = DatasetValidationResult(
is_valid=True, dataset_id=10, dataset_name="ds", warnings=[]
)
mock_validate.return_value = CompileResult(
success=False,
error="Column 'bad_col' does not exist",
error_code="CHART_VALIDATION_FAILED",
tier="validation",
error_obj=ChartGenerationError(
error_type="invalid_column",
message="Column 'bad_col' does not exist",
details="Available: a, b, c",
suggestions=["a"],
error_code="CHART_VALIDATION_FAILED",
),
)
request = {
"identifier": 42,
"generate_preview": False,
"config": {
"chart_type": "table",
"columns": [{"name": "bad_col"}],
},
}
async with Client(mcp) as client:
result = await client.call_tool("update_chart", {"request": request})
assert result.structured_content["success"] is False
error = result.structured_content["error"]
assert error["error_code"] == "CHART_VALIDATION_FAILED"
mock_update_cmd_cls.assert_not_called()
@@ -23,7 +23,9 @@ import importlib
from unittest.mock import Mock, patch
import pytest
from fastmcp import Client
from superset.mcp_service.app import mcp
from superset.mcp_service.chart.schemas import (
AxisConfig,
ColumnRef,
@@ -34,11 +36,48 @@ from superset.mcp_service.chart.schemas import (
XYChartConfig,
)
# The package ``__init__.py`` re-exports the ``update_chart_preview`` tool
# function under the same dotted path as the module, so mock.patch's string
# lookup of ``...update_chart_preview.<attr>`` can resolve to the function on
# some Python versions. Hold a direct module reference for ``patch.object``.
update_chart_preview_module = importlib.import_module(
"superset.mcp_service.chart.tool.update_chart_preview"
)
@pytest.fixture
def mcp_server():
return mcp
@pytest.fixture
def mock_auth():
"""Mock authentication for tool-invocation tests."""
with patch("superset.mcp_service.auth.get_user_from_request") as mock_get_user:
user = Mock()
user.id = 1
user.username = "admin"
mock_get_user.return_value = user
yield mock_get_user
def _mock_dataset(id: int = 1) -> Mock:
"""Mock SqlaTable with the attributes the tool reads."""
column = Mock()
column.column_name = "ds"
column.type = "TIMESTAMP"
database = Mock()
database.database_name = "main"
dataset = Mock()
dataset.id = id
dataset.table_name = "birth_names"
dataset.schema = None
dataset.columns = [column]
dataset.metrics = []
dataset.database = database
return dataset
class TestUpdateChartPreview:
"""Tests for update_chart_preview MCP tool."""
@@ -528,6 +567,9 @@ class TestUpdateChartPreview:
assert result is None
@patch.object(update_chart_preview_module, "validate_and_compile")
@patch.object(update_chart_preview_module, "has_dataset_access", return_value=True)
@patch("superset.daos.dataset.DatasetDAO.find_by_id")
@patch.object(update_chart_preview_module, "analyze_chart_semantics")
@patch.object(update_chart_preview_module, "analyze_chart_capabilities")
@patch.object(update_chart_preview_module, "generate_explore_link")
@@ -541,11 +583,16 @@ class TestUpdateChartPreview:
mock_generate_explore_link,
mock_analyze_chart_capabilities,
mock_analyze_chart_semantics,
mock_find_by_id,
unused_access_mock,
mock_validate_and_compile,
) -> None:
"""Invalid previous form_data_key is warning-only for preview updates."""
mock_user = Mock()
mock_user.id = 1
mock_get_user_from_request.return_value = mock_user
mock_find_by_id.return_value = _mock_dataset(id=3)
mock_validate_and_compile.return_value = Mock(success=True)
mock_get_previous_form_data.return_value = None
mock_generate_explore_link.return_value = (
"http://localhost:8088/explore/?form_data_key=new_preview_key"
@@ -581,6 +628,9 @@ class TestUpdateChartPreview:
]
mock_get_previous_form_data.assert_called_once_with("nonexistent_key_12345")
@patch.object(update_chart_preview_module, "validate_and_compile")
@patch.object(update_chart_preview_module, "has_dataset_access", return_value=True)
@patch("superset.daos.dataset.DatasetDAO.find_by_id")
@patch.object(update_chart_preview_module, "analyze_chart_semantics")
@patch.object(update_chart_preview_module, "analyze_chart_capabilities")
@patch.object(update_chart_preview_module, "generate_explore_link")
@@ -594,11 +644,16 @@ class TestUpdateChartPreview:
mock_generate_explore_link,
mock_analyze_chart_capabilities,
mock_analyze_chart_semantics,
mock_find_by_id,
unused_access_mock,
mock_validate_and_compile,
) -> None:
"""Valid previous form_data preserves filters without a cache warning."""
mock_user = Mock()
mock_user.id = 1
mock_get_user_from_request.return_value = mock_user
mock_find_by_id.return_value = _mock_dataset(id=3)
mock_validate_and_compile.return_value = Mock(success=True)
cached_adhoc_filters = [
{
"clause": "WHERE",
@@ -642,3 +697,101 @@ class TestUpdateChartPreview:
assert result["error"] is None
assert result["warnings"] == []
mock_get_previous_form_data.assert_called_once_with("valid_key_12345")
class TestUpdateChartPreviewValidation:
"""Tier-1 validation gate and dataset access checks."""
@patch.object(update_chart_preview_module, "has_dataset_access", return_value=True)
@patch("superset.daos.dataset.DatasetDAO.find_by_id")
@patch.object(update_chart_preview_module, "validate_and_compile")
@patch(
"superset.mcp_service.commands.create_form_data.MCPCreateFormDataCommand.run"
)
@pytest.mark.asyncio
async def test_validation_failure_skips_cache_write(
self,
mock_create_form_data,
mock_validate,
mock_find_dataset,
unused_access_mock,
mcp_server,
mock_auth,
):
"""Bad column ref → structured error with suggestions, no cache write."""
from superset.mcp_service.chart.compile import CompileResult
from superset.mcp_service.common.error_schemas import ChartGenerationError
mock_find_dataset.return_value = _mock_dataset(id=3)
mock_validate.return_value = CompileResult(
success=False,
error="Column 'num_boys' does not exist in dataset",
error_code="CHART_VALIDATION_FAILED",
tier="validation",
error_obj=ChartGenerationError(
error_type="invalid_column",
message="Column 'num_boys' does not exist in dataset",
details="Available columns: ds, gender, name, num, sum_boys",
suggestions=["sum_boys"],
error_code="CHART_VALIDATION_FAILED",
),
)
config = XYChartConfig(
chart_type="xy",
x=ColumnRef(name="ds"),
y=[ColumnRef(name="num_boys", aggregate="SUM")],
kind="line",
)
request = UpdateChartPreviewRequest(
form_data_key="prev_key", dataset_id="3", config=config
)
async with Client(mcp_server) as client:
result = await client.call_tool(
"update_chart_preview", {"request": request.model_dump()}
)
assert result.data["success"] is False
assert result.data["chart"] is None
error = result.data["error"]
assert isinstance(error, dict)
assert error["error_code"] == "CHART_VALIDATION_FAILED"
assert "sum_boys" in error["suggestions"]
mock_create_form_data.assert_not_called()
@patch.object(update_chart_preview_module, "has_dataset_access", return_value=False)
@patch("superset.daos.dataset.DatasetDAO.find_by_id")
@patch(
"superset.mcp_service.commands.create_form_data.MCPCreateFormDataCommand.run"
)
@pytest.mark.asyncio
async def test_dataset_access_denied_short_circuits(
self,
mock_create_form_data,
mock_find_dataset,
unused_access_mock,
mcp_server,
mock_auth,
):
"""has_dataset_access=False → DatasetNotAccessible, no cache write."""
mock_find_dataset.return_value = _mock_dataset(id=3)
config = TableChartConfig(
chart_type="table", columns=[ColumnRef(name="region")]
)
request = UpdateChartPreviewRequest(
form_data_key="prev_key", dataset_id="3", config=config
)
async with Client(mcp_server) as client:
result = await client.call_tool(
"update_chart_preview", {"request": request.model_dump()}
)
assert result.data["success"] is False
assert result.data["chart"] is None
error = result.data["error"]
assert isinstance(error, dict)
assert error["error_type"] == "DatasetNotAccessible"
mock_create_form_data.assert_not_called()
@@ -19,6 +19,7 @@
Comprehensive unit tests for MCP generate_explore_link tool
"""
import importlib
import logging
from unittest.mock import Mock, patch
@@ -37,6 +38,14 @@ from superset.mcp_service.chart.schemas import (
)
from superset.mcp_service.common.error_schemas import DatasetContext
# The package ``__init__.py`` re-exports the ``generate_explore_link`` tool
# function under the same dotted path as the module, so mock.patch's string
# lookup of ``...generate_explore_link.<attr>`` can resolve to the function
# on some Python versions. Hold a direct module reference for ``patch.object``.
generate_explore_link_module = importlib.import_module(
"superset.mcp_service.explore.tool.generate_explore_link"
)
logging.basicConfig(level=logging.DEBUG)
logger = logging.getLogger(__name__)
@@ -57,6 +66,29 @@ def mock_auth():
yield mock_get_user
@pytest.fixture(autouse=True)
def mock_dataset_access_granted():
"""Grant dataset access by default; tests that need a denial override this."""
with patch.object(
generate_explore_link_module, "has_dataset_access", return_value=True
):
yield
@pytest.fixture(autouse=True)
def mock_validation_passes():
"""Skip Tier-1 dataset validation by default so Mock datasets don't trip the
real validator. Individual tests that exercise validation override this."""
from superset.mcp_service.chart.compile import CompileResult
with patch.object(
generate_explore_link_module,
"validate_and_compile",
return_value=CompileResult(success=True),
):
yield
@pytest.fixture(autouse=True)
def mock_webdriver_baseurl(app_context):
"""Mock WEBDRIVER_BASEURL_USER_FRIENDLY for consistent test URLs."""
@@ -922,3 +954,102 @@ class TestGenerateExploreLinkColumnNormalization:
assert result.data["error"] is None
# original names should pass through unchanged
assert result.data["form_data"]["x_axis"] == "orderdate"
class TestGenerateExploreLinkValidation:
"""Tier-1 validation gate (DatasetValidator) and dataset access checks."""
@pytest.fixture(autouse=True)
def mock_validation_passes(self):
"""Override the module-level autouse patch so each test in this class
can stub ``validate_and_compile`` itself. The fixture name MUST match
the module-level fixture for pytest's override-by-name to take effect.
"""
return
@patch.object(generate_explore_link_module, "validate_and_compile")
@patch("superset.daos.dataset.DatasetDAO.find_by_id")
@patch(
"superset.mcp_service.commands.create_form_data.MCPCreateFormDataCommand.run"
)
@pytest.mark.asyncio
async def test_validation_failure_returns_structured_error(
self,
mock_create_form_data,
mock_find_dataset,
mock_validate,
mcp_server,
):
"""Non-existent column → structured ChartGenerationError with suggestions,
and MCPCreateFormDataCommand must NOT be called (no cache write)."""
from superset.mcp_service.chart.compile import CompileResult
from superset.mcp_service.common.error_schemas import ChartGenerationError
mock_find_dataset.return_value = _mock_dataset(id=3)
mock_validate.return_value = CompileResult(
success=False,
error="Column 'num_boys' does not exist in dataset",
error_code="CHART_VALIDATION_FAILED",
tier="validation",
error_obj=ChartGenerationError(
error_type="invalid_column",
message="Column 'num_boys' does not exist in dataset",
details="Available columns: ds, gender, name, num, sum_boys",
suggestions=["sum_boys"],
error_code="CHART_VALIDATION_FAILED",
),
)
config = XYChartConfig(
chart_type="xy",
x=ColumnRef(name="ds"),
y=[ColumnRef(name="num_boys", aggregate="SUM")],
kind="line",
)
request = GenerateExploreLinkRequest(dataset_id="3", config=config)
async with Client(mcp_server) as client:
result = await client.call_tool(
"generate_explore_link", {"request": request.model_dump()}
)
assert result.data["url"] == ""
assert result.data["form_data_key"] is None
error = result.data["error"]
assert isinstance(error, dict)
assert error["error_code"] == "CHART_VALIDATION_FAILED"
assert "sum_boys" in error["suggestions"]
mock_create_form_data.assert_not_called()
@patch.object(
generate_explore_link_module, "has_dataset_access", return_value=False
)
@patch("superset.daos.dataset.DatasetDAO.find_by_id")
@patch(
"superset.mcp_service.commands.create_form_data.MCPCreateFormDataCommand.run"
)
@pytest.mark.asyncio
async def test_dataset_access_denied_short_circuits(
self,
mock_create_form_data,
mock_find_dataset,
unused_access_mock,
mcp_server,
):
"""has_dataset_access=False blocks the tool before any cache write."""
mock_find_dataset.return_value = _mock_dataset(id=3)
config = TableChartConfig(
chart_type="table", columns=[ColumnRef(name="region")]
)
request = GenerateExploreLinkRequest(dataset_id="3", config=config)
async with Client(mcp_server) as client:
result = await client.call_tool(
"generate_explore_link", {"request": request.model_dump()}
)
assert result.data["url"] == ""
# Surface as "not found" rather than leaking that the dataset exists.
assert "Dataset not found" in result.data["error"]
mock_create_form_data.assert_not_called()
@@ -298,7 +298,8 @@ class TestOpenSqlLabWithContext:
field_path=("title",),
)
assert response.error == sanitize_for_llm_context(
"Database with ID 404 not found",
"Database with ID 404 not found."
" Use list_databases to get valid database IDs.",
field_path=("error",),
)
finally:
+138 -2
View File
@@ -19,6 +19,7 @@
from datetime import datetime
import pytest
from flask import current_app
from pytest_mock import MockerFixture
from sqlalchemy import (
Column,
@@ -28,6 +29,7 @@ from sqlalchemy import (
Table as SqlalchemyTable,
)
from sqlalchemy.engine.reflection import Inspector
from sqlalchemy.engine.url import make_url
from sqlalchemy.orm.session import Session
from sqlalchemy.sql import Select
@@ -523,6 +525,60 @@ def test_get_all_materialized_view_names_in_schema_needs_oauth2(
assert excinfo.value.error.error_type == SupersetErrorType.OAUTH2_REDIRECT
def test_get_sqla_engine(mocker: MockerFixture) -> None:
"""
Test `_get_sqla_engine`.
"""
from superset.models.core import Database
user = mocker.MagicMock()
user.email = "alice.doe@example.org"
mocker.patch(
"superset.models.core.security_manager.find_user",
return_value=user,
)
mocker.patch("superset.models.core.get_username", return_value="alice")
create_engine = mocker.patch("superset.models.core.create_engine")
database = Database(database_name="my_db", sqlalchemy_uri="trino://")
database._get_sqla_engine(nullpool=False)
create_engine.assert_called_with(
make_url("trino:///"),
connect_args={"source": "Apache Superset"},
)
def test_get_sqla_engine_user_impersonation(mocker: MockerFixture) -> None:
"""
Test user impersonation in `_get_sqla_engine`.
"""
from superset.models.core import Database
user = mocker.MagicMock()
user.email = "alice.doe@example.org"
mocker.patch(
"superset.models.core.security_manager.find_user",
return_value=user,
)
mocker.patch("superset.models.core.get_username", return_value="alice")
create_engine = mocker.patch("superset.models.core.create_engine")
database = Database(
database_name="my_db",
sqlalchemy_uri="trino://",
impersonate_user=True,
)
database._get_sqla_engine(nullpool=False)
create_engine.assert_called_with(
make_url("trino:///"),
connect_args={"user": "alice", "source": "Apache Superset"},
)
def test_add_database_to_signature():
args = ["param1", "param2"]
@@ -548,6 +604,36 @@ def test_add_database_to_signature():
assert args3 == ["param1", "param2", database]
@with_feature_flags(IMPERSONATE_WITH_EMAIL_PREFIX=True)
def test_get_sqla_engine_user_impersonation_email(mocker: MockerFixture) -> None:
"""
Test user impersonation in `_get_sqla_engine` with `username_from_email`.
"""
from superset.models.core import Database
user = mocker.MagicMock()
user.email = "alice.doe@example.org"
mocker.patch(
"superset.models.core.security_manager.find_user",
return_value=user,
)
mocker.patch("superset.models.core.get_username", return_value="alice")
create_engine = mocker.patch("superset.models.core.create_engine")
database = Database(
database_name="my_db",
sqlalchemy_uri="trino://",
impersonate_user=True,
)
database._get_sqla_engine(nullpool=False)
create_engine.assert_called_with(
make_url("trino:///"),
connect_args={"user": "alice.doe", "source": "Apache Superset"},
)
def test_is_oauth2_enabled() -> None:
"""
Test the `is_oauth2_enabled` method.
@@ -689,8 +775,8 @@ def test_raw_connection_oauth_engine(mocker: MockerFixture) -> None:
encrypted_extra=json.dumps(oauth2_client_info),
)
database.db_engine_spec.oauth2_exception = OAuth2Error
create_engine = mocker.patch("superset.engines.manager.create_engine")
create_engine.side_effect = OAuth2Error("OAuth2 required")
_get_sqla_engine = mocker.patch.object(database, "_get_sqla_engine")
_get_sqla_engine.side_effect = OAuth2Error("OAuth2 required")
with pytest.raises(OAuth2RedirectError) as excinfo:
with database.get_raw_connection() as conn:
@@ -793,6 +879,56 @@ def test_get_schema_access_for_file_upload() -> None:
assert database.get_schema_access_for_file_upload() == {"public"}
def test_engine_context_manager(mocker: MockerFixture, app_context: None) -> None:
"""
Test the engine context manager.
"""
from unittest.mock import MagicMock
engine_context_manager = MagicMock()
mocker.patch.dict(
current_app.config,
{"ENGINE_CONTEXT_MANAGER": engine_context_manager},
)
_get_sqla_engine = mocker.patch.object(Database, "_get_sqla_engine")
database = Database(database_name="my_db", sqlalchemy_uri="trino://")
with database.get_sqla_engine("catalog", "schema"):
pass
engine_context_manager.assert_called_once_with(database, "catalog", "schema")
engine_context_manager().__enter__.assert_called_once()
engine_context_manager().__exit__.assert_called_once_with(None, None, None)
_get_sqla_engine.assert_called_once_with(
catalog="catalog",
schema="schema",
nullpool=True,
source=None,
sqlalchemy_uri="trino://",
)
def test_engine_oauth2(mocker: MockerFixture) -> None:
"""
Test that we handle OAuth2 when `create_engine` fails.
"""
database = Database(database_name="my_db", sqlalchemy_uri="trino://")
mocker.patch.object(database, "_get_sqla_engine", side_effect=Exception)
mocker.patch.object(database, "is_oauth2_enabled", return_value=True)
mocker.patch.object(database.db_engine_spec, "needs_oauth2", return_value=True)
start_oauth2_dance = mocker.patch.object(
database.db_engine_spec,
"start_oauth2_dance",
side_effect=OAuth2Error("OAuth2 required"),
)
with pytest.raises(OAuth2Error):
with database.get_sqla_engine("catalog", "schema"):
pass
start_oauth2_dance.assert_called_with(database)
def test_purge_oauth2_tokens(session: Session) -> None:
"""
Test the `purge_oauth2_tokens` method.
+229
View File
@@ -1937,6 +1937,235 @@ def test_extras_having_is_parenthesized(
)
def test_calculated_column_filter_is_parenthesized(
database: Database,
) -> None:
"""
Test that calculated column expressions containing OR are wrapped in
parentheses when used in WHERE filters.
Without parentheses, a calculated column expression like
``status = 'active' OR status = 'pending'`` combined with other filters
via AND would produce unexpected evaluation order due to SQL operator
precedence (AND binds tighter than OR), potentially dropping time range
and other filters. Same class of bug as fixed in PR #38183 for
extras.where/having, but on the calculated column filter path.
"""
from superset.connectors.sqla.models import SqlaTable, TableColumn
table = SqlaTable(
database=database,
schema=None,
table_name="t",
columns=[
TableColumn(column_name="a", type="INTEGER"),
TableColumn(
column_name="is_active",
expression="status = 'active' OR status = 'pending'",
type="BOOLEAN",
),
],
)
sqla_query = table.get_sqla_query(
columns=["a"],
filter=[
{
"col": "is_active",
"op": "IS TRUE",
"val": None,
},
],
extras={},
is_timeseries=False,
metrics=[],
)
with database.get_sqla_engine() as engine:
sql = str(
sqla_query.sqla_query.compile(
dialect=engine.dialect,
compile_kwargs={"literal_binds": True},
)
)
assert "(status = 'active' OR status = 'pending')" in sql, (
f"Calculated column expression should be wrapped in parentheses. "
f"Generated SQL: {sql}"
)
def test_calculated_column_nested_or_and_is_parenthesized(
database: Database,
) -> None:
"""
Test that calculated column expressions with nested OR/AND combinations
are correctly parenthesized as a single unit in WHERE filters.
"""
from superset.connectors.sqla.models import SqlaTable, TableColumn
table = SqlaTable(
database=database,
schema=None,
table_name="t",
columns=[
TableColumn(column_name="a", type="INTEGER"),
TableColumn(
column_name="is_target",
expression=(
"(status = 'active' AND region = 'US') "
"OR (status = 'pending' AND region = 'EU')"
),
type="BOOLEAN",
),
],
)
sqla_query = table.get_sqla_query(
columns=["a"],
filter=[
{
"col": "is_target",
"op": "IS TRUE",
"val": None,
},
],
extras={},
is_timeseries=False,
metrics=[],
)
with database.get_sqla_engine() as engine:
sql = str(
sqla_query.sqla_query.compile(
dialect=engine.dialect,
compile_kwargs={"literal_binds": True},
)
)
assert (
"((status = 'active' AND region = 'US') "
"OR (status = 'pending' AND region = 'EU'))"
) in sql, (
f"Nested OR/AND expression should be wrapped in parentheses. "
f"Generated SQL: {sql}"
)
def test_calculated_column_non_boolean_filter_is_parenthesized(
database: Database,
) -> None:
"""
Test that non-boolean calculated column expressions are parenthesized
when used with IN filters.
"""
from superset.connectors.sqla.models import SqlaTable, TableColumn
table = SqlaTable(
database=database,
schema=None,
table_name="t",
columns=[
TableColumn(column_name="a", type="INTEGER"),
TableColumn(
column_name="full_name",
expression="first_name || ' ' || last_name",
type="TEXT",
),
],
)
sqla_query = table.get_sqla_query(
columns=["a"],
filter=[
{
"col": "full_name",
"op": "IN",
"val": ["John Doe", "Jane Doe"],
},
],
extras={},
is_timeseries=False,
metrics=[],
)
with database.get_sqla_engine() as engine:
sql = str(
sqla_query.sqla_query.compile(
dialect=engine.dialect,
compile_kwargs={"literal_binds": True},
)
)
assert "(first_name || ' ' || last_name)" in sql, (
f"Non-boolean calculated column should be wrapped in parentheses. "
f"Generated SQL: {sql}"
)
def test_multiple_calculated_columns_each_parenthesized(
database: Database,
) -> None:
"""
Test that multiple calculated columns used as filters are each
independently wrapped in parentheses.
"""
from superset.connectors.sqla.models import SqlaTable, TableColumn
table = SqlaTable(
database=database,
schema=None,
table_name="t",
columns=[
TableColumn(column_name="a", type="INTEGER"),
TableColumn(
column_name="is_active",
expression="status = 'active' OR status = 'pending'",
type="BOOLEAN",
),
TableColumn(
column_name="is_premium",
expression="tier = 'gold' OR tier = 'platinum'",
type="BOOLEAN",
),
],
)
sqla_query = table.get_sqla_query(
columns=["a"],
filter=[
{
"col": "is_active",
"op": "IS TRUE",
"val": None,
},
{
"col": "is_premium",
"op": "IS TRUE",
"val": None,
},
],
extras={},
is_timeseries=False,
metrics=[],
)
with database.get_sqla_engine() as engine:
sql = str(
sqla_query.sqla_query.compile(
dialect=engine.dialect,
compile_kwargs={"literal_binds": True},
)
)
assert "(status = 'active' OR status = 'pending')" in sql, (
f"First calculated column should be parenthesized. Generated SQL: {sql}"
)
assert "(tier = 'gold' OR tier = 'platinum')" in sql, (
f"Second calculated column should be parenthesized. Generated SQL: {sql}"
)
def _run_probe(
database: Database,
type_probe_needs_row: bool = False,
@@ -202,6 +202,7 @@ def setup_mock_raw_connection(
def _raw_connection(
catalog: str | None = None,
schema: str | None = None,
nullpool: bool = True,
source: Any | None = None,
):
yield mock_connection
@@ -396,3 +396,127 @@ def test_transpile_unknown_source_engine_uses_generic() -> None:
"SELECT * FROM orders", "postgresql", "unknown_engine"
)
assert result == "SELECT * FROM orders"
# Tests for identify=True (identifier quoting)
@pytest.mark.parametrize(
"sql,dialect,expected",
[
# PostgreSQL - double-quoted identifiers
(
"STATE ILIKE '%AL%'",
"postgresql",
"\"STATE\" ILIKE '%AL%'",
),
# MySQL - backtick-quoted identifiers, ILIKE transpiled
(
"STATE ILIKE '%AL%'",
"mysql",
"LOWER(`STATE`) LIKE LOWER('%AL%')",
),
# BigQuery - backtick-quoted identifiers, ILIKE transpiled
(
"STATE ILIKE '%AL%'",
"bigquery",
"LOWER(`STATE`) LIKE LOWER('%AL%')",
),
# Snowflake - double-quoted identifiers
(
"STATE ILIKE '%AL%'",
"snowflake",
"\"STATE\" ILIKE '%AL%'",
),
# MSSQL - bracket-quoted identifiers
(
"STATE = 'CA'",
"mssql",
"[STATE] = 'CA'",
),
# Compound filter with multiple identifiers
(
"STATE = 'CA' AND AIRLINE = 'Delta'",
"postgresql",
"\"STATE\" = 'CA' AND \"AIRLINE\" = 'Delta'",
),
# Lowercase identifiers also get quoted
(
"name = 'test'",
"postgresql",
"\"name\" = 'test'",
),
],
)
def test_identify_quotes_identifiers(sql: str, dialect: str, expected: str) -> None:
"""Test that identify=True quotes identifiers per target dialect."""
assert transpile_to_dialect(sql, dialect, identify=True) == expected
def test_identify_unknown_engine_returns_unchanged() -> None:
"""Test that identify=True has no effect on unknown engines."""
sql = "STATE = 'CA'"
assert transpile_to_dialect(sql, "unknown_engine", identify=True) == sql
@pytest.mark.parametrize(
"sql,engine,expected",
[
(
"STATE ILIKE '%AL%'",
"postgresql",
"\"STATE\" ILIKE '%AL%'",
),
(
"country ILIKE '%Italy%'",
"bigquery",
"LOWER(`country`) LIKE LOWER('%Italy%')",
),
],
)
def test_identify_with_source_engine(sql: str, engine: str, expected: str) -> None:
"""Test identify=True with source_engine matching target engine."""
result = transpile_to_dialect(sql, engine, source_engine=engine, identify=True)
assert result == expected
@pytest.mark.parametrize(
"engine",
["postgresql", "bigquery", "mysql", "snowflake"],
)
def test_identify_transpilation_is_idempotent(engine: str) -> None:
"""Test that transpiling twice produces the same result (idempotent).
This matters because _sanitize_filters() can be called multiple times
via validate().
"""
clause = "STATE ILIKE '%AL%'"
pass1 = transpile_to_dialect(clause, engine, source_engine=engine, identify=True)
pass2 = transpile_to_dialect(pass1, engine, source_engine=engine, identify=True)
assert pass1 == pass2
def test_sanitize_filters_writes_back_transpiled_clause() -> None:
"""Test that _sanitize_filters always persists the transpiled clause.
Regression test: a previous conditional `if sanitized_clause != clause`
skipped the write-back when transpile_to_dialect had already modified
the clause, leaving the original unquoted value in extras.
"""
from unittest.mock import MagicMock
from superset.common.query_object import QueryObject
mock_datasource = MagicMock()
mock_datasource.database.db_engine_spec.engine = "postgresql"
query_obj = QueryObject(
datasource=mock_datasource,
columns=["STATE"],
metrics=[],
extras={
"where": "STATE = 'CA'",
"transpile_to_dialect": True,
},
)
query_obj.validate()
assert '"STATE"' in query_obj.extras["where"]
+173 -1
View File
@@ -21,7 +21,12 @@ from flask_babel import lazy_gettext as _
from superset.commands.chart.exceptions import ChartDataQueryFailedError
from superset.errors import ErrorLevel, SupersetError, SupersetErrorType
from superset.exceptions import SupersetErrorException, SupersetErrorsException
from superset.exceptions import (
OAuth2RedirectError,
SupersetErrorException,
SupersetErrorsException,
SupersetVizException,
)
@mock.patch("superset.tasks.async_queries.security_manager")
@@ -149,3 +154,170 @@ def test_load_chart_data_into_cache_with_superset_errors_exception(
assert errors[1]["message"] == "Table not found"
assert errors[1]["error_type"] == SupersetErrorType.TABLE_DOES_NOT_EXIST_ERROR
assert errors[1]["level"] == ErrorLevel.WARNING
@mock.patch("superset.tasks.async_queries.security_manager")
@mock.patch("superset.tasks.async_queries.async_query_manager")
@mock.patch("superset.tasks.async_queries.get_viz")
@mock.patch("superset.tasks.async_queries.get_datasource_info")
def test_load_explore_json_into_cache_preserves_oauth2_redirect_error(
mock_get_datasource_info,
mock_get_viz,
mock_async_query_manager,
mock_security_manager,
):
"""
OAuth2RedirectError raised by ``viz_obj.get_payload`` must reach the async
job's errors list as a structured SIP-40 envelope so the frontend can
render the OAuth2 banner identically to the sync legacy path.
"""
from superset.tasks.async_queries import load_explore_json_into_cache
job_metadata = {"user_id": 1}
form_data: dict = {}
mock_get_datasource_info.return_value = (1, "table")
mock_security_manager.get_user_by_id.return_value = mock.MagicMock()
mock_async_query_manager.STATUS_ERROR = "error"
viz_obj = mock.MagicMock()
viz_obj.get_payload.side_effect = OAuth2RedirectError(
url="https://accounts.example.com/o/oauth2/v2/auth?...",
tab_id="tab-123",
redirect_uri="https://superset.example.com/oauth2/redirect",
)
mock_get_viz.return_value = viz_obj
with pytest.raises(OAuth2RedirectError):
load_explore_json_into_cache(job_metadata, form_data)
call_args = mock_async_query_manager.update_job.call_args
assert call_args[0] == (job_metadata, "error")
errors = call_args[1]["errors"]
assert len(errors) == 1
assert errors[0]["error_type"] == SupersetErrorType.OAUTH2_REDIRECT
assert errors[0]["extra"] == {
"url": "https://accounts.example.com/o/oauth2/v2/auth?...",
"tab_id": "tab-123",
"redirect_uri": "https://superset.example.com/oauth2/redirect",
}
@mock.patch("superset.tasks.async_queries.security_manager")
@mock.patch("superset.tasks.async_queries.async_query_manager")
@mock.patch("superset.tasks.async_queries.get_viz")
@mock.patch("superset.tasks.async_queries.get_datasource_info")
def test_load_explore_json_into_cache_preserves_superset_errors_exception(
mock_get_datasource_info,
mock_get_viz,
mock_async_query_manager,
mock_security_manager,
):
"""SupersetErrorsException must be preserved as a list of SIP-40 dicts."""
from superset.tasks.async_queries import load_explore_json_into_cache
job_metadata = {"user_id": 1}
form_data: dict = {}
mock_get_datasource_info.return_value = (1, "table")
mock_security_manager.get_user_by_id.return_value = mock.MagicMock()
mock_async_query_manager.STATUS_ERROR = "error"
viz_obj = mock.MagicMock()
viz_obj.get_payload.side_effect = SupersetErrorsException(
[
SupersetError(
message="Column not found",
error_type=SupersetErrorType.COLUMN_DOES_NOT_EXIST_ERROR,
level=ErrorLevel.ERROR,
),
SupersetError(
message="Table not found",
error_type=SupersetErrorType.TABLE_DOES_NOT_EXIST_ERROR,
level=ErrorLevel.WARNING,
),
]
)
mock_get_viz.return_value = viz_obj
with pytest.raises(SupersetErrorsException):
load_explore_json_into_cache(job_metadata, form_data)
errors = mock_async_query_manager.update_job.call_args[1]["errors"]
assert len(errors) == 2
assert errors[0]["error_type"] == SupersetErrorType.COLUMN_DOES_NOT_EXIST_ERROR
assert errors[1]["error_type"] == SupersetErrorType.TABLE_DOES_NOT_EXIST_ERROR
@mock.patch("superset.tasks.async_queries.security_manager")
@mock.patch("superset.tasks.async_queries.async_query_manager")
@mock.patch("superset.tasks.async_queries.get_viz")
@mock.patch("superset.tasks.async_queries.get_datasource_info")
def test_load_explore_json_into_cache_preserves_superset_viz_exception(
mock_get_datasource_info,
mock_get_viz,
mock_async_query_manager,
mock_security_manager,
):
"""
Test that SupersetVizException passes ``ex.errors`` straight through.
"""
from superset.tasks.async_queries import load_explore_json_into_cache
job_metadata = {"user_id": 1}
form_data: dict = {}
mock_get_datasource_info.return_value = (1, "table")
mock_security_manager.get_user_by_id.return_value = mock.MagicMock()
mock_async_query_manager.STATUS_ERROR = "error"
payload_errors = [
{
"message": "Bad column",
"error_type": SupersetErrorType.VIZ_GET_DF_ERROR,
"level": ErrorLevel.ERROR,
}
]
viz_obj = mock.MagicMock()
viz_obj.get_payload.return_value = {"errors": payload_errors}
viz_obj.has_error.return_value = True
mock_get_viz.return_value = viz_obj
with pytest.raises(SupersetVizException):
load_explore_json_into_cache(job_metadata, form_data)
errors = mock_async_query_manager.update_job.call_args[1]["errors"]
assert errors == payload_errors
@mock.patch("superset.tasks.async_queries.security_manager")
@mock.patch("superset.tasks.async_queries.async_query_manager")
@mock.patch("superset.tasks.async_queries.get_viz")
@mock.patch("superset.tasks.async_queries.get_datasource_info")
def test_load_explore_json_into_cache_falls_back_to_string_for_generic_exception(
mock_get_datasource_info,
mock_get_viz,
mock_async_query_manager,
mock_security_manager,
):
"""
Test that Non-Superset exception are passed as plain-string error.
"""
from superset.tasks.async_queries import load_explore_json_into_cache
job_metadata = {"user_id": 1}
form_data: dict = {}
mock_get_datasource_info.return_value = (1, "table")
mock_security_manager.get_user_by_id.return_value = mock.MagicMock()
mock_async_query_manager.STATUS_ERROR = "error"
viz_obj = mock.MagicMock()
viz_obj.get_payload.side_effect = RuntimeError("boom")
mock_get_viz.return_value = viz_obj
with pytest.raises(RuntimeError):
load_explore_json_into_cache(job_metadata, form_data)
errors = mock_async_query_manager.update_job.call_args[1]["errors"]
assert errors == ["boom"]
+111
View File
@@ -0,0 +1,111 @@
# Licensed to the Apache Software Foundation (ASF) under one
# or more contributor license agreements. See the NOTICE file
# distributed with this work for additional information
# regarding copyright ownership. The ASF licenses this file
# to you under the Apache License, Version 2.0 (the
# "License"); you may not use this file except in compliance
# with the License. You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing,
# software distributed under the License is distributed on an
# "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
# KIND, either express or implied. See the License for the
# specific language governing permissions and limitations
# under the License.
from typing import Any
from unittest.mock import patch
import pytest
from superset import viz
from superset.common.db_query_status import QueryStatus
from superset.connectors.sqla.models import SqlaTable
from superset.errors import SupersetErrorType
from superset.exceptions import (
OAuth2RedirectError,
QueryObjectValidationError,
)
from superset.models.core import Database
QUERY_OBJ: dict[str, Any] = {"row_limit": 100, "from_dttm": None, "to_dttm": None}
def _viz() -> viz.BaseViz:
database = Database(database_name="d", sqlalchemy_uri="sqlite://")
datasource = SqlaTable(
table_name="t",
columns=[],
metrics=[],
main_dttm_col=None,
database=database,
)
# ``force=True`` skips the data cache lookup branch so ``get_df`` is always
# invoked, which is what we want to assert error-handling against.
return viz.BaseViz(
datasource=datasource,
form_data={"viz_type": "table"},
force=True,
)
def test_get_df_payload_propagates_oauth2_redirect_error() -> None:
"""
OAuth2RedirectError (a SupersetErrorException) must propagate out of
``get_df_payload`` so the global Flask error handler can serialize it.
"""
obj = _viz()
oauth_exc = OAuth2RedirectError(
url="https://accounts.example.com/o/oauth2/v2/auth?...",
tab_id="tab-123",
redirect_uri="https://superset.example.com/oauth2/redirect",
)
with patch.object(viz.BaseViz, "get_df", side_effect=oauth_exc):
with pytest.raises(OAuth2RedirectError) as exc_info:
obj.get_df_payload(QUERY_OBJ)
assert exc_info.value.error.error_type == SupersetErrorType.OAUTH2_REDIRECT
assert exc_info.value.error.extra == {
"url": "https://accounts.example.com/o/oauth2/v2/auth?...",
"tab_id": "tab-123",
"redirect_uri": "https://superset.example.com/oauth2/redirect",
}
def test_get_df_payload_captures_generic_exception_as_viz_get_df_error() -> None:
"""
Non-Superset exception raised by ``get_df`` are downgraded to a
``VIZ_GET_DF_ERROR`` entry on ``self.errors``.
"""
obj = _viz()
with patch.object(viz.BaseViz, "get_df", side_effect=RuntimeError("boom")):
payload = obj.get_df_payload(QUERY_OBJ)
assert obj.status == QueryStatus.FAILED
assert payload["status"] == QueryStatus.FAILED
assert len(obj.errors) == 1
assert obj.errors[0]["error_type"] == SupersetErrorType.VIZ_GET_DF_ERROR
assert obj.errors[0]["message"] == "boom"
def test_get_df_payload_captures_query_object_validation_error() -> None:
"""
``QueryObjectValidationError`` is reported as ``VIZ_GET_DF_ERROR``.
"""
obj = _viz()
with patch.object(
viz.BaseViz,
"get_df",
side_effect=QueryObjectValidationError("bad query"),
):
payload = obj.get_df_payload(QUERY_OBJ)
assert obj.status == QueryStatus.FAILED
assert payload["status"] == QueryStatus.FAILED
assert len(obj.errors) == 1
assert obj.errors[0]["error_type"] == SupersetErrorType.VIZ_GET_DF_ERROR
assert obj.errors[0]["message"] == "bad query"