Files
superset2/superset/common/query_actions.py
T

383 lines
14 KiB
Python

# 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 __future__ import annotations
import copy
import time
from typing import Any, Callable, TYPE_CHECKING
from flask_babel import _
from superset.common.chart_data import ChartDataResultType
from superset.common.chart_data_timing import (
QueryAcquisitionTiming,
QueryDataResult,
QueryTiming,
)
from superset.common.db_query_status import QueryStatus
from superset.exceptions import QueryObjectValidationError, SupersetParseError
from superset.explorables.base import Explorable
from superset.utils.core import (
extract_column_dtype,
extract_dataframe_dtypes,
ExtraFiltersReasonType,
get_column_name,
get_time_filter_status,
)
from superset.utils.currency import (
detect_currency,
detect_currency_from_df,
has_auto_currency_in_column_config,
)
if TYPE_CHECKING:
from superset.common.query_context import QueryContext
from superset.common.query_object import QueryObject
def _get_datasource(query_context: QueryContext, query_obj: QueryObject) -> Explorable:
return query_obj.datasource or query_context.datasource
def _get_columns(
query_context: QueryContext, query_obj: QueryObject, _: bool
) -> dict[str, Any]:
datasource = _get_datasource(query_context, query_obj)
return {
"data": [
{
"column_name": col.column_name,
"verbose_name": col.verbose_name,
"dtype": extract_column_dtype(col),
}
for col in datasource.columns
]
}
def _get_timegrains(
query_context: QueryContext, query_obj: QueryObject, _: bool
) -> dict[str, Any]:
datasource = _get_datasource(query_context, query_obj)
# Use the new get_time_grains() method from Explorable protocol
grains = datasource.get_time_grains()
return {"data": grains}
def _get_query(
query_context: QueryContext,
query_obj: QueryObject,
_: bool,
) -> dict[str, Any]:
datasource = _get_datasource(query_context, query_obj)
result = {"language": datasource.query_language}
try:
result["query"] = datasource.get_query_str(query_obj.to_dict())
except QueryObjectValidationError as err:
# Validation errors (missing required fields, invalid config)
# No SQL was generated
result["error"] = err.message
except SupersetParseError as err:
# Parsing errors (SQL optimization/parsing failed)
# SQL was generated but couldn't be optimized - show both
if err.error.extra and (sql := err.error.extra.get("sql")) is not None:
result["query"] = sql
result["error"] = err.error.message
return result
def _detect_currency(
query_context: QueryContext,
query_obj: QueryObject,
datasource: Explorable,
df: Any = None,
) -> str | None:
"""
Detect currency from filtered data for AUTO mode currency formatting.
First attempts to detect from the provided dataframe if the currency
column is present. Falls back to a separate query only when needed.
Checks both top-level currency_format (used by Pie, Timeseries, etc.)
and column_config (used by Table charts) for AUTO currency settings.
:param query_context: The query context with form_data containing currency_format
:param query_obj: The original query object with filters
:param datasource: The datasource being queried
:param df: Optional dataframe to detect currency from (avoids extra query)
:return: ISO 4217 currency code (e.g., "USD") or None
"""
form_data = query_context.form_data or {}
# Check top-level currency_format (for Pie, Timeseries, etc.)
currency_format = form_data.get("currency_format", {})
top_level_auto = (
isinstance(currency_format, dict) and currency_format.get("symbol") == "AUTO"
)
# Check column_config (for Table charts)
column_config_auto = has_auto_currency_in_column_config(form_data)
# Only detect if AUTO is configured somewhere
if not top_level_auto and not column_config_auto:
return None
currency_column = getattr(datasource, "currency_code_column", None)
if not currency_column:
return None
if df is not None and currency_column in df.columns:
return detect_currency_from_df(df, currency_column)
return detect_currency(
datasource=datasource,
filters=query_obj.filter,
granularity=query_obj.granularity,
from_dttm=query_obj.from_dttm,
to_dttm=query_obj.to_dttm,
extras=query_obj.extras,
)
def _get_full_with_timing(
query_context: QueryContext,
query_obj: QueryObject,
force_cached: bool | None = False,
) -> tuple[dict[str, Any], QueryAcquisitionTiming, int]:
acquired = query_context.get_df_payload_result(query_obj, force_cached=force_cached)
payload_assembly_start_ns = time.perf_counter_ns()
payload = _materialize_full_payload(query_context, query_obj, acquired.payload)
payload_assembly_ns = max(0, time.perf_counter_ns() - payload_assembly_start_ns)
return payload, acquired.timing, payload_assembly_ns
def _materialize_full_payload(
query_context: QueryContext,
query_obj: QueryObject,
payload: dict[str, Any],
) -> dict[str, Any]:
datasource = _get_datasource(query_context, query_obj)
result_type = query_obj.result_type or query_context.result_type
df = payload["df"]
status = payload["status"]
if status != QueryStatus.FAILED:
payload["colnames"] = list(df.columns)
payload["indexnames"] = list(df.index)
payload["coltypes"] = extract_dataframe_dtypes(df, datasource)
payload["data"] = query_context.get_data(df, payload["coltypes"])
payload["result_format"] = query_context.result_format
payload["detected_currency"] = _detect_currency(
query_context, query_obj, datasource, df
)
del payload["df"]
applied_time_columns, rejected_time_columns = get_time_filter_status(
datasource, query_obj.applied_time_extras
)
applied_filter_columns = payload.get("applied_filter_columns", [])
rejected_filter_columns = payload.get("rejected_filter_columns", [])
del payload["applied_filter_columns"]
del payload["rejected_filter_columns"]
payload["applied_filters"] = [
{"column": get_column_name(col)} for col in applied_filter_columns
] + applied_time_columns
payload["rejected_filters"] = [
{
"reason": ExtraFiltersReasonType.COL_NOT_IN_DATASOURCE,
"column": get_column_name(col),
}
for col in rejected_filter_columns
] + rejected_time_columns
if result_type == ChartDataResultType.RESULTS and status != QueryStatus.FAILED:
return {
"data": payload.get("data"),
"colnames": payload.get("colnames"),
"coltypes": payload.get("coltypes"),
"rowcount": payload.get("rowcount"),
"sql_rowcount": payload.get("sql_rowcount"),
"detected_currency": payload.get("detected_currency"),
}
return payload
def _prepare_samples_query(
query_context: QueryContext,
query_obj: QueryObject,
) -> QueryObject:
datasource = _get_datasource(query_context, query_obj)
query_obj = copy.copy(query_obj)
query_obj.is_timeseries = False
query_obj.orderby = []
query_obj.metrics = None
query_obj.post_processing = []
# Mark the query as system-authored sampling so query generation may apply
# the engine's bounded-read override (physical datasets only). The shallow
# copy above shares the extras dict with the original query object, so
# build a new dict instead of mutating in place.
query_obj.extras = {**(query_obj.extras or {}), "system_sampling": True}
qry_obj_cols = []
for o in datasource.columns:
if isinstance(o, dict):
if column_name := o.get("column_name"):
qry_obj_cols.append(column_name)
else:
qry_obj_cols.append(o.column_name)
query_obj.columns = qry_obj_cols
query_obj.from_dttm = None
query_obj.to_dttm = None
return query_obj
def _prepare_drill_detail_query(
query_context: QueryContext,
query_obj: QueryObject,
) -> QueryObject:
# todo(yongjie): Remove this function,
# when determining whether samples should be applied to the time filter.
datasource = _get_datasource(query_context, query_obj)
# Refuse for datasource types that don't model raw rows (e.g. semantic
# views). Mirrors the ``supports_samples`` gate on the ``/samples``
# endpoint so drill-detail is hard-blocked on the backend, not just
# hidden in the frontend menu. Defaults to ``True`` for any datasource
# class that doesn't explicitly opt out.
if not getattr(datasource, "supports_drill_to_detail", True):
raise QueryObjectValidationError(
_("Drill to detail is not available for this datasource type.")
)
query_obj = copy.copy(query_obj)
query_obj.is_timeseries = False
query_obj.metrics = None
query_obj.post_processing = []
qry_obj_cols = []
for o in datasource.columns:
if isinstance(o, dict):
if column_name := o.get("column_name"):
qry_obj_cols.append(column_name)
else:
qry_obj_cols.append(o.column_name)
query_obj.columns = qry_obj_cols
query_obj.orderby = [(query_obj.columns[0], True)]
return query_obj
_metadata_result_type_functions: dict[
ChartDataResultType, Callable[[QueryContext, QueryObject, bool], dict[str, Any]]
] = {
ChartDataResultType.COLUMNS: _get_columns,
ChartDataResultType.TIMEGRAINS: _get_timegrains,
ChartDataResultType.QUERY: _get_query,
}
_data_result_type_preparers: dict[
ChartDataResultType,
Callable[[QueryContext, QueryObject], QueryObject] | None,
] = {
ChartDataResultType.SAMPLES: _prepare_samples_query,
ChartDataResultType.FULL: None,
ChartDataResultType.RESULTS: None,
# Post-processing is visualization-specific, so full data is returned and
# transformed later with the chart context.
ChartDataResultType.POST_PROCESSED: None,
ChartDataResultType.DRILL_DETAIL: _prepare_drill_detail_query,
}
def _metadata_timing(total_ns: int) -> QueryTiming:
return QueryTiming(
query_planning_ns=None,
cache_resolution_ns=None,
data_acquisition_ns=None,
payload_assembly_ns=None,
total_ns=total_ns,
)
def get_query_results(
result_type: ChartDataResultType,
query_context: QueryContext,
query_obj: QueryObject,
force_cached: bool,
) -> dict[str, Any]:
"""Return the legacy payload-only view of a chart-data result.
This compatibility wrapper deliberately discards the typed timing sidecar
returned by :func:`get_query_results_with_timing`.
:param result_type: the type of result to return
:param query_context: query context to which the query object belongs
:param query_obj: query object for which to retrieve the results
:param force_cached: should results be forcefully retrieved from cache
:raises QueryObjectValidationError: if an unsupported result type is requested
:return: JSON serializable result payload
"""
return get_query_results_with_timing(
result_type,
query_context,
query_obj,
force_cached,
).payload
def get_query_results_with_timing(
result_type: ChartDataResultType,
query_context: QueryContext,
query_obj: QueryObject,
force_cached: bool,
) -> QueryDataResult:
"""
Return result payload and timing without storing timing in the payload.
The total interval begins before result-family preparation and dispatch.
Metadata result types do not acquire dataframe state, so their phase values
are null while the measured total remains available.
"""
started_ns = time.perf_counter_ns()
if result_func := _metadata_result_type_functions.get(result_type):
payload = result_func(query_context, query_obj, force_cached)
total_ns = max(0, time.perf_counter_ns() - started_ns)
return QueryDataResult(payload=payload, timing=_metadata_timing(total_ns))
if result_type not in _data_result_type_preparers:
raise QueryObjectValidationError(
_("Invalid result type: %(result_type)s", result_type=result_type)
)
if preparer := _data_result_type_preparers[result_type]:
query_obj = preparer(query_context, query_obj)
payload, acquisition_timing, action_assembly_ns = _get_full_with_timing(
query_context,
query_obj,
force_cached,
)
total_ns = max(0, time.perf_counter_ns() - started_ns)
return QueryDataResult(
payload=payload,
timing=QueryTiming(
query_planning_ns=acquisition_timing.query_planning_ns,
cache_resolution_ns=acquisition_timing.cache_resolution_ns,
data_acquisition_ns=acquisition_timing.data_acquisition_ns,
payload_assembly_ns=(
acquisition_timing.payload_assembly_ns + action_assembly_ns
),
total_ns=total_ns,
),
)