mirror of
https://github.com/apache/superset.git
synced 2026-09-01 04:51:23 +00:00
Co-authored-by: Evan Rusackas <evan@preset.io> Co-authored-by: Evan Rusackas <evan@rusackas.com> Co-authored-by: Claude <noreply@anthropic.com> Co-authored-by: Joe Li <joe@preset.io>
766 lines
25 KiB
Python
766 lines
25 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.
|
|
|
|
# pylint: disable=import-outside-toplevel
|
|
|
|
from datetime import datetime
|
|
from typing import Optional
|
|
from unittest import mock
|
|
|
|
import pytest
|
|
from pytest_mock import MockerFixture
|
|
from sqlalchemy.engine.url import make_url, URL
|
|
|
|
from superset.app import SupersetApp
|
|
from superset.errors import ErrorLevel, SupersetError, SupersetErrorType
|
|
from superset.superset_typing import OAuth2ClientConfig
|
|
from superset.utils import json
|
|
from tests.unit_tests.db_engine_specs.utils import assert_convert_dttm
|
|
from tests.unit_tests.fixtures.common import dttm # noqa: F401
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
"target_type,expected_result",
|
|
[
|
|
("Date", "TO_DATE('2019-01-02')"),
|
|
("DateTime", "CAST('2019-01-02T03:04:05.678900' AS DATETIME)"),
|
|
("TimeStamp", "TO_TIMESTAMP('2019-01-02T03:04:05.678900')"),
|
|
("TIMESTAMP_NTZ", "TO_TIMESTAMP('2019-01-02T03:04:05.678900')"),
|
|
("TIMESTAMP_LTZ", "TO_TIMESTAMP('2019-01-02T03:04:05.678900')"),
|
|
("TIMESTAMP_TZ", "TO_TIMESTAMP('2019-01-02T03:04:05.678900')"),
|
|
("TIMESTAMPLTZ", "TO_TIMESTAMP('2019-01-02T03:04:05.678900')"),
|
|
("TIMESTAMPNTZ", "TO_TIMESTAMP('2019-01-02T03:04:05.678900')"),
|
|
("TIMESTAMPTZ", "TO_TIMESTAMP('2019-01-02T03:04:05.678900')"),
|
|
(
|
|
"TIMESTAMP WITH LOCAL TIME ZONE",
|
|
"TO_TIMESTAMP('2019-01-02T03:04:05.678900')",
|
|
),
|
|
("TIMESTAMP WITHOUT TIME ZONE", "TO_TIMESTAMP('2019-01-02T03:04:05.678900')"),
|
|
("UnknownType", None),
|
|
],
|
|
)
|
|
def test_convert_dttm(
|
|
target_type: str,
|
|
expected_result: Optional[str],
|
|
dttm: datetime, # noqa: F811
|
|
) -> None:
|
|
from superset.db_engine_specs.snowflake import (
|
|
SnowflakeEngineSpec as spec, # noqa: N813
|
|
)
|
|
|
|
assert_convert_dttm(spec, target_type, expected_result, dttm)
|
|
|
|
|
|
def test_database_connection_test_mutator() -> None:
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
from superset.models.core import Database
|
|
|
|
database = Database(sqlalchemy_uri="snowflake://abc")
|
|
SnowflakeEngineSpec.mutate_db_for_connection_test(database)
|
|
engine_params = json.loads(database.extra or "{}")
|
|
|
|
assert {
|
|
"engine_params": {"connect_args": {"validate_default_parameters": True}}
|
|
} == engine_params
|
|
|
|
|
|
def test_extract_errors() -> None:
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
|
|
msg = "Object dumbBrick does not exist or not authorized."
|
|
result = SnowflakeEngineSpec.extract_errors(Exception(msg))
|
|
assert result == [
|
|
SupersetError(
|
|
message="dumbBrick does not exist in this database.",
|
|
error_type=SupersetErrorType.OBJECT_DOES_NOT_EXIST_ERROR,
|
|
level=ErrorLevel.ERROR,
|
|
extra={
|
|
"engine_name": "Snowflake",
|
|
"issue_codes": [
|
|
{
|
|
"code": 1029,
|
|
"message": "Issue 1029 - The object does not exist in the given database.", # noqa: E501
|
|
}
|
|
],
|
|
},
|
|
)
|
|
]
|
|
|
|
msg = "syntax error line 1 at position 10 unexpected 'limited'."
|
|
result = SnowflakeEngineSpec.extract_errors(Exception(msg))
|
|
assert result == [
|
|
SupersetError(
|
|
message='Please check your query for syntax errors at or near "limited". Then, try running your query again.', # noqa: E501
|
|
error_type=SupersetErrorType.SYNTAX_ERROR,
|
|
level=ErrorLevel.ERROR,
|
|
extra={
|
|
"engine_name": "Snowflake",
|
|
"issue_codes": [
|
|
{
|
|
"code": 1030,
|
|
"message": "Issue 1030 - The query has a syntax error.",
|
|
}
|
|
],
|
|
},
|
|
)
|
|
]
|
|
|
|
|
|
@mock.patch("sqlalchemy.engine.Engine.connect")
|
|
def test_get_cancel_query_id(engine_mock: mock.Mock) -> None:
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
from superset.models.sql_lab import Query
|
|
|
|
query = Query()
|
|
cursor_mock = engine_mock.return_value.__enter__.return_value
|
|
cursor_mock.fetchone.return_value = [123]
|
|
assert SnowflakeEngineSpec.get_cancel_query_id(cursor_mock, query) == 123
|
|
|
|
|
|
@mock.patch("sqlalchemy.engine.Engine.connect")
|
|
def test_cancel_query(engine_mock: mock.Mock) -> None:
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
from superset.models.sql_lab import Query
|
|
|
|
query = Query()
|
|
cursor_mock = engine_mock.return_value.__enter__.return_value
|
|
assert SnowflakeEngineSpec.cancel_query(cursor_mock, query, "123") is True
|
|
|
|
|
|
@mock.patch("sqlalchemy.engine.Engine.connect")
|
|
def test_cancel_query_failed(engine_mock: mock.Mock) -> None:
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
from superset.models.sql_lab import Query
|
|
|
|
query = Query()
|
|
cursor_mock = engine_mock.raiseError.side_effect = Exception()
|
|
assert SnowflakeEngineSpec.cancel_query(cursor_mock, query, "123") is False
|
|
|
|
|
|
def test_get_extra_params(mocker: MockerFixture) -> None:
|
|
"""
|
|
Test the ``get_extra_params`` method.
|
|
"""
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
|
|
database = mocker.MagicMock()
|
|
|
|
database.extra = {}
|
|
assert SnowflakeEngineSpec.get_extra_params(database) == {
|
|
"engine_params": {"connect_args": {"application": "Apache Superset"}}
|
|
}
|
|
|
|
database.extra = json.dumps(
|
|
{
|
|
"engine_params": {
|
|
"connect_args": {"application": "Custom user agent", "foo": "bar"}
|
|
}
|
|
}
|
|
)
|
|
assert SnowflakeEngineSpec.get_extra_params(database) == {
|
|
"engine_params": {
|
|
"connect_args": {"application": "Custom user agent", "foo": "bar"}
|
|
}
|
|
}
|
|
|
|
|
|
def test_get_schema_from_engine_params() -> None:
|
|
"""
|
|
Test the ``get_schema_from_engine_params`` method.
|
|
"""
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
|
|
assert (
|
|
SnowflakeEngineSpec.get_schema_from_engine_params(
|
|
make_url("snowflake://user:pass@account/database_name/default"),
|
|
{},
|
|
)
|
|
== "default"
|
|
)
|
|
|
|
assert (
|
|
SnowflakeEngineSpec.get_schema_from_engine_params(
|
|
make_url("snowflake://user:pass@account/database_name"),
|
|
{},
|
|
)
|
|
is None
|
|
)
|
|
|
|
assert (
|
|
SnowflakeEngineSpec.get_schema_from_engine_params(
|
|
make_url("snowflake://user:pass@account/"),
|
|
{},
|
|
)
|
|
is None
|
|
)
|
|
|
|
|
|
def test_adjust_engine_params_fully_qualified() -> None:
|
|
"""
|
|
Test the ``adjust_engine_params`` method when the URL has catalog and schema.
|
|
"""
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
|
|
url = make_url("snowflake://user:pass@account/database_name/default")
|
|
|
|
uri = SnowflakeEngineSpec.adjust_engine_params(url, {})[0]
|
|
assert (
|
|
uri.render_as_string(hide_password=False)
|
|
== "snowflake://user:pass@account/database_name/default"
|
|
)
|
|
|
|
uri = SnowflakeEngineSpec.adjust_engine_params(
|
|
url,
|
|
{},
|
|
schema="new_schema",
|
|
)[0]
|
|
assert (
|
|
uri.render_as_string(hide_password=False)
|
|
== "snowflake://user:pass@account/database_name/new_schema"
|
|
)
|
|
|
|
uri = SnowflakeEngineSpec.adjust_engine_params(
|
|
url,
|
|
{},
|
|
catalog="new_catalog",
|
|
)[0]
|
|
assert (
|
|
uri.render_as_string(hide_password=False)
|
|
== "snowflake://user:pass@account/new_catalog/default"
|
|
)
|
|
|
|
uri = SnowflakeEngineSpec.adjust_engine_params(
|
|
url,
|
|
{},
|
|
catalog="new_catalog",
|
|
schema="new_schema",
|
|
)[0]
|
|
assert (
|
|
uri.render_as_string(hide_password=False)
|
|
== "snowflake://user:pass@account/new_catalog/new_schema"
|
|
)
|
|
|
|
|
|
def test_adjust_engine_params_catalog_only() -> None:
|
|
"""
|
|
Test the ``adjust_engine_params`` method when the URL has only the catalog.
|
|
"""
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
|
|
url = make_url("snowflake://user:pass@account/database_name")
|
|
|
|
uri = SnowflakeEngineSpec.adjust_engine_params(url, {})[0]
|
|
assert (
|
|
uri.render_as_string(hide_password=False)
|
|
== "snowflake://user:pass@account/database_name"
|
|
)
|
|
|
|
uri = SnowflakeEngineSpec.adjust_engine_params(
|
|
url,
|
|
{},
|
|
schema="new_schema",
|
|
)[0]
|
|
assert (
|
|
uri.render_as_string(hide_password=False)
|
|
== "snowflake://user:pass@account/database_name/new_schema"
|
|
)
|
|
|
|
uri = SnowflakeEngineSpec.adjust_engine_params(
|
|
url,
|
|
{},
|
|
catalog="new_catalog",
|
|
)[0]
|
|
assert (
|
|
uri.render_as_string(hide_password=False)
|
|
== "snowflake://user:pass@account/new_catalog"
|
|
)
|
|
|
|
uri = SnowflakeEngineSpec.adjust_engine_params(
|
|
url,
|
|
{},
|
|
catalog="new_catalog",
|
|
schema="new_schema",
|
|
)[0]
|
|
assert (
|
|
uri.render_as_string(hide_password=False)
|
|
== "snowflake://user:pass@account/new_catalog/new_schema"
|
|
)
|
|
|
|
|
|
def test_get_default_catalog() -> None:
|
|
"""
|
|
Test the ``get_default_catalog`` method.
|
|
"""
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
from superset.models.core import Database
|
|
|
|
database = Database(
|
|
database_name="my_db",
|
|
sqlalchemy_uri="snowflake://user:pass@account/database_name",
|
|
)
|
|
assert SnowflakeEngineSpec.get_default_catalog(database) == "database_name"
|
|
|
|
database = Database(
|
|
database_name="my_db",
|
|
sqlalchemy_uri="snowflake://user:pass@account/database_name/default",
|
|
)
|
|
assert SnowflakeEngineSpec.get_default_catalog(database) == "database_name"
|
|
|
|
|
|
def test_mask_encrypted_extra() -> None:
|
|
"""
|
|
Test that the private keys are masked when the database is edited.
|
|
"""
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
|
|
config = json.dumps(
|
|
{
|
|
"auth_method": "keypair",
|
|
"auth_params": {
|
|
"privatekey_body": (
|
|
"-----BEGIN ENCRYPTED PRIVATE KEY-----"
|
|
"..."
|
|
"-----END ENCRYPTED PRIVATE KEY-----"
|
|
),
|
|
"privatekey_pass": "my_password",
|
|
},
|
|
}
|
|
)
|
|
|
|
assert SnowflakeEngineSpec.mask_encrypted_extra(config) == json.dumps(
|
|
{
|
|
"auth_method": "keypair",
|
|
"auth_params": {
|
|
"privatekey_body": "XXXXXXXXXX",
|
|
"privatekey_pass": "XXXXXXXXXX",
|
|
},
|
|
}
|
|
)
|
|
|
|
|
|
def test_mask_encrypted_extra_oauth2_client_secret() -> None:
|
|
"""
|
|
The database-level OAuth2 client secret must be masked in
|
|
``masked_encrypted_extra``, matching the other engine specs supporting
|
|
the same ``oauth2_client_info`` path (gsheets, trino) -- otherwise a
|
|
database editor can read it back unmasked.
|
|
"""
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
|
|
config = json.dumps(
|
|
{
|
|
"auth_method": "oauth2",
|
|
"oauth2_client_info": {"id": "client-id", "secret": "my-secret"},
|
|
}
|
|
)
|
|
|
|
assert SnowflakeEngineSpec.mask_encrypted_extra(config) == json.dumps(
|
|
{
|
|
"auth_method": "oauth2",
|
|
"oauth2_client_info": {"id": "client-id", "secret": "XXXXXXXXXX"},
|
|
}
|
|
)
|
|
|
|
|
|
def test_mask_encrypted_extra_no_fields() -> None:
|
|
"""
|
|
Test that the private key is masked when the database is edited.
|
|
"""
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
|
|
config = json.dumps(
|
|
{
|
|
# this is a fake example and the fields are made up
|
|
"auth_method": "token",
|
|
"auth_params": {
|
|
"jwt": "SECRET",
|
|
},
|
|
}
|
|
)
|
|
|
|
assert SnowflakeEngineSpec.mask_encrypted_extra(config) == json.dumps(
|
|
{
|
|
"auth_method": "token",
|
|
"auth_params": {
|
|
"jwt": "SECRET",
|
|
},
|
|
}
|
|
)
|
|
|
|
|
|
def test_handle_boolean_filter() -> None:
|
|
"""
|
|
Test that Snowflake uses equality operators for boolean filters instead of IS.
|
|
"""
|
|
from sqlalchemy import Boolean, Column
|
|
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
|
|
# Create a mock SQLAlchemy column
|
|
bool_col = Column("test_col", Boolean)
|
|
|
|
# Test IS_TRUE filter - use actual FilterOperator values
|
|
from superset.utils.core import FilterOperator
|
|
|
|
result_true = SnowflakeEngineSpec.handle_boolean_filter(
|
|
bool_col, FilterOperator.IS_TRUE, True
|
|
)
|
|
# The result should be a equality comparison, not an IS comparison
|
|
assert (
|
|
str(result_true.compile(compile_kwargs={"literal_binds": True}))
|
|
== "test_col = true"
|
|
)
|
|
|
|
# Test IS_FALSE filter
|
|
result_false = SnowflakeEngineSpec.handle_boolean_filter(
|
|
bool_col, FilterOperator.IS_FALSE, False
|
|
)
|
|
assert (
|
|
str(result_false.compile(compile_kwargs={"literal_binds": True}))
|
|
== "test_col = false"
|
|
)
|
|
|
|
|
|
def test_use_equality_for_boolean_filters_property() -> None:
|
|
"""
|
|
Test that Snowflake has the use_equality_for_boolean_filters property set to True.
|
|
"""
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
|
|
assert SnowflakeEngineSpec.use_equality_for_boolean_filters is True
|
|
|
|
|
|
def test_unmask_encrypted_extra() -> None:
|
|
"""
|
|
Test that the private keys can be reused from the previous `encrypted_extra`.
|
|
"""
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
|
|
old = json.dumps(
|
|
{
|
|
"auth_method": "keypair",
|
|
"auth_params": {
|
|
"privatekey_body": (
|
|
"-----BEGIN ENCRYPTED PRIVATE KEY-----"
|
|
"..."
|
|
"-----END ENCRYPTED PRIVATE KEY-----"
|
|
),
|
|
"privatekey_pass": "my_password",
|
|
},
|
|
}
|
|
)
|
|
new = json.dumps(
|
|
{
|
|
"foo": "bar",
|
|
"auth_method": "keypair",
|
|
"auth_params": {
|
|
"privatekey_body": "XXXXXXXXXX",
|
|
"privatekey_pass": "XXXXXXXXXX",
|
|
},
|
|
}
|
|
)
|
|
|
|
assert SnowflakeEngineSpec.unmask_encrypted_extra(old, new) == json.dumps(
|
|
{
|
|
"foo": "bar",
|
|
"auth_method": "keypair",
|
|
"auth_params": {
|
|
"privatekey_body": (
|
|
"-----BEGIN ENCRYPTED PRIVATE KEY-----"
|
|
"..."
|
|
"-----END ENCRYPTED PRIVATE KEY-----"
|
|
),
|
|
"privatekey_pass": "my_password",
|
|
},
|
|
}
|
|
)
|
|
|
|
|
|
@pytest.fixture
|
|
def oauth2_config() -> OAuth2ClientConfig:
|
|
"""
|
|
Config for Snowflake OAuth2.
|
|
"""
|
|
return {
|
|
"id": "snowflake-oauth2-client-id",
|
|
"secret": "snowflake-oauth2-client-secret",
|
|
"scope": "refresh_token",
|
|
"redirect_uri": "http://localhost:8088/api/v1/database/oauth2/",
|
|
"authorization_request_uri": "https://snowflake.oauth2.example/oauth/authorize",
|
|
"token_request_uri": "https://snowflake.oauth2.example/oauth/token-request",
|
|
"request_content_type": "data",
|
|
}
|
|
|
|
|
|
def test_get_oauth2_token(
|
|
mocker: MockerFixture,
|
|
oauth2_config: OAuth2ClientConfig,
|
|
) -> None:
|
|
"""
|
|
Test `get_oauth2_token`.
|
|
"""
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
|
|
requests: mock.MagicMock = mocker.patch("superset.db_engine_specs.base.requests")
|
|
requests.post().json.return_value = {
|
|
"access_token": "access-token",
|
|
"expires_in": 3600,
|
|
"scope": "scope",
|
|
"token_type": "Bearer",
|
|
"refresh_token": "refresh-token",
|
|
}
|
|
|
|
assert SnowflakeEngineSpec.get_oauth2_token(oauth2_config, "code") == {
|
|
"access_token": "access-token",
|
|
"expires_in": 3600,
|
|
"scope": "scope",
|
|
"token_type": "Bearer",
|
|
"refresh_token": "refresh-token",
|
|
}
|
|
requests.post.assert_called_with(
|
|
"https://snowflake.oauth2.example/oauth/token-request",
|
|
data={
|
|
"code": "code",
|
|
"client_id": "snowflake-oauth2-client-id",
|
|
"client_secret": "snowflake-oauth2-client-secret",
|
|
"redirect_uri": "http://localhost:8088/api/v1/database/oauth2/",
|
|
"grant_type": "authorization_code",
|
|
},
|
|
timeout=30.0,
|
|
)
|
|
|
|
|
|
def test_impersonate_user(app: SupersetApp, mocker: MockerFixture) -> None:
|
|
"""
|
|
Test that Snowflake supports user impersonation.
|
|
|
|
Impersonation only applies within a request context (see
|
|
``test_impersonate_user_outside_request_context`` below for the
|
|
background-execution case), so these assertions run inside one.
|
|
"""
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
from superset.models.core import Database
|
|
|
|
database: Database = Database(sqlalchemy_uri="snowflake://abc")
|
|
|
|
mocker.patch(
|
|
"superset.db_engine_specs.snowflake.SnowflakeEngineSpec.is_oauth2_enabled",
|
|
return_value=True,
|
|
)
|
|
|
|
with app.test_request_context("/some/place/"):
|
|
assert SnowflakeEngineSpec.impersonate_user(
|
|
database=database,
|
|
username=None,
|
|
user_token=None,
|
|
url=make_url("snowflake://user:pass@account/database_name/default"),
|
|
engine_kwargs={
|
|
"connect_args": {
|
|
"validate_default_parameters": True,
|
|
},
|
|
},
|
|
) == (
|
|
make_url("snowflake://user:pass@account/database_name/default"),
|
|
{"connect_args": {"validate_default_parameters": True}},
|
|
)
|
|
|
|
assert SnowflakeEngineSpec.impersonate_user(
|
|
database=database,
|
|
username=None,
|
|
user_token=None,
|
|
url=make_url("snowflake://user:pass@account/database_name/default"),
|
|
engine_kwargs={},
|
|
) == (
|
|
make_url(
|
|
"snowflake://user:pass@account/database_name/default?authenticator=oauth"
|
|
),
|
|
{"connect_args": {"authenticator": "oauth"}},
|
|
)
|
|
|
|
mocker.patch(
|
|
"superset.db_engine_specs.snowflake.is_feature_enabled",
|
|
return_value=True,
|
|
)
|
|
|
|
mocker.patch(
|
|
"superset.security_manager.find_user",
|
|
return_value=mocker.MagicMock(email="impersonated_user@example.com"),
|
|
)
|
|
assert SnowflakeEngineSpec.impersonate_user(
|
|
database=database,
|
|
username="impersonated_user",
|
|
user_token="test_token", # noqa: S106
|
|
url=make_url("snowflake://user:pass@account/database_name/default"),
|
|
engine_kwargs={},
|
|
) == (
|
|
make_url(
|
|
"snowflake://impersonated_user:pass@account/database_name/default?authenticator=oauth&token=test_token"
|
|
),
|
|
{"connect_args": {"authenticator": "oauth"}},
|
|
)
|
|
|
|
|
|
def test_impersonate_user_email_prefix_uses_username_directly(
|
|
app: SupersetApp, mocker: MockerFixture
|
|
) -> None:
|
|
"""
|
|
With IMPERSONATE_WITH_EMAIL_PREFIX enabled, ``Database._get_sqla_engine()``
|
|
has already substituted the email prefix for the login username before
|
|
calling ``impersonate_user`` -- the value it passes in is no longer a
|
|
lookupable login. Re-looking it up as a username (the pre-fix behavior)
|
|
fails whenever the login differs from the prefix, silently leaving the
|
|
default/service-account username paired with the impersonated user's
|
|
OAuth token instead of failing loudly. The fixed code must use the given
|
|
value directly and must not call ``find_user`` at all in this branch.
|
|
"""
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
from superset.models.core import Database
|
|
|
|
database: Database = Database(sqlalchemy_uri="snowflake://abc")
|
|
|
|
mocker.patch(
|
|
"superset.db_engine_specs.snowflake.SnowflakeEngineSpec.is_oauth2_enabled",
|
|
return_value=True,
|
|
)
|
|
mocker.patch(
|
|
"superset.db_engine_specs.snowflake.is_feature_enabled",
|
|
return_value=True,
|
|
)
|
|
find_user = mocker.patch("superset.security_manager.find_user")
|
|
|
|
with app.test_request_context("/some/place/"):
|
|
# "jdoe" is the email prefix Database._get_sqla_engine() already
|
|
# derived; the login it derived it from ("jdoe123", say) is gone by
|
|
# this point and must not be re-derived here.
|
|
result = SnowflakeEngineSpec.impersonate_user(
|
|
database=database,
|
|
username="jdoe",
|
|
user_token="test_token", # noqa: S106
|
|
url=make_url("snowflake://user:pass@account/database_name/default"),
|
|
engine_kwargs={},
|
|
)
|
|
|
|
assert result == (
|
|
make_url(
|
|
"snowflake://jdoe:pass@account/database_name/default?authenticator=oauth&token=test_token"
|
|
),
|
|
{"connect_args": {"authenticator": "oauth"}},
|
|
)
|
|
find_user.assert_not_called()
|
|
|
|
|
|
def test_impersonate_user_outside_request_context(mocker: MockerFixture) -> None:
|
|
"""
|
|
Background executions (alerts/reports) have no per-user token, so OAuth
|
|
impersonation must not engage outside a request context — even when
|
|
``database.is_oauth2_enabled()`` returns True because of a
|
|
database-level OAuth2 client config, which (unlike the app-config-based
|
|
check) isn't itself request-context-aware.
|
|
"""
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
from superset.models.core import Database
|
|
|
|
database: Database = Database(sqlalchemy_uri="snowflake://abc")
|
|
mocker.patch.object(Database, "is_oauth2_enabled", return_value=True)
|
|
|
|
url: URL = make_url("snowflake://user:pass@account/database_name/default")
|
|
assert SnowflakeEngineSpec.impersonate_user(
|
|
database=database,
|
|
username=None,
|
|
user_token="test_token", # noqa: S106
|
|
url=url,
|
|
engine_kwargs={},
|
|
) == (url, {"connect_args": {}})
|
|
|
|
|
|
def test_custom_snowflake_auth_error_matches_raw_dbapi_exception() -> None:
|
|
"""
|
|
`BaseEngineSpec.execute()` runs against a bare DBAPI cursor, so the
|
|
exception it sees is the raw Snowflake error, never wrapped by
|
|
SQLAlchemy. `CustomSnowflakeAuthError` must still recognize it so the
|
|
OAuth2 re-auth dance triggers for SQL Lab queries.
|
|
"""
|
|
from superset.db_engine_specs.snowflake import (
|
|
CustomSnowflakeAuthError,
|
|
DatabaseError,
|
|
)
|
|
|
|
raw_error: Exception = DatabaseError("250001: Invalid OAuth access token.")
|
|
assert isinstance(raw_error, CustomSnowflakeAuthError)
|
|
|
|
|
|
def test_custom_snowflake_auth_error_matches_sqlalchemy_wrapped_exception() -> None:
|
|
"""
|
|
Some call sites execute through SQLAlchemy's `Engine`, which wraps the
|
|
original DBAPI exception in `sqlalchemy.exc.DatabaseError.orig`.
|
|
`CustomSnowflakeAuthError` must keep matching this shape too.
|
|
"""
|
|
from sqlalchemy.exc import DatabaseError as SqlalchemyDatabaseError
|
|
|
|
from superset.db_engine_specs.snowflake import (
|
|
CustomSnowflakeAuthError,
|
|
DatabaseError,
|
|
)
|
|
|
|
wrapped_error: SqlalchemyDatabaseError = SqlalchemyDatabaseError(
|
|
statement="SELECT 1",
|
|
params=None,
|
|
orig=DatabaseError("250001: Invalid OAuth access token."),
|
|
)
|
|
assert isinstance(wrapped_error, CustomSnowflakeAuthError)
|
|
|
|
|
|
def test_custom_snowflake_auth_error_does_not_match_unrelated_errors() -> None:
|
|
"""
|
|
Other Snowflake DB errors, and non-Snowflake exceptions, must not be
|
|
mistaken for an expired OAuth token.
|
|
"""
|
|
from superset.db_engine_specs.snowflake import (
|
|
CustomSnowflakeAuthError,
|
|
DatabaseError,
|
|
)
|
|
|
|
assert not isinstance(
|
|
DatabaseError("Object FOO does not exist."), CustomSnowflakeAuthError
|
|
)
|
|
assert not isinstance(
|
|
ValueError("Invalid OAuth access token."), CustomSnowflakeAuthError
|
|
)
|
|
|
|
|
|
def test_snowflake_oauth2_exception_catches_refresh_token_error() -> None:
|
|
"""
|
|
`refresh_oauth2_token()` catches failures from the (unoverridden) base
|
|
`get_oauth2_fresh_token()` with `except db_engine_spec.oauth2_exception`.
|
|
That base method raises `OAuth2TokenRefreshError`, which isn't related to
|
|
`CustomSnowflakeAuthError` by real subclassing, so `oauth2_exception` must
|
|
include it directly -- an `except` clause never triggers the metaclass's
|
|
`__instancecheck__`, unlike `isinstance()`.
|
|
"""
|
|
from superset.db_engine_specs.snowflake import SnowflakeEngineSpec
|
|
from superset.exceptions import OAuth2TokenRefreshError
|
|
|
|
try:
|
|
raise OAuth2TokenRefreshError("refresh token revoked")
|
|
except SnowflakeEngineSpec.oauth2_exception:
|
|
pass
|
|
else:
|
|
pytest.fail(
|
|
"OAuth2TokenRefreshError must be caught by "
|
|
"SnowflakeEngineSpec.oauth2_exception"
|
|
)
|