From 3fee7d3b3c8263a587dc38df5b07ff2ff127d52f Mon Sep 17 00:00:00 2001 From: vaggelisd Date: Mon, 24 Feb 2025 15:01:49 +0200 Subject: [PATCH 1/4] Feat: Set normalization strategy according to used gateway --- sqlmesh/core/config/gateway.py | 2 ++ sqlmesh/core/config/root.py | 5 +++++ tests/core/test_context.py | 31 ++++++++++++++++++++++++++++++- 3 files changed, 37 insertions(+), 1 deletion(-) diff --git a/sqlmesh/core/config/gateway.py b/sqlmesh/core/config/gateway.py index 6184d500d2..a51557c4d7 100644 --- a/sqlmesh/core/config/gateway.py +++ b/sqlmesh/core/config/gateway.py @@ -4,6 +4,7 @@ from sqlmesh.core import constants as c from sqlmesh.core.config.base import BaseConfig +from sqlmesh.core.config.model import ModelDefaultsConfig from sqlmesh.core.config.common import variables_validator from sqlmesh.core.config.connection import ( SerializableConnectionConfig, @@ -34,6 +35,7 @@ class GatewayConfig(BaseConfig): scheduler: t.Optional[SchedulerConfig] = None state_schema: t.Optional[str] = c.SQLMESH variables: t.Dict[str, t.Any] = {} + model_defaults: t.Optional[ModelDefaultsConfig] = None _connection_config_validator = connection_config_validator _scheduler_config_validator = scheduler_config_validator diff --git a/sqlmesh/core/config/root.py b/sqlmesh/core/config/root.py index 7b0881df67..a13f04ec11 100644 --- a/sqlmesh/core/config/root.py +++ b/sqlmesh/core/config/root.py @@ -287,6 +287,11 @@ def default_gateway_name(self) -> str: @property def dialect(self) -> t.Optional[str]: + if self.default_gateway: + gateway_config = self.gateways[self.default_gateway] + if gateway_config.model_defaults: + return gateway_config.model_defaults.dialect + return self.model_defaults.dialect @property diff --git a/tests/core/test_context.py b/tests/core/test_context.py index a693d869f6..8fde3d7b05 100644 --- a/tests/core/test_context.py +++ b/tests/core/test_context.py @@ -11,7 +11,7 @@ import pandas as pd from pathlib import Path from pytest_mock.plugin import MockerFixture -from sqlglot import exp, parse_one +from sqlglot import exp, parse_one, Dialect from sqlglot.errors import SchemaError from sqlmesh.core.config.gateway import GatewayConfig @@ -1103,6 +1103,35 @@ def test_override_dialect_normalization_strategy(): DuckDB.NORMALIZATION_STRATEGY = NormalizationStrategy.CASE_INSENSITIVE +def test_different_gateway_normalization_strategy(tmp_path: pathlib.Path): + config = Config( + gateways={ + "duckdb": GatewayConfig( + connection=DuckDBConnectionConfig(database="db.db"), + model_defaults=ModelDefaultsConfig( + dialect="snowflake, normalization_strategy=case_insensitive" + ), + ) + }, + model_defaults=ModelDefaultsConfig(dialect="snowflake"), + default_gateway="duckdb", + ) + + from sqlglot.dialects import Snowflake + from sqlglot.dialects.dialect import NormalizationStrategy + + assert Snowflake.NORMALIZATION_STRATEGY == NormalizationStrategy.UPPERCASE + + ctx = Context(paths=tmp_path, config=config, gateway="duckdb") + + dialect = Dialect.get_or_raise(ctx.config.dialect) + + assert dialect == "snowflake" + assert Snowflake.NORMALIZATION_STRATEGY == NormalizationStrategy.CASE_INSENSITIVE + + Snowflake.NORMALIZATION_STRATEGY = NormalizationStrategy.UPPERCASE + + def test_access_self_columns_to_types_in_macro(tmp_path: pathlib.Path): create_temp_file( tmp_path, From 2f23206ec0ca60f0739eb5f73d6d743f499c5635 Mon Sep 17 00:00:00 2001 From: vaggelisd Date: Mon, 24 Feb 2025 22:10:08 +0200 Subject: [PATCH 2/4] Generalize for all model defaults --- sqlmesh/core/config/root.py | 5 ----- sqlmesh/core/context.py | 12 ++++++++++++ tests/core/test_config.py | 26 ++++++++++++++++++++++++++ tests/core/test_context.py | 2 +- 4 files changed, 39 insertions(+), 6 deletions(-) diff --git a/sqlmesh/core/config/root.py b/sqlmesh/core/config/root.py index a13f04ec11..7b0881df67 100644 --- a/sqlmesh/core/config/root.py +++ b/sqlmesh/core/config/root.py @@ -287,11 +287,6 @@ def default_gateway_name(self) -> str: @property def dialect(self) -> t.Optional[str]: - if self.default_gateway: - gateway_config = self.gateways[self.default_gateway] - if gateway_config.model_defaults: - return gateway_config.model_defaults.dialect - return self.model_defaults.dialect @property diff --git a/sqlmesh/core/context.py b/sqlmesh/core/context.py index f2ae2cd4c3..2b0a6dcf32 100644 --- a/sqlmesh/core/context.py +++ b/sqlmesh/core/context.py @@ -78,6 +78,7 @@ from sqlmesh.core.macros import ExecutableOrMacro, macro from sqlmesh.core.metric import Metric, rewrite from sqlmesh.core.model import Model, update_model_schemas +from sqlmesh.core.config.model import ModelDefaultsConfig from sqlmesh.core.notification_target import ( NotificationEvent, NotificationTarget, @@ -382,6 +383,17 @@ def __init__( self.selected_gateway: self._connection_config.create_engine_adapter() } + if self.selected_gateway: + gw_model_defaults = self.config.gateways[self.selected_gateway] + if gw_model_defaults.model_defaults: + # Merge global model defaults with the selected gateway's, if it's overriden + global_defaults = self.config.model_defaults.model_dump(exclude_unset=True) + gateway_defaults = gw_model_defaults.model_defaults.model_dump(exclude_unset=True) + + self.config.model_defaults = ModelDefaultsConfig( + **{**global_defaults, **gateway_defaults} + ) + self._snapshot_evaluator: t.Optional[SnapshotEvaluator] = None self.console = get_console() diff --git a/tests/core/test_config.py b/tests/core/test_config.py index 28646c341f..ccc8bdb599 100644 --- a/tests/core/test_config.py +++ b/tests/core/test_config.py @@ -850,3 +850,29 @@ def test_gcp_postgres_ip_and_scopes(tmp_path): assert conn.scopes[0] == "https://www.googleapis.com/auth/cloud-platform" assert conn.scopes[1] == "https://www.googleapis.com/auth/sqlservice.admin" assert conn.ip_type == "private" + + +def test_gateway_model_defaults(tmp_path): + global_defaults = ModelDefaultsConfig( + dialect="snowflake", owner="foo", optimize_query=True, enabled=True, cron="@daily" + ) + gateway_defaults = ModelDefaultsConfig(dialect="duckdb", owner="baz", optimize_query=False) + + config = Config( + gateways={ + "duckdb": GatewayConfig( + connection=DuckDBConnectionConfig(database="db.db"), + model_defaults=gateway_defaults, + ) + }, + model_defaults=global_defaults, + default_gateway="duckdb", + ) + + ctx = Context(paths=tmp_path, config=config, gateway="duckdb") + + expected = ModelDefaultsConfig( + dialect="duckdb", owner="baz", optimize_query=False, enabled=True, cron="@daily" + ) + + assert ctx.config.model_defaults == expected diff --git a/tests/core/test_context.py b/tests/core/test_context.py index 8fde3d7b05..e33040f11a 100644 --- a/tests/core/test_context.py +++ b/tests/core/test_context.py @@ -1127,7 +1127,7 @@ def test_different_gateway_normalization_strategy(tmp_path: pathlib.Path): dialect = Dialect.get_or_raise(ctx.config.dialect) assert dialect == "snowflake" - assert Snowflake.NORMALIZATION_STRATEGY == NormalizationStrategy.CASE_INSENSITIVE + assert dialect.NORMALIZATION_STRATEGY == NormalizationStrategy.CASE_INSENSITIVE Snowflake.NORMALIZATION_STRATEGY = NormalizationStrategy.UPPERCASE From 2e8c08594e380f1c393879705aa686c771dc6a20 Mon Sep 17 00:00:00 2001 From: vaggelisd Date: Mon, 24 Feb 2025 22:22:23 +0200 Subject: [PATCH 3/4] Simplify --- sqlmesh/core/context.py | 6 +++--- tests/core/test_context.py | 2 +- 2 files changed, 4 insertions(+), 4 deletions(-) diff --git a/sqlmesh/core/context.py b/sqlmesh/core/context.py index 2b0a6dcf32..1067c45c29 100644 --- a/sqlmesh/core/context.py +++ b/sqlmesh/core/context.py @@ -384,11 +384,11 @@ def __init__( } if self.selected_gateway: - gw_model_defaults = self.config.gateways[self.selected_gateway] - if gw_model_defaults.model_defaults: + gw_model_defaults = self.config.gateways[self.selected_gateway].model_defaults + if gw_model_defaults: # Merge global model defaults with the selected gateway's, if it's overriden global_defaults = self.config.model_defaults.model_dump(exclude_unset=True) - gateway_defaults = gw_model_defaults.model_defaults.model_dump(exclude_unset=True) + gateway_defaults = gw_model_defaults.model_dump(exclude_unset=True) self.config.model_defaults = ModelDefaultsConfig( **{**global_defaults, **gateway_defaults} diff --git a/tests/core/test_context.py b/tests/core/test_context.py index e33040f11a..3d90f11018 100644 --- a/tests/core/test_context.py +++ b/tests/core/test_context.py @@ -1127,7 +1127,7 @@ def test_different_gateway_normalization_strategy(tmp_path: pathlib.Path): dialect = Dialect.get_or_raise(ctx.config.dialect) assert dialect == "snowflake" - assert dialect.NORMALIZATION_STRATEGY == NormalizationStrategy.CASE_INSENSITIVE + assert dialect.normalization_strategy == NormalizationStrategy.CASE_INSENSITIVE Snowflake.NORMALIZATION_STRATEGY = NormalizationStrategy.UPPERCASE From 5b9e24609d4eda40e363bbb0aac772a72482eee3 Mon Sep 17 00:00:00 2001 From: vaggelisd Date: Mon, 24 Feb 2025 22:41:41 +0200 Subject: [PATCH 4/4] Move above patching --- sqlmesh/core/context.py | 35 +++++++++++++++++------------------ tests/core/test_context.py | 2 +- 2 files changed, 18 insertions(+), 19 deletions(-) diff --git a/sqlmesh/core/context.py b/sqlmesh/core/context.py index 1067c45c29..c656735341 100644 --- a/sqlmesh/core/context.py +++ b/sqlmesh/core/context.py @@ -354,13 +354,6 @@ def __init__( self._all_dialects: t.Set[str] = {self.config.dialect or ""} - # This allows overriding the default dialect's normalization strategy, so for example - # one can do `dialect="duckdb,normalization_strategy=lowercase"` and this will be - # applied to the DuckDB dialect globally - if "normalization_strategy" in str(self.config.dialect): - dialect = Dialect.get_or_raise(self.config.dialect) - type(dialect).NORMALIZATION_STRATEGY = dialect.normalization_strategy - if self.config.disable_anonymized_analytics: analytics.disable_analytics() @@ -371,6 +364,23 @@ def __init__( self.auto_categorize_changes = self.config.plan.auto_categorize_changes self.selected_gateway = gateway or self.config.default_gateway_name + gw_model_defaults = self.config.gateways[self.selected_gateway].model_defaults + if gw_model_defaults: + # Merge global model defaults with the selected gateway's, if it's overriden + global_defaults = self.config.model_defaults.model_dump(exclude_unset=True) + gateway_defaults = gw_model_defaults.model_dump(exclude_unset=True) + + self.config.model_defaults = ModelDefaultsConfig( + **{**global_defaults, **gateway_defaults} + ) + + # This allows overriding the default dialect's normalization strategy, so for example + # one can do `dialect="duckdb,normalization_strategy=lowercase"` and this will be + # applied to the DuckDB dialect globally + if "normalization_strategy" in str(self.config.dialect): + dialect = Dialect.get_or_raise(self.config.dialect) + type(dialect).NORMALIZATION_STRATEGY = dialect.normalization_strategy + self._loaders = [ (loader or config.loader)(self, path, **config.loader_kwargs) for path, config in self.configs.items() @@ -383,17 +393,6 @@ def __init__( self.selected_gateway: self._connection_config.create_engine_adapter() } - if self.selected_gateway: - gw_model_defaults = self.config.gateways[self.selected_gateway].model_defaults - if gw_model_defaults: - # Merge global model defaults with the selected gateway's, if it's overriden - global_defaults = self.config.model_defaults.model_dump(exclude_unset=True) - gateway_defaults = gw_model_defaults.model_dump(exclude_unset=True) - - self.config.model_defaults = ModelDefaultsConfig( - **{**global_defaults, **gateway_defaults} - ) - self._snapshot_evaluator: t.Optional[SnapshotEvaluator] = None self.console = get_console() diff --git a/tests/core/test_context.py b/tests/core/test_context.py index 3d90f11018..8fde3d7b05 100644 --- a/tests/core/test_context.py +++ b/tests/core/test_context.py @@ -1127,7 +1127,7 @@ def test_different_gateway_normalization_strategy(tmp_path: pathlib.Path): dialect = Dialect.get_or_raise(ctx.config.dialect) assert dialect == "snowflake" - assert dialect.normalization_strategy == NormalizationStrategy.CASE_INSENSITIVE + assert Snowflake.NORMALIZATION_STRATEGY == NormalizationStrategy.CASE_INSENSITIVE Snowflake.NORMALIZATION_STRATEGY = NormalizationStrategy.UPPERCASE