From 99798619b6538881b1009e92e60cec32bacc2993 Mon Sep 17 00:00:00 2001 From: Themis Valtinos <73662635+themisvaltinos@users.noreply.github.com> Date: Mon, 7 Apr 2025 10:56:56 +0300 Subject: [PATCH 01/14] Feat!: Introduce multiple gateway virtual layer --- docs/reference/configuration.md | 1 + sqlmesh/core/config/root.py | 2 + sqlmesh/core/context.py | 29 +++++++++------ sqlmesh/core/context_diff.py | 10 ++++- sqlmesh/core/environment.py | 9 +++++ sqlmesh/core/loader.py | 2 + sqlmesh/core/model/definition.py | 7 ++++ sqlmesh/core/plan/builder.py | 1 + sqlmesh/core/plan/evaluator.py | 10 ++++- sqlmesh/core/snapshot/evaluator.py | 37 ++++++++++++++----- sqlmesh/core/state_sync/base.py | 2 +- sqlmesh/core/state_sync/cache.py | 2 +- sqlmesh/core/state_sync/db/environment.py | 2 + sqlmesh/core/state_sync/db/facade.py | 16 +++++--- sqlmesh/core/state_sync/db/snapshot.py | 10 +++-- .../migrations/v0056_restore_table_indexes.py | 2 +- ...v0078_add_gateway_managed_virtual_layer.py | 28 ++++++++++++++ sqlmesh/schedulers/airflow/state_sync.py | 2 +- tests/core/test_config.py | 6 --- tests/core/test_context.py | 6 --- tests/schedulers/airflow/test_client.py | 1 + 21 files changed, 137 insertions(+), 48 deletions(-) create mode 100644 sqlmesh/migrations/v0078_add_gateway_managed_virtual_layer.py diff --git a/docs/reference/configuration.md b/docs/reference/configuration.md index 8956f700b9..e9d45810a8 100644 --- a/docs/reference/configuration.md +++ b/docs/reference/configuration.md @@ -35,6 +35,7 @@ Configuration options for SQLMesh environment creation and promotion. | `physical_schema_override` | (Deprecated) Use `physical_schema_mapping` instead. A mapping from model schema names to names of schemas in which physical tables for the corresponding models will be placed. | dict[string, string] | N | | `physical_schema_mapping` | A mapping from regular expressions to names of schemas in which physical tables for the corresponding models [will be placed](../guides/configuration.md#physical-table-schemas). (Default physical schema name: `sqlmesh__[model schema]`) | dict[string, string] | N | | `environment_suffix_target` | Whether SQLMesh views should append their environment name to the `schema` or `table` - [additional details](../guides/configuration.md#view-schema-override). (Default: `schema`) | string | N | +| `gateway_managed_virtual_layer` | Whether SQLMesh views of the virtual layer will be created by the default gateway or model specified gateways - [additional details](../guides/configuration.md#view-schema-override). (Default: False) | boolean | N | | `environment_catalog_mapping` | A mapping from regular expressions to catalog names. The catalog name is used to determine the target catalog for a given environment. | dict[string, string] | N | | `log_limit` | The default number of logs to keep (Default: `20`) | int | N | diff --git a/sqlmesh/core/config/root.py b/sqlmesh/core/config/root.py index bdb8815f37..84841fce8f 100644 --- a/sqlmesh/core/config/root.py +++ b/sqlmesh/core/config/root.py @@ -72,6 +72,7 @@ class Config(BaseConfig): model_defaults: Default values for model definitions. physical_schema_mapping: A mapping from regular expressions to names of schemas in which physical tables for corresponding models will be placed. environment_suffix_target: Indicates whether to append the environment name to the schema or table name. + gateway_managed_virtual_layer: Whether the models' views in the virtual layer are created by the model-specific gateway rather than the default gateway. environment_catalog_mapping: A mapping from regular expressions to catalog names. The catalog name is used to determine the target catalog for a given environment. default_target_environment: The name of the environment that will be the default target for the `sqlmesh plan` and `sqlmesh run` commands. log_limit: The default number of logs to keep. @@ -110,6 +111,7 @@ class Config(BaseConfig): environment_suffix_target: EnvironmentSuffixTarget = Field( default=EnvironmentSuffixTarget.default ) + gateway_managed_virtual_layer: bool = False environment_catalog_mapping: t.Dict[re.Pattern, str] = {} default_target_environment: str = c.PROD log_limit: int = c.DEFAULT_LOG_LIMIT diff --git a/sqlmesh/core/context.py b/sqlmesh/core/context.py index ec88ecea14..dce6ccaca6 100644 --- a/sqlmesh/core/context.py +++ b/sqlmesh/core/context.py @@ -380,6 +380,8 @@ def __init__( self.pinned_environments = Environment.sanitize_names(self.config.pinned_environments) self.auto_categorize_changes = self.config.plan.auto_categorize_changes self.selected_gateway = gateway or self.config.default_gateway_name + self.gateway_managed_virtual_layer = self.config.gateway_managed_virtual_layer + self.catalogs: t.Dict[str, str] = {} gw_model_defaults = self.config.gateways[self.selected_gateway].model_defaults if gw_model_defaults: @@ -585,6 +587,15 @@ def load(self, update_schemas: bool = True) -> GenericContext[C]: """Load all files in the context's path.""" load_start_ts = time.perf_counter() + # In a multi virtual layer setup we need the catalog of each engine + # since there is a possibility of not a shared catalog between them + if self.gateway_managed_virtual_layer: + self.catalogs = { + name: adapter.default_catalog + for name, adapter in self.engine_adapters.items() + if adapter.default_catalog + } + loaded_projects = [loader.load() for loader in self._loaders] self.dag = DAG() @@ -2212,16 +2223,6 @@ def _model_tables(self) -> t.Dict[str, str]: for fqn, snapshot in self.snapshots.items() } - @property - def _snapshot_gateways(self) -> t.Dict[str, str]: - """Mapping of snapshot name to the gateway if specified in the model.""" - - return { - fqn: snapshot.model.gateway - for fqn, snapshot in self.snapshots.items() - if snapshot.is_model and snapshot.model.gateway - } - @cached_property def engine_adapters(self) -> t.Dict[str, EngineAdapter]: """Returns all the engine adapters for the gateways defined in the configuration.""" @@ -2292,14 +2293,18 @@ def _context_diff( ensure_finalized_snapshots=ensure_finalized_snapshots, diff_rendered=diff_rendered, environment_statements=self._environment_statements, + gateway_managed_virtual_layer=self.gateway_managed_virtual_layer, ) def _run_janitor(self, ignore_ttl: bool = False) -> None: self._cleanup_environments() - expired_snapshots = self.state_sync.delete_expired_snapshots(ignore_ttl=ignore_ttl) + expired_snapshots, snapshot_gateways = self.state_sync.delete_expired_snapshots( + ignore_ttl=ignore_ttl + ) + self.snapshot_evaluator.cleanup( expired_snapshots, - self._snapshot_gateways, + snapshot_gateways, on_complete=self.console.update_cleanup_progress, ) diff --git a/sqlmesh/core/context_diff.py b/sqlmesh/core/context_diff.py index 88d4c5cf32..82d2a52cdc 100644 --- a/sqlmesh/core/context_diff.py +++ b/sqlmesh/core/context_diff.py @@ -53,6 +53,8 @@ class ContextDiff(PydanticModel): """Whether the currently stored environment record is in unfinalized state.""" normalize_environment_name: bool """Whether the environment name should be normalized.""" + gateway_managed_virtual_layer: bool = False + """Whether the virtual layer's views will be created by the model specified gateways.""" create_from: str """The name of the environment the target environment will be created from if new.""" create_from_env_exists: bool @@ -96,6 +98,7 @@ def create( excluded_requirements: t.Optional[t.Set[str]] = None, diff_rendered: bool = False, environment_statements: t.Optional[t.List[EnvironmentStatements]] = [], + gateway_managed_virtual_layer: bool = False, ) -> ContextDiff: """Create a ContextDiff object. @@ -118,7 +121,11 @@ def create( env = state_reader.get_environment(environment) create_from_env_exists = False - if env is None or env.expired: + if ( + env is None + or env.expired + or env.gateway_managed_virtual_layer != gateway_managed_virtual_layer + ): env = state_reader.get_environment(create_from.lower()) if not env and create_from != c.PROD: @@ -226,6 +233,7 @@ def create( diff_rendered=diff_rendered, previous_environment_statements=previous_environment_statements, environment_statements=environment_statements, + gateway_managed_virtual_layer=gateway_managed_virtual_layer, ) @classmethod diff --git a/sqlmesh/core/environment.py b/sqlmesh/core/environment.py index 6afec51f73..95744be343 100644 --- a/sqlmesh/core/environment.py +++ b/sqlmesh/core/environment.py @@ -32,12 +32,15 @@ class EnvironmentNamingInfo(PydanticModel): catalog_name_override: The name of the catalog to use for this environment if an override was provided normalize_name: Indicates whether the environment's name will be normalized. For example, if it's `dev`, then it will become `DEV` when targeting Snowflake. + gateway_managed_virtual_layer: Determines whether the virtual layer's views are created by the model-specific + gateways, otherwise the default gateway is used. Default: False. """ name: str = c.PROD suffix_target: EnvironmentSuffixTarget = Field(default=EnvironmentSuffixTarget.SCHEMA) catalog_name_override: t.Optional[str] = None normalize_name: bool = True + gateway_managed_virtual_layer: bool = False @field_validator("name", mode="before") @classmethod @@ -49,6 +52,11 @@ def _sanitize_name(cls, v: str) -> str: def _validate_normalize_name(cls, v: t.Any) -> bool: return True if v is None else bool(v) + @field_validator("gateway_managed_virtual_layer", mode="before") + @classmethod + def _validate_gateway_managed_virtual_layer(cls, v: t.Any) -> bool: + return False if v is None else bool(v) + @t.overload @classmethod def sanitize_name(cls, v: str) -> str: ... @@ -194,6 +202,7 @@ def naming_info(self) -> EnvironmentNamingInfo: suffix_target=self.suffix_target, catalog_name_override=self.catalog_name_override, normalize_name=self.normalize_name, + gateway_managed_virtual_layer=self.gateway_managed_virtual_layer, ) @property diff --git a/sqlmesh/core/loader.py b/sqlmesh/core/loader.py index edd4911156..e9ee4282f6 100644 --- a/sqlmesh/core/loader.py +++ b/sqlmesh/core/loader.py @@ -468,6 +468,7 @@ def _load() -> t.List[Model]: default_catalog=self.context.default_catalog, infer_names=self.config.model_naming.infer_names, signal_definitions=signals, + catalogs=self.context.catalogs, ) except Exception as ex: raise ConfigError(f"Failed to load model definition at '{path}'.\n{ex}") @@ -525,6 +526,7 @@ def _load_python_models( default_catalog=self.context.default_catalog, infer_names=self.config.model_naming.infer_names, audit_definitions=audits, + catalogs=self.context.catalogs, ): if model.enabled: models[model.fqn] = model diff --git a/sqlmesh/core/model/definition.py b/sqlmesh/core/model/definition.py index 5930d8b835..ab3eeec808 100644 --- a/sqlmesh/core/model/definition.py +++ b/sqlmesh/core/model/definition.py @@ -1907,6 +1907,13 @@ def create_models_from_blueprints( else: gateway_name = None + if ( + (catalogs := loader_kwargs.pop("catalogs", None)) + and gateway_name + and (catalog := catalogs.get(gateway_name)) + ): + loader_kwargs["default_catalog"] = catalog + model_blueprints.append( loader( path=path, diff --git a/sqlmesh/core/plan/builder.py b/sqlmesh/core/plan/builder.py index ccb854d974..bf3e8fd543 100644 --- a/sqlmesh/core/plan/builder.py +++ b/sqlmesh/core/plan/builder.py @@ -151,6 +151,7 @@ def __init__( name=self._context_diff.environment, suffix_target=environment_suffix_target, normalize_name=self._context_diff.normalize_environment_name, + gateway_managed_virtual_layer=self._context_diff.gateway_managed_virtual_layer, ) self._latest_plan: t.Optional[Plan] = None diff --git a/sqlmesh/core/plan/evaluator.py b/sqlmesh/core/plan/evaluator.py index 6faed438f5..685cf3e396 100644 --- a/sqlmesh/core/plan/evaluator.py +++ b/sqlmesh/core/plan/evaluator.py @@ -423,8 +423,16 @@ def _demote_snapshots( environment_naming_info: EnvironmentNamingInfo, on_complete: t.Optional[t.Callable[[SnapshotInfoLike], None]] = None, ) -> None: + # In a multi virtual layer setup we need the gateway info from the snapshots for demotion + snapshots_to_demote: t.List[Snapshot] = [] + if environment_naming_info.gateway_managed_virtual_layer: + removed_snapshots = self.state_sync.get_snapshots(target_snapshots) + snapshots_to_demote = [removed_snapshots[s.snapshot_id] for s in target_snapshots] + self.snapshot_evaluator.demote( - target_snapshots, environment_naming_info, on_complete=on_complete + snapshots_to_demote or target_snapshots, + environment_naming_info, + on_complete=on_complete, ) def _restate(self, plan: EvaluatablePlan, snapshots_by_name: t.Dict[str, Snapshot]) -> None: diff --git a/sqlmesh/core/snapshot/evaluator.py b/sqlmesh/core/snapshot/evaluator.py index 22246bd875..d770584d74 100644 --- a/sqlmesh/core/snapshot/evaluator.py +++ b/sqlmesh/core/snapshot/evaluator.py @@ -226,15 +226,23 @@ def promote( deployability_index: Determines snapshots that are deployable in the context of this promotion. on_complete: A callback to call on each successfully promoted snapshot. """ - self._create_schemas( - [ - s.qualified_view_name.table_for_environment( - environment_naming_info, dialect=self.adapter.dialect + + gateway_by_schema: t.Dict[t.Any, str] = {} + tables: t.List[t.Any] = [] + for snapshot in target_snapshots: + if snapshot.is_model and not snapshot.is_symbolic: + table = snapshot.qualified_view_name.table_for_environment( + environment_naming_info, + dialect=self._get_adapter(snapshot.model_gateway).dialect + if environment_naming_info.gateway_managed_virtual_layer + else self.adapter.dialect, ) - for s in target_snapshots - if s.is_model and not s.is_symbolic - ] - ) + tables.append(table) + if environment_naming_info.gateway_managed_virtual_layer: + table_schema = d.schema_(table.db, catalog=table.catalog) + gateway_by_schema[table_schema] = snapshot.model_gateway or "" + self._create_schemas(tables=tables, gateways=gateway_by_schema) + deployability_index = deployability_index or DeployabilityIndex.all_deployable() with self.concurrent_context(): concurrent_apply_to_snapshots( @@ -923,7 +931,11 @@ def _promote_snapshot( table_mapping: t.Optional[t.Dict[str, str]] = None, ) -> None: if snapshot.is_model: - adapter = self.adapter + adapter = ( + self._get_adapter(snapshot.model_gateway) + if environment_naming_info.gateway_managed_virtual_layer + else self.adapter + ) table_name = snapshot.table_name(deployability_index.is_representative(snapshot)) view_name = snapshot.qualified_view_name.for_environment( environment_naming_info, dialect=adapter.dialect @@ -956,7 +968,12 @@ def _demote_snapshot( environment_naming_info: EnvironmentNamingInfo, on_complete: t.Optional[t.Callable[[SnapshotInfoLike], None]], ) -> None: - adapter = self.adapter + adapter = ( + self._get_adapter(snapshot.model_gateway) + if environment_naming_info.gateway_managed_virtual_layer + and isinstance(snapshot, Snapshot) + else self.adapter + ) view_name = snapshot.qualified_view_name.for_environment( environment_naming_info, dialect=adapter.dialect ) diff --git a/sqlmesh/core/state_sync/base.py b/sqlmesh/core/state_sync/base.py index 771dd94172..d976029f14 100644 --- a/sqlmesh/core/state_sync/base.py +++ b/sqlmesh/core/state_sync/base.py @@ -305,7 +305,7 @@ def delete_snapshots(self, snapshot_ids: t.Iterable[SnapshotIdLike]) -> None: @abc.abstractmethod def delete_expired_snapshots( self, ignore_ttl: bool = False - ) -> t.List[SnapshotTableCleanupTask]: + ) -> t.Tuple[t.List[SnapshotTableCleanupTask], t.Dict[str, str]]: """Removes expired snapshots. Expired snapshots are snapshots that have exceeded their time-to-live diff --git a/sqlmesh/core/state_sync/cache.py b/sqlmesh/core/state_sync/cache.py index ab40f186e9..04f4d67a97 100644 --- a/sqlmesh/core/state_sync/cache.py +++ b/sqlmesh/core/state_sync/cache.py @@ -112,7 +112,7 @@ def delete_snapshots(self, snapshot_ids: t.Iterable[SnapshotIdLike]) -> None: def delete_expired_snapshots( self, ignore_ttl: bool = False - ) -> t.List[SnapshotTableCleanupTask]: + ) -> t.Tuple[t.List[SnapshotTableCleanupTask], t.Dict[str, str]]: self.snapshot_cache.clear() return self.state_sync.delete_expired_snapshots(ignore_ttl=ignore_ttl) diff --git a/sqlmesh/core/state_sync/db/environment.py b/sqlmesh/core/state_sync/db/environment.py index c8f641f354..9fe5216800 100644 --- a/sqlmesh/core/state_sync/db/environment.py +++ b/sqlmesh/core/state_sync/db/environment.py @@ -48,6 +48,7 @@ def __init__( "catalog_name_override": exp.DataType.build("text"), "previous_finalized_snapshots": exp.DataType.build(blob_type), "normalize_name": exp.DataType.build("boolean"), + "gateway_managed_virtual_layer": exp.DataType.build("boolean"), "requirements": exp.DataType.build(blob_type), } @@ -328,6 +329,7 @@ def _environment_to_df(environment: Environment) -> pd.DataFrame: else None ), "normalize_name": environment.normalize_name, + "gateway_managed_virtual_layer": environment.gateway_managed_virtual_layer, "requirements": json.dumps(environment.requirements), } ] diff --git a/sqlmesh/core/state_sync/db/facade.py b/sqlmesh/core/state_sync/db/facade.py index 884955d98e..da16eba1e6 100644 --- a/sqlmesh/core/state_sync/db/facade.py +++ b/sqlmesh/core/state_sync/db/facade.py @@ -199,7 +199,11 @@ def promote( ) != table_infos[name].qualified_view_name.for_environment(environment.naming_info) } - if not existing_environment.expired: + if ( + not existing_environment.expired + and existing_environment.gateway_managed_virtual_layer + == environment.gateway_managed_virtual_layer + ): if environment.previous_plan_id != existing_environment.plan_id: raise ConflictingPlanError( f"Plan '{environment.plan_id}' is no longer valid for the target environment '{environment.name}'. " @@ -273,12 +277,14 @@ def invalidate_environment(self, name: str) -> None: @transactional() def delete_expired_snapshots( self, ignore_ttl: bool = False - ) -> t.List[SnapshotTableCleanupTask]: - expired_snapshot_ids, cleanup_targets = self.snapshot_state.delete_expired_snapshots( - self.environment_state.get_environments(), ignore_ttl=ignore_ttl + ) -> t.Tuple[t.List[SnapshotTableCleanupTask], t.Dict[str, str]]: + expired_snapshot_ids, cleanup_targets, gateways = ( + self.snapshot_state.delete_expired_snapshots( + self.environment_state.get_environments(), ignore_ttl=ignore_ttl + ) ) self.interval_state.cleanup_intervals(cleanup_targets, expired_snapshot_ids) - return cleanup_targets + return cleanup_targets, gateways @transactional() def delete_expired_environments(self) -> t.List[Environment]: diff --git a/sqlmesh/core/state_sync/db/snapshot.py b/sqlmesh/core/state_sync/db/snapshot.py index e46f7d0151..c1f4e1a17c 100644 --- a/sqlmesh/core/state_sync/db/snapshot.py +++ b/sqlmesh/core/state_sync/db/snapshot.py @@ -195,7 +195,7 @@ def unpause_snapshots( def delete_expired_snapshots( self, environments: t.Iterable[Environment], ignore_ttl: bool = False - ) -> t.Tuple[t.Set[SnapshotId], t.List[SnapshotTableCleanupTask]]: + ) -> t.Tuple[t.Set[SnapshotId], t.List[SnapshotTableCleanupTask], t.Dict[str, str]]: """Deletes expired snapshots. Args: @@ -220,7 +220,7 @@ def delete_expired_snapshots( for name, identifier, version in fetchall(self.engine_adapter, expired_query) } if not expired_candidates: - return set(), [] + return set(), [], {} promoted_snapshot_ids = { snapshot.snapshot_id @@ -240,6 +240,7 @@ def _is_snapshot_used(snapshot: SharedVersionSnapshot) -> bool: ) cleanup_targets = [] expired_snapshot_ids = set() + gateways = {} for versions_batch in version_batches: snapshots = self._get_snapshots_with_same_version(versions_batch) @@ -253,6 +254,9 @@ def _is_snapshot_used(snapshot: SharedVersionSnapshot) -> bool: expired_snapshot_ids.update([s.snapshot_id for s in expired_snapshots]) for snapshot in expired_snapshots: + if (node := snapshot.raw_snapshot.get("node")) and (gateway := node.get("gateway")): + gateways[snapshot.snapshot_id.name] = gateway + shared_version_snapshots = snapshots_by_version[(snapshot.name, snapshot.version)] shared_version_snapshots.discard(snapshot.snapshot_id) @@ -271,7 +275,7 @@ def _is_snapshot_used(snapshot: SharedVersionSnapshot) -> bool: if expired_snapshot_ids: self.delete_snapshots(expired_snapshot_ids) - return expired_snapshot_ids, cleanup_targets + return expired_snapshot_ids, cleanup_targets, gateways def delete_snapshots(self, snapshot_ids: t.Iterable[SnapshotIdLike]) -> None: """Deletes snapshots. diff --git a/sqlmesh/migrations/v0056_restore_table_indexes.py b/sqlmesh/migrations/v0056_restore_table_indexes.py index d6fab1669b..9d290e6b3c 100644 --- a/sqlmesh/migrations/v0056_restore_table_indexes.py +++ b/sqlmesh/migrations/v0056_restore_table_indexes.py @@ -1,4 +1,4 @@ -"""Readds indexes and primary keys in case tables were restored from a backup.""" +"""Reads indexes and primary keys in case tables were restored from a backup.""" from sqlglot import exp from sqlmesh.utils import random_id diff --git a/sqlmesh/migrations/v0078_add_gateway_managed_virtual_layer.py b/sqlmesh/migrations/v0078_add_gateway_managed_virtual_layer.py new file mode 100644 index 0000000000..bb43d27a7e --- /dev/null +++ b/sqlmesh/migrations/v0078_add_gateway_managed_virtual_layer.py @@ -0,0 +1,28 @@ +"""Add flag that controls whether the virtual layer's views will be created by the model specified gateway rather than the default gateway.""" + +from sqlglot import exp + + +def migrate(state_sync, **kwargs): # type: ignore + engine_adapter = state_sync.engine_adapter + environments_table = "_environments" + if state_sync.schema: + environments_table = f"{state_sync.schema}.{environments_table}" + + alter_table_exp = exp.Alter( + this=exp.to_table(environments_table), + kind="TABLE", + actions=[ + exp.ColumnDef( + this=exp.to_column("gateway_managed_virtual_layer"), + kind=exp.DataType.build("boolean"), + ) + ], + ) + engine_adapter.execute(alter_table_exp) + + state_sync.engine_adapter.update_table( + environments_table, + {"gateway_managed_virtual_layer": False}, + where=exp.true(), + ) diff --git a/sqlmesh/schedulers/airflow/state_sync.py b/sqlmesh/schedulers/airflow/state_sync.py index fecf15199a..217c66f334 100644 --- a/sqlmesh/schedulers/airflow/state_sync.py +++ b/sqlmesh/schedulers/airflow/state_sync.py @@ -172,7 +172,7 @@ def delete_snapshots(self, snapshot_ids: t.Iterable[SnapshotIdLike]) -> None: def delete_expired_snapshots( self, ignore_ttl: bool = False - ) -> t.List[SnapshotTableCleanupTask]: + ) -> t.Tuple[t.List[SnapshotTableCleanupTask], t.Dict[str, str]]: """Removes expired snapshots. Expired snapshots are snapshots that have exceeded their time-to-live diff --git a/tests/core/test_config.py b/tests/core/test_config.py index f920fc8dc9..2fd4640626 100644 --- a/tests/core/test_config.py +++ b/tests/core/test_config.py @@ -762,12 +762,6 @@ def test_multi_gateway_config(tmp_path, mocker: MockerFixture): ctx = Context(paths=tmp_path, config=config) - mocker.patch.object( - Context, - "_snapshot_gateways", - new_callable=mocker.PropertyMock(return_value={"snapshot": "athena"}), - ) - assert isinstance(ctx._connection_config, RedshiftConnectionConfig) assert len(ctx.engine_adapters) == 2 assert isinstance(ctx.engine_adapters["athena"], AthenaEngineAdapter) diff --git a/tests/core/test_context.py b/tests/core/test_context.py index 7c05dade1e..5712bb8c7b 100644 --- a/tests/core/test_context.py +++ b/tests/core/test_context.py @@ -340,12 +340,6 @@ def test_gateway_specific_adapters(copy_to_temp_path, mocker): assert len(ctx._engine_adapters) == 1 assert ctx.engine_adapter == ctx._engine_adapters["dev"] - mocker.patch.object( - Context, - "_snapshot_gateways", - new_callable=mocker.PropertyMock(return_value={"test_snapshot": "test"}), - ) - ctx = Context(paths=path, config="isolated_systems_config") assert len(ctx.engine_adapters) == 3 diff --git a/tests/schedulers/airflow/test_client.py b/tests/schedulers/airflow/test_client.py index 01ffd214be..92a92dc964 100644 --- a/tests/schedulers/airflow/test_client.py +++ b/tests/schedulers/airflow/test_client.py @@ -170,6 +170,7 @@ def test_apply_plan(mocker: MockerFixture, snapshot: Snapshot): ], "start_at": "2022-01-01", "end_at": "2022-01-01", + "gateway_managed_virtual_layer": False, "plan_id": "test_plan_id", "previous_plan_id": "previous_plan_id", "promoted_snapshot_ids": [ From 846e7a66438e60f5e67419f2f581cdbab3c5dfd8 Mon Sep 17 00:00:00 2001 From: Themis Valtinos <73662635+themisvaltinos@users.noreply.github.com> Date: Tue, 8 Apr 2025 23:07:44 +0300 Subject: [PATCH 02/14] Revise clean up logic --- examples/multi_virtual_layer/audits/.gitkeep | 0 examples/multi_virtual_layer/config.yaml | 28 ++++++ examples/multi_virtual_layer/macros/.gitkeep | 0 .../multi_virtual_layer/macros/__init__.py | 6 ++ examples/multi_virtual_layer/models/.gitkeep | 0 .../models/local_schema/model_one.sql | 8 ++ .../models/local_schema/model_two.sql | 9 ++ .../models/memory_schema/model_one.sql | 9 ++ .../models/memory_schema/model_two.sql | 10 +++ examples/multi_virtual_layer/tests/.gitkeep | 0 sqlmesh/core/context.py | 68 ++++++++++---- sqlmesh/core/state_sync/base.py | 46 +++++++++- sqlmesh/core/state_sync/cache.py | 2 +- sqlmesh/core/state_sync/common.py | 48 +++++++--- sqlmesh/core/state_sync/db/environment.py | 52 ++++++++--- sqlmesh/core/state_sync/db/facade.py | 37 ++++++-- sqlmesh/core/state_sync/db/snapshot.py | 31 +++++-- sqlmesh/schedulers/airflow/state_sync.py | 49 +++++++++- tests/core/test_context.py | 50 ++++++----- tests/core/test_integration.py | 79 ++++++++++++++++ tests/core/test_snapshot_evaluator.py | 89 ++++++++++++++++++- 21 files changed, 543 insertions(+), 78 deletions(-) create mode 100644 examples/multi_virtual_layer/audits/.gitkeep create mode 100644 examples/multi_virtual_layer/config.yaml create mode 100644 examples/multi_virtual_layer/macros/.gitkeep create mode 100644 examples/multi_virtual_layer/macros/__init__.py create mode 100644 examples/multi_virtual_layer/models/.gitkeep create mode 100644 examples/multi_virtual_layer/models/local_schema/model_one.sql create mode 100644 examples/multi_virtual_layer/models/local_schema/model_two.sql create mode 100644 examples/multi_virtual_layer/models/memory_schema/model_one.sql create mode 100644 examples/multi_virtual_layer/models/memory_schema/model_two.sql create mode 100644 examples/multi_virtual_layer/tests/.gitkeep diff --git a/examples/multi_virtual_layer/audits/.gitkeep b/examples/multi_virtual_layer/audits/.gitkeep new file mode 100644 index 0000000000..e69de29bb2 diff --git a/examples/multi_virtual_layer/config.yaml b/examples/multi_virtual_layer/config.yaml new file mode 100644 index 0000000000..483472f16b --- /dev/null +++ b/examples/multi_virtual_layer/config.yaml @@ -0,0 +1,28 @@ +gateways: + local: + connection: + type: duckdb + database: db.duckdb + variables: + overriden_var: 'gateway_1' + memory: + connection: + type: duckdb + variables: + overriden_var: 'gateway_2' + +default_gateway: local + +model_defaults: + dialect: 'duckdb' + +model_naming: + infer_names: True + +gateway_managed_virtual_layer: True + +variables: + overriden_var: 'global' + global_one: 88 + + diff --git a/examples/multi_virtual_layer/macros/.gitkeep b/examples/multi_virtual_layer/macros/.gitkeep new file mode 100644 index 0000000000..e69de29bb2 diff --git a/examples/multi_virtual_layer/macros/__init__.py b/examples/multi_virtual_layer/macros/__init__.py new file mode 100644 index 0000000000..1e9de0aa75 --- /dev/null +++ b/examples/multi_virtual_layer/macros/__init__.py @@ -0,0 +1,6 @@ +from sqlmesh import macro + + +@macro() +def one(context): + return 1 diff --git a/examples/multi_virtual_layer/models/.gitkeep b/examples/multi_virtual_layer/models/.gitkeep new file mode 100644 index 0000000000..e69de29bb2 diff --git a/examples/multi_virtual_layer/models/local_schema/model_one.sql b/examples/multi_virtual_layer/models/local_schema/model_one.sql new file mode 100644 index 0000000000..1bb062a80b --- /dev/null +++ b/examples/multi_virtual_layer/models/local_schema/model_one.sql @@ -0,0 +1,8 @@ +MODEL ( + kind FULL, +); + +SELECT + @overriden_var as item_id, + @global_one as global_one, + @one() AS macro_one \ No newline at end of file diff --git a/examples/multi_virtual_layer/models/local_schema/model_two.sql b/examples/multi_virtual_layer/models/local_schema/model_two.sql new file mode 100644 index 0000000000..93b927eee8 --- /dev/null +++ b/examples/multi_virtual_layer/models/local_schema/model_two.sql @@ -0,0 +1,9 @@ +MODEL ( + kind FULL, +); + +SELECT + item_id, + global_one +FROM + local_schema.model_one; \ No newline at end of file diff --git a/examples/multi_virtual_layer/models/memory_schema/model_one.sql b/examples/multi_virtual_layer/models/memory_schema/model_one.sql new file mode 100644 index 0000000000..9c0fac206a --- /dev/null +++ b/examples/multi_virtual_layer/models/memory_schema/model_one.sql @@ -0,0 +1,9 @@ +MODEL ( + kind FULL, + gateway memory +); + +SELECT + @overriden_var as item_id, + @global_one as global_one, + @one() AS macro_one diff --git a/examples/multi_virtual_layer/models/memory_schema/model_two.sql b/examples/multi_virtual_layer/models/memory_schema/model_two.sql new file mode 100644 index 0000000000..c249aa37f6 --- /dev/null +++ b/examples/multi_virtual_layer/models/memory_schema/model_two.sql @@ -0,0 +1,10 @@ +MODEL ( + kind FULL, + gateway memory +); + +SELECT + item_id, + global_one +FROM + memory_schema.model_one; \ No newline at end of file diff --git a/examples/multi_virtual_layer/tests/.gitkeep b/examples/multi_virtual_layer/tests/.gitkeep new file mode 100644 index 0000000000..e69de29bb2 diff --git a/sqlmesh/core/context.py b/sqlmesh/core/context.py index dce6ccaca6..a4862deafa 100644 --- a/sqlmesh/core/context.py +++ b/sqlmesh/core/context.py @@ -364,6 +364,7 @@ def __init__( self._environment_statements: t.List[EnvironmentStatements] = [] self._excluded_requirements: t.Set[str] = set() self._default_catalog: t.Optional[str] = None + self._catalogs: t.Dict[str, str] = {} self._linters: t.Dict[str, Linter] = {} self._loaded: bool = False @@ -381,7 +382,6 @@ def __init__( self.auto_categorize_changes = self.config.plan.auto_categorize_changes self.selected_gateway = gateway or self.config.default_gateway_name self.gateway_managed_virtual_layer = self.config.gateway_managed_virtual_layer - self.catalogs: t.Dict[str, str] = {} gw_model_defaults = self.config.gateways[self.selected_gateway].model_defaults if gw_model_defaults: @@ -587,15 +587,6 @@ def load(self, update_schemas: bool = True) -> GenericContext[C]: """Load all files in the context's path.""" load_start_ts = time.perf_counter() - # In a multi virtual layer setup we need the catalog of each engine - # since there is a possibility of not a shared catalog between them - if self.gateway_managed_virtual_layer: - self.catalogs = { - name: adapter.default_catalog - for name, adapter in self.engine_adapters.items() - if adapter.default_catalog - } - loaded_projects = [loader.load() for loader in self._loaders] self.dag = DAG() @@ -2233,6 +2224,17 @@ def engine_adapters(self) -> t.Dict[str, EngineAdapter]: self._engine_adapters[gateway_name] = adapter return self._engine_adapters + @cached_property + def catalogs(self) -> t.Dict[str, str]: + """Returns the catalogs for each engine adapter in a multi virtual layer setup when the catalog isn't shared.""" + if self.gateway_managed_virtual_layer: + self._catalogs = { + name: adapter.default_catalog + for name, adapter in self.engine_adapters.items() + if adapter.default_catalog + } + return self._catalogs + def _get_engine_adapter(self, gateway: t.Optional[str] = None) -> EngineAdapter: if gateway: if adapter := self.engine_adapters.get(gateway): @@ -2297,22 +2299,52 @@ def _context_diff( ) def _run_janitor(self, ignore_ttl: bool = False) -> None: - self._cleanup_environments() - expired_snapshots, snapshot_gateways = self.state_sync.delete_expired_snapshots( - ignore_ttl=ignore_ttl + # Get expired environments and removes their views and schemas + expired_environments, filter_expr = self._cleanup_environments() + + # Get expired snapshots and corresponding gateways per snapshot when applied + expired_snapshots_ids, cleanup_targets, snapshot_gateways = ( + self.state_sync.get_expired_snapshots(ignore_ttl=ignore_ttl) ) + # Clean up intervals from the state sync + self.state_sync.cleanup_intervals(cleanup_targets, expired_snapshots_ids) + + # Remove the expired snapshots tables self.snapshot_evaluator.cleanup( - expired_snapshots, - snapshot_gateways, + target_snapshots=cleanup_targets, + snapshot_gateways=snapshot_gateways, on_complete=self.console.update_cleanup_progress, ) + # Finally, remove the expired environments and snapshots from the state sync + self.state_sync.delete_environments(expired_environments, filter_expr) + self.state_sync.delete_snapshots(expired_snapshots_ids) self.state_sync.compact_intervals() - def _cleanup_environments(self) -> None: - expired_environments = self.state_sync.delete_expired_environments() - cleanup_expired_views(self.engine_adapter, expired_environments, console=self.console) + def _cleanup_environments(self) -> t.Tuple[t.List[Environment], exp.LTE]: + expired_environments, filter_expr = self.state_sync.get_expired_environments() + + environment_snapshot_adapters: t.Dict[str, t.Dict[str, EngineAdapter]] = {} + for environment in expired_environments: + snapshot_adapters: t.Dict[str, EngineAdapter] = {} + if environment.gateway_managed_virtual_layer: + snapshots = self.state_sync.get_snapshots(environment.snapshots) + for snapshot_id, snapshot in snapshots.items(): + if snapshot.is_model and not snapshot.is_symbolic: + snapshot_adapters[snapshot_id.name] = self._get_engine_adapter( + snapshot.model_gateway + ) + environment_snapshot_adapters[environment.name] = snapshot_adapters + + cleanup_expired_views( + adapter=self.engine_adapter, + environments=expired_environments, + console=self.console, + environment_snapshot_adapters=environment_snapshot_adapters, + ) + + return expired_environments, filter_expr def _try_connection(self, connection_name: str, validator: t.Callable[[], None]) -> None: connection_name = connection_name.capitalize() diff --git a/sqlmesh/core/state_sync/base.py b/sqlmesh/core/state_sync/base.py index d976029f14..44471a41a9 100644 --- a/sqlmesh/core/state_sync/base.py +++ b/sqlmesh/core/state_sync/base.py @@ -7,6 +7,7 @@ import typing as t from sqlglot import __version__ as SQLGLOT_VERSION +from sqlglot import exp from sqlmesh import migrations from sqlmesh.core.environment import Environment, EnvironmentNamingInfo, EnvironmentStatements @@ -305,7 +306,7 @@ def delete_snapshots(self, snapshot_ids: t.Iterable[SnapshotIdLike]) -> None: @abc.abstractmethod def delete_expired_snapshots( self, ignore_ttl: bool = False - ) -> t.Tuple[t.List[SnapshotTableCleanupTask], t.Dict[str, str]]: + ) -> t.List[SnapshotTableCleanupTask]: """Removes expired snapshots. Expired snapshots are snapshots that have exceeded their time-to-live @@ -319,6 +320,23 @@ def delete_expired_snapshots( The list of table cleanup tasks. """ + @abc.abstractmethod + def get_expired_snapshots( + self, ignore_ttl: bool = False + ) -> t.Tuple[t.Set[SnapshotId], t.List[SnapshotTableCleanupTask], t.Dict[str, str]]: + """Gets expired snapshots. + + Expired snapshots are snapshots that have exceeded their time-to-live + and are no longer in use within an environment. + + Args: + ignore_ttl: Ignore the TTL on the snapshot when considering it expired. This has the effect of deleting + all snapshots that are not referenced in any environment + + Returns: + A tuple of expired snapshot IDs, cleanup targets and gateway per snapshot dictionary. + """ + @abc.abstractmethod def invalidate_environment(self, name: str) -> None: """Invalidates the target environment by setting its expiration timestamp to now. @@ -327,6 +345,14 @@ def invalidate_environment(self, name: str) -> None: name: The name of the environment to invalidate. """ + @abc.abstractmethod + def cleanup_intervals( + self, + cleanup_targets: t.List[SnapshotTableCleanupTask], + expired_snapshot_ids: t.Set[SnapshotId], + ) -> None: + """Cleans up intervals.""" + @abc.abstractmethod def remove_intervals( self, @@ -395,6 +421,24 @@ def delete_expired_environments(self) -> t.List[Environment]: The list of removed environments. """ + @abc.abstractmethod + def get_expired_environments(self) -> t.Tuple[t.List[Environment], exp.LTE]: + """Returns the expired environments. + + Expired environments are environments that have exceeded their time-to-live value. + + Returns: + The list of environments to remove, the filter to remove environments. + """ + + @abc.abstractmethod + def delete_environments(self, environments: t.List[Environment], filter_expr: exp.LTE) -> None: + """Removes environments. + + Returns: + The list of removed environments. + """ + @abc.abstractmethod def unpause_snapshots( self, snapshots: t.Collection[SnapshotInfoLike], unpaused_dt: TimeLike diff --git a/sqlmesh/core/state_sync/cache.py b/sqlmesh/core/state_sync/cache.py index 04f4d67a97..ab40f186e9 100644 --- a/sqlmesh/core/state_sync/cache.py +++ b/sqlmesh/core/state_sync/cache.py @@ -112,7 +112,7 @@ def delete_snapshots(self, snapshot_ids: t.Iterable[SnapshotIdLike]) -> None: def delete_expired_snapshots( self, ignore_ttl: bool = False - ) -> t.Tuple[t.List[SnapshotTableCleanupTask], t.Dict[str, str]]: + ) -> t.List[SnapshotTableCleanupTask]: self.snapshot_cache.clear() return self.state_sync.delete_expired_snapshots(ignore_ttl=ignore_ttl) diff --git a/sqlmesh/core/state_sync/common.py b/sqlmesh/core/state_sync/common.py index 90ab67989c..ae7d36cc09 100644 --- a/sqlmesh/core/state_sync/common.py +++ b/sqlmesh/core/state_sync/common.py @@ -21,7 +21,10 @@ def cleanup_expired_views( - adapter: EngineAdapter, environments: t.List[Environment], console: t.Optional[Console] = None + adapter: EngineAdapter, + environments: t.List[Environment], + environment_snapshot_adapters: t.Optional[t.Dict[str, t.Dict[str, EngineAdapter]]] = None, + console: t.Optional[Console] = None, ) -> None: expired_schema_environments = [ environment for environment in environments if environment.suffix_target.is_schema @@ -29,13 +32,25 @@ def cleanup_expired_views( expired_table_environments = [ environment for environment in environments if environment.suffix_target.is_table ] - for expired_catalog, expired_schema in { + + # Drop the schemas for the expired environments + # Note: We have to use the corresponding adapter if it is a gateway managed virtual layer + for engine_adapter, expired_catalog, expired_schema in { ( + ( + engine_adapter := ( + (environment_dict.get(snapshot.name) or adapter) + if environment.gateway_managed_virtual_layer + and environment_snapshot_adapters + and (environment_dict := environment_snapshot_adapters.get(environment.name)) + else adapter + ) + ), snapshot.qualified_view_name.catalog_for_environment( - environment.naming_info, dialect=adapter.dialect + environment.naming_info, dialect=engine_adapter.dialect ), snapshot.qualified_view_name.schema_for_environment( - environment.naming_info, dialect=adapter.dialect + environment.naming_info, dialect=engine_adapter.dialect ), ) for environment in expired_schema_environments @@ -44,27 +59,40 @@ def cleanup_expired_views( }: schema = schema_(expired_schema, expired_catalog) try: - adapter.drop_schema( + engine_adapter.drop_schema( schema, ignore_if_not_exists=True, cascade=True, ) if console: - console.update_cleanup_progress(schema.sql(dialect=adapter.dialect)) + console.update_cleanup_progress(schema.sql(dialect=engine_adapter.dialect)) except Exception as e: raise SQLMeshError( f"Failed to drop the expired environment schema '{schema}': {e}" ) from e - for expired_view in { - snapshot.qualified_view_name.for_environment( - environment.naming_info, dialect=adapter.dialect + + # Drop the views for the expired environments + for engine_adapter, expired_view in { + ( + ( + engine_adapter := ( + (environment_dict.get(snapshot.name) or adapter) + if environment.gateway_managed_virtual_layer + and environment_snapshot_adapters + and (environment_dict := environment_snapshot_adapters.get(environment.name)) + else adapter + ) + ), + snapshot.qualified_view_name.for_environment( + environment.naming_info, dialect=engine_adapter.dialect + ), ) for environment in expired_table_environments for snapshot in environment.snapshots if snapshot.is_model and not snapshot.is_symbolic }: try: - adapter.drop_view(expired_view, ignore_if_not_exists=True) + engine_adapter.drop_view(expired_view, ignore_if_not_exists=True) if console: console.update_cleanup_progress(expired_view) except Exception as e: diff --git a/sqlmesh/core/state_sync/db/environment.py b/sqlmesh/core/state_sync/db/environment.py index 9fe5216800..3a4cbfa8c6 100644 --- a/sqlmesh/core/state_sync/db/environment.py +++ b/sqlmesh/core/state_sync/db/environment.py @@ -168,21 +168,20 @@ def delete_expired_environments(self) -> t.List[Environment]: Returns: A list of deleted environments. """ - now_ts = now_timestamp() - filter_expr = exp.LTE( - this=exp.column("expiration_ts"), - expression=exp.Literal.number(now_ts), - ) + environments, filter_expr = self.get_expired_environments() + self.delete_environments(environments=environments, filter_expr=filter_expr) + return environments - rows = fetchall( - self.engine_adapter, - self._environments_query( - where=filter_expr, - lock_for_update=True, - ), - ) - environments = [self._environment_from_row(r) for r in rows] + def delete_environments(self, environments: t.List[Environment], filter_expr: exp.LTE) -> None: + """Deletes environments and corresponding environment statements. + Returns: + The list of deleted environments. + """ + if not environments: + return + + # Delete the expired environments self.engine_adapter.delete_from( self.environments_table, where=filter_expr, @@ -199,7 +198,32 @@ def delete_expired_environments(self) -> t.List[Environment]: where=exp.or_(*expired_environments), ) - return environments + def get_expired_environments(self) -> t.Tuple[t.List[Environment], exp.LTE]: + """Fetches expired environments along with their filter expression and snapshot gateways. + + Returns: + A tuple containing: + - A list of expired environments. + - The filter expression used to identify expired environments. + """ + + now_ts = now_timestamp() + filter_expr = exp.LTE( + this=exp.column("expiration_ts"), + expression=exp.Literal.number(now_ts), + ) + + rows = fetchall( + self.engine_adapter, + self._environments_query( + where=filter_expr, + lock_for_update=True, + ), + ) + + environments = [self._environment_from_row(r) for r in rows] + + return environments, filter_expr def get_environments(self) -> t.List[Environment]: """Fetches all environments. diff --git a/sqlmesh/core/state_sync/db/facade.py b/sqlmesh/core/state_sync/db/facade.py index da16eba1e6..0fe1a74749 100644 --- a/sqlmesh/core/state_sync/db/facade.py +++ b/sqlmesh/core/state_sync/db/facade.py @@ -277,19 +277,44 @@ def invalidate_environment(self, name: str) -> None: @transactional() def delete_expired_snapshots( self, ignore_ttl: bool = False - ) -> t.Tuple[t.List[SnapshotTableCleanupTask], t.Dict[str, str]]: - expired_snapshot_ids, cleanup_targets, gateways = ( - self.snapshot_state.delete_expired_snapshots( - self.environment_state.get_environments(), ignore_ttl=ignore_ttl - ) + ) -> t.List[SnapshotTableCleanupTask]: + expired_snapshot_ids, cleanup_targets = self.snapshot_state.delete_expired_snapshots( + self.environment_state.get_environments(), ignore_ttl=ignore_ttl ) self.interval_state.cleanup_intervals(cleanup_targets, expired_snapshot_ids) - return cleanup_targets, gateways + return cleanup_targets @transactional() def delete_expired_environments(self) -> t.List[Environment]: return self.environment_state.delete_expired_environments() + @transactional() + def delete_environments(self, environments: t.List[Environment], filter_expr: exp.LTE) -> None: + self.environment_state.delete_environments( + environments=environments, filter_expr=filter_expr + ) + + @transactional() + def cleanup_intervals( + self, + cleanup_targets: t.List[SnapshotTableCleanupTask], + expired_snapshot_ids: t.Set[SnapshotId], + ) -> None: + self.interval_state.cleanup_intervals(cleanup_targets, expired_snapshot_ids) + + def get_expired_snapshots( + self, ignore_ttl: bool = False + ) -> t.Tuple[t.Set[SnapshotId], t.List[SnapshotTableCleanupTask], t.Dict[str, str]]: + expired_snapshot_ids, cleanup_targets, snapshot_gateways = ( + self.snapshot_state.get_expired_snapshots( + self.environment_state.get_environments(), ignore_ttl=ignore_ttl + ) + ) + return expired_snapshot_ids, cleanup_targets, snapshot_gateways + + def get_expired_environments(self) -> t.Tuple[t.List[Environment], exp.LTE]: + return self.environment_state.get_expired_environments() + def delete_snapshots(self, snapshot_ids: t.Iterable[SnapshotIdLike]) -> None: self.snapshot_state.delete_snapshots(snapshot_ids) diff --git a/sqlmesh/core/state_sync/db/snapshot.py b/sqlmesh/core/state_sync/db/snapshot.py index c1f4e1a17c..8c5c4ac33a 100644 --- a/sqlmesh/core/state_sync/db/snapshot.py +++ b/sqlmesh/core/state_sync/db/snapshot.py @@ -195,7 +195,7 @@ def unpause_snapshots( def delete_expired_snapshots( self, environments: t.Iterable[Environment], ignore_ttl: bool = False - ) -> t.Tuple[t.Set[SnapshotId], t.List[SnapshotTableCleanupTask], t.Dict[str, str]]: + ) -> t.Tuple[t.Set[SnapshotId], t.List[SnapshotTableCleanupTask]]: """Deletes expired snapshots. Args: @@ -204,6 +204,27 @@ def delete_expired_snapshots( Returns: A tuple of expired snapshot IDs and cleanup targets. """ + + expired_snapshot_ids, cleanup_targets, _ = self.get_expired_snapshots( + environments=environments, ignore_ttl=ignore_ttl + ) + + if expired_snapshot_ids: + self.delete_snapshots(expired_snapshot_ids) + + return expired_snapshot_ids, cleanup_targets + + def get_expired_snapshots( + self, environments: t.Iterable[Environment], ignore_ttl: bool = False + ) -> t.Tuple[t.Set[SnapshotId], t.List[SnapshotTableCleanupTask], t.Dict[str, str]]: + """Gets the expired snapshots. + + Args: + ignore_ttl: Whether to ignore the TTL of the snapshots. + + Returns: + A tuple of expired snapshot IDs, cleanup targets and gateway per snapshot dictionary. + """ current_ts = now_timestamp(minute_floor=False) expired_query = exp.select("name", "identifier", "version").from_(self.snapshots_table) @@ -240,7 +261,7 @@ def _is_snapshot_used(snapshot: SharedVersionSnapshot) -> bool: ) cleanup_targets = [] expired_snapshot_ids = set() - gateways = {} + snapshot_gateways: t.Dict[str, str] = {} for versions_batch in version_batches: snapshots = self._get_snapshots_with_same_version(versions_batch) @@ -255,7 +276,7 @@ def _is_snapshot_used(snapshot: SharedVersionSnapshot) -> bool: for snapshot in expired_snapshots: if (node := snapshot.raw_snapshot.get("node")) and (gateway := node.get("gateway")): - gateways[snapshot.snapshot_id.name] = gateway + snapshot_gateways[snapshot.snapshot_id.name] = gateway shared_version_snapshots = snapshots_by_version[(snapshot.name, snapshot.version)] shared_version_snapshots.discard(snapshot.snapshot_id) @@ -273,9 +294,7 @@ def _is_snapshot_used(snapshot: SharedVersionSnapshot) -> bool: ) ) - if expired_snapshot_ids: - self.delete_snapshots(expired_snapshot_ids) - return expired_snapshot_ids, cleanup_targets, gateways + return expired_snapshot_ids, cleanup_targets, snapshot_gateways def delete_snapshots(self, snapshot_ids: t.Iterable[SnapshotIdLike]) -> None: """Deletes snapshots. diff --git a/sqlmesh/schedulers/airflow/state_sync.py b/sqlmesh/schedulers/airflow/state_sync.py index 217c66f334..7df3d855b1 100644 --- a/sqlmesh/schedulers/airflow/state_sync.py +++ b/sqlmesh/schedulers/airflow/state_sync.py @@ -3,6 +3,7 @@ import logging import typing as t +from sqlglot import exp from sqlmesh.core.console import Console from sqlmesh.core.environment import Environment, EnvironmentStatements from sqlmesh.core.snapshot import ( @@ -172,7 +173,7 @@ def delete_snapshots(self, snapshot_ids: t.Iterable[SnapshotIdLike]) -> None: def delete_expired_snapshots( self, ignore_ttl: bool = False - ) -> t.Tuple[t.List[SnapshotTableCleanupTask], t.Dict[str, str]]: + ) -> t.List[SnapshotTableCleanupTask]: """Removes expired snapshots. Expired snapshots are snapshots that have exceeded their time-to-live @@ -293,6 +294,41 @@ def delete_expired_environments(self) -> t.List[Environment]: "Deleting expired environments is not supported by the Airflow state sync." ) + def get_expired_environments(self) -> t.Tuple[t.List[Environment], exp.LTE]: + """Returns the expired environments. + + Expired environments are environments that have exceeded their time-to-live value. + + """ + raise NotImplementedError( + "Getting expired environments is not supported by the Airflow state sync." + ) + + def delete_environments(self, environments: t.List[Environment], filter_expr: exp.LTE) -> None: + """Removes environments. + + Returns: + The list of removed environments. + """ + raise NotImplementedError( + "Deleting environments is not supported by the Airflow state sync." + ) + + def get_expired_snapshots( + self, ignore_ttl: bool = False + ) -> t.Tuple[t.Set[SnapshotId], t.List[SnapshotTableCleanupTask], t.Dict[str, str]]: + """Gets expired snapshots. + + Expired snapshots are snapshots that have exceeded their time-to-live + and are no longer in use within an environment. + + Returns: + A tuple of expired snapshot IDs, cleanup targets and gateway per snapshot dictionary. + """ + raise NotImplementedError( + "Getting expired snapshots is not supported by the Airflow state sync." + ) + def unpause_snapshots( self, snapshots: t.Collection[SnapshotInfoLike], unpaused_dt: TimeLike ) -> None: @@ -318,6 +354,17 @@ def compact_intervals(self) -> None: "Compacting intervals is not supported by the Airflow state sync." ) + def cleanup_intervals( + self, + cleanup_targets: t.List[SnapshotTableCleanupTask], + expired_snapshot_ids: t.Set[SnapshotId], + ) -> None: + """Cleans up intervals.""" + + raise NotImplementedError( + "Cleaning up intervals is not supported by the Airflow state sync." + ) + def update_auto_restatements( self, next_auto_restatement_ts: t.Dict[SnapshotNameVersion, t.Optional[int]] ) -> None: diff --git a/tests/core/test_context.py b/tests/core/test_context.py index 5712bb8c7b..16cc969a4c 100644 --- a/tests/core/test_context.py +++ b/tests/core/test_context.py @@ -43,6 +43,7 @@ from sqlmesh.utils.date import ( make_inclusive_end, now, + now_timestamp, to_date, to_timestamp, yesterday_ds, @@ -779,28 +780,37 @@ def test_janitor(sushi_context, mocker: MockerFixture) -> None: adapter_mock = mocker.MagicMock() adapter_mock.dialect = "duckdb" state_sync_mock = mocker.MagicMock() - state_sync_mock.delete_expired_environments.return_value = [ - Environment( - name="test_environment", - suffix_target=EnvironmentSuffixTarget.TABLE, - snapshots=[x.table_info for x in sushi_context.snapshots.values()], - start_at="2022-01-01", - end_at="2022-01-01", - plan_id="test_plan_id", - previous_plan_id="test_plan_id", - ), - Environment( - name="test_environment", - suffix_target=EnvironmentSuffixTarget.SCHEMA, - snapshots=[x.table_info for x in sushi_context.snapshots.values()], - start_at="2022-01-01", - end_at="2022-01-01", - plan_id="test_plan_id", - previous_plan_id="test_plan_id", - ), - ] + now_ts = now_timestamp() + filter_expr = exp.LTE( + this=exp.column("expiration_ts"), + expression=exp.Literal.number(now_ts), + ) + state_sync_mock.get_expired_environments.return_value = ( + [ + Environment( + name="test_environment", + suffix_target=EnvironmentSuffixTarget.TABLE, + snapshots=[x.table_info for x in sushi_context.snapshots.values()], + start_at="2022-01-01", + end_at="2022-01-01", + plan_id="test_plan_id", + previous_plan_id="test_plan_id", + ), + Environment( + name="test_environment", + suffix_target=EnvironmentSuffixTarget.SCHEMA, + snapshots=[x.table_info for x in sushi_context.snapshots.values()], + start_at="2022-01-01", + end_at="2022-01-01", + plan_id="test_plan_id", + previous_plan_id="test_plan_id", + ), + ], + filter_expr, + ) sushi_context._engine_adapters = {sushi_context.config.default_gateway: adapter_mock} sushi_context._state_sync = state_sync_mock + state_sync_mock.get_expired_snapshots.return_value = (set({}), [], {}) sushi_context._run_janitor() # Assert that the schemas are dropped just twice for the schema based environment # Make sure that external model schemas/tables are not dropped diff --git a/tests/core/test_integration.py b/tests/core/test_integration.py index 5376e18b16..5af1de83c4 100644 --- a/tests/core/test_integration.py +++ b/tests/core/test_integration.py @@ -10,6 +10,7 @@ import pandas as pd import pytest from pathlib import Path +import os import time_machine from pytest_mock.plugin import MockerFixture from sqlglot import exp @@ -4494,6 +4495,84 @@ def test_multi(mocker): ] +@use_terminal_console +def test_multi_virtual_layer(mocker): + context = Context(paths=["examples/multi_virtual_layer"]) + + local_db = "db.duckdb" + if os.path.exists(local_db): + os.remove(local_db) + + # For the model without gateway the default should be used and the gateway variable should overide the global + assert ( + context.render("local_schema.model_one").sql() + == 'SELECT \'gateway_1\' AS "item_id", 88 AS "global_one", 1 AS "macro_one"' + ) + + # For model with gateway specified the appropriate variable should be used to overide + assert ( + context.render("memory.memory_schema.model_one").sql() + == 'SELECT \'gateway_2\' AS "item_id", 88 AS "global_one", 1 AS "macro_one"' + ) + + # context._new_state_sync().reset(default_catalog=context.default_catalog) + plan = context.plan_builder().build() + assert len(plan.new_snapshots) == 4 + context.apply(plan) + + # Validate the tables that source from the first tables are correct as well with evaluate + assert ( + context.evaluate( + "local_schema.model_two", start=now(), end=now(), execution_time=now() + ).to_string() + == " item_id global_one\n0 gateway_1 88" + ) + assert ( + context.evaluate( + "memory.memory_schema.model_two", start=now(), end=now(), execution_time=now() + ).to_string() + == " item_id global_one\n0 gateway_2 88" + ) + + assert sorted(set(snapshot.name for snapshot in plan.directly_modified)) == [ + '"db"."local_schema"."model_one"', + '"db"."local_schema"."model_two"', + '"memory"."memory_schema"."model_one"', + '"memory"."memory_schema"."model_two"', + ] + + model = context.get_model("memory.memory_schema.model_one") + + context.upsert_model(model.copy(update={"query": model.query.select("'c' AS extra")})) + plan = context.plan_builder().build() + context.apply(plan) + + state_environments = context.state_reader.get_environments() + state_snapshots = context.state_reader.get_snapshots(context.snapshots.values()) + + assert state_environments[0].gateway_managed_virtual_layer + assert len(state_snapshots) == len(state_environments[0].snapshots) + + assert [snapshot.name for snapshot in plan.directly_modified] == [ + '"memory"."memory_schema"."model_one"' + ] + assert [x.name for x in list(plan.indirectly_modified.values())[0]] == [ + '"memory"."memory_schema"."model_two"' + ] + + assert len(plan.missing_intervals) == 1 + + assert ( + context.evaluate( + "memory.memory_schema.model_one", start=now(), end=now(), execution_time=now() + ).to_string() + == " item_id global_one macro_one extra\n0 gateway_2 88 1 c" + ) + + if os.path.exists(local_db): + os.remove(local_db) + + def test_multi_dbt(mocker): context = Context(paths=["examples/multi_dbt/bronze", "examples/multi_dbt/silver"]) context._new_state_sync().reset(default_catalog=context.default_catalog) diff --git a/tests/core/test_snapshot_evaluator.py b/tests/core/test_snapshot_evaluator.py index 5cc4364fc5..8eb8d8087f 100644 --- a/tests/core/test_snapshot_evaluator.py +++ b/tests/core/test_snapshot_evaluator.py @@ -4038,7 +4038,8 @@ def model_with_statements(context, **kwargs): assert len(create_args) == 1 assert create_args[0][0] == (f"sqlmesh__db.db__multi_engine_test_model__{snapshot.version}",) - evaluator.promote([snapshot], EnvironmentNamingInfo(name="test_env")) + environment_naming_info = EnvironmentNamingInfo(name="test_env") + evaluator.promote([snapshot], environment_naming_info) # Verify that the default gateway creates the view for the virtual layer engine_adapters["secondary"].create_view.assert_not_called() @@ -4053,3 +4054,89 @@ def model_with_statements(context, **kwargs): # Validate that the get_catalog_type method was called only on the secondary engine from the macro evaluator engine_adapters["default"].get_catalog_type.assert_not_called() assert len(engine_adapters["secondary"].get_catalog_type.call_args_list) == 2 + + evaluator.demote([snapshot], environment_naming_info) + engine_adapters["default"].drop_view.assert_called_once_with( + "db__test_env.multi_engine_test_model", + cascade=False, + ) + + environment_naming_info_gw = EnvironmentNamingInfo( + name="test_env", gateway_managed_virtual_layer=True + ) + # Validate that promoting with gateway_managed_virtual_layer leads to this gateway being used for virtual layer + evaluator.promote([snapshot], environment_naming_info_gw) + view_args = engine_adapters["secondary"].create_view.call_args_list + assert len(view_args) == 1 + assert view_args[0][0][0] == "db__test_env.multi_engine_test_model" + + # Similarly for demotion + evaluator.demote([snapshot], environment_naming_info_gw) + engine_adapters["secondary"].drop_view.assert_called_once_with( + "db__test_env.multi_engine_test_model", + cascade=False, + ) + + +def test_multiple_engine_virtual_layer(snapshot: Snapshot, adapters, make_snapshot): + engine_adapters = {"default": adapters[0], "secondary": adapters[1], "third": adapters[2]} + evaluator = SnapshotEvaluator(engine_adapters) + + model = load_sql_based_model( + parse( # type: ignore + """ + MODEL ( + name test_schema.test_model, + kind FULL, + gateway secondary, + dialect postgres, + ); + SELECT a::int FROM tbl; + CREATE INDEX IF NOT EXISTS test_idx ON test_schema.test_model(a); + """ + ), + ) + + snapshot_2 = make_snapshot(model) + snapshot.categorize_as(SnapshotChangeCategory.BREAKING) + snapshot_2.categorize_as(SnapshotChangeCategory.BREAKING) + evaluator.create([snapshot_2, snapshot], {}, DeployabilityIndex.all_deployable()) + + # Default gateway adapter to create table without gateway + create_args = engine_adapters["default"].create_table.call_args_list + assert len(create_args) == 1 + assert create_args[0][0] == (f"sqlmesh__db.db__model__{snapshot.version}",) + + # Secondary gateway for gateway-specicied model + create_args_2 = engine_adapters["secondary"].create_table.call_args_list + assert len(create_args_2) == 1 + assert create_args_2[0][0] == ( + f"sqlmesh__test_schema.test_schema__test_model__{snapshot_2.version}", + ) + + environment_naming_info = EnvironmentNamingInfo( + name="test_env", gateway_managed_virtual_layer=True + ) + engine_adapters["third"].create_table.assert_not_called() + evaluator.promote([snapshot, snapshot_2], environment_naming_info) + + # Virtual layer will use the model-specified gateway adapter for the second model and default otherwise + view_args_default = engine_adapters["default"].create_view.call_args_list + engine_adapters["third"].create_view.assert_not_called() + view_args_secondary = engine_adapters["secondary"].create_view.call_args_list + + assert len(view_args_default) == 1 + assert view_args_default[0][0][0] == "db__test_env.model" + assert len(view_args_secondary) == 1 + assert view_args_secondary[0][0][0] == "test_schema__test_env.test_model" + + # Demotion will follow with the same pattern + evaluator.demote([snapshot_2, snapshot], environment_naming_info) + engine_adapters["default"].drop_view.assert_called_once_with( + "db__test_env.model", + cascade=False, + ) + engine_adapters["secondary"].drop_view.assert_called_once_with( + "test_schema__test_env.test_model", + cascade=False, + ) From f37ad6c916fd75a1c87528673de106012ff17f30 Mon Sep 17 00:00:00 2001 From: Themis Valtinos <73662635+themisvaltinos@users.noreply.github.com> Date: Tue, 8 Apr 2025 23:22:16 +0300 Subject: [PATCH 03/14] Revert to deleting first from state envs and snapshots --- sqlmesh/core/context.py | 14 ++++------- tests/core/test_context.py | 51 ++++++++++++++++---------------------- 2 files changed, 27 insertions(+), 38 deletions(-) diff --git a/sqlmesh/core/context.py b/sqlmesh/core/context.py index a4862deafa..650395b75a 100644 --- a/sqlmesh/core/context.py +++ b/sqlmesh/core/context.py @@ -2300,14 +2300,15 @@ def _context_diff( def _run_janitor(self, ignore_ttl: bool = False) -> None: # Get expired environments and removes their views and schemas - expired_environments, filter_expr = self._cleanup_environments() + self._cleanup_environments() # Get expired snapshots and corresponding gateways per snapshot when applied expired_snapshots_ids, cleanup_targets, snapshot_gateways = ( self.state_sync.get_expired_snapshots(ignore_ttl=ignore_ttl) ) - # Clean up intervals from the state sync + # Clean up snapshots and intervals from the state sync + self.state_sync.delete_snapshots(expired_snapshots_ids) self.state_sync.cleanup_intervals(cleanup_targets, expired_snapshots_ids) # Remove the expired snapshots tables @@ -2317,13 +2318,10 @@ def _run_janitor(self, ignore_ttl: bool = False) -> None: on_complete=self.console.update_cleanup_progress, ) - # Finally, remove the expired environments and snapshots from the state sync - self.state_sync.delete_environments(expired_environments, filter_expr) - self.state_sync.delete_snapshots(expired_snapshots_ids) self.state_sync.compact_intervals() - def _cleanup_environments(self) -> t.Tuple[t.List[Environment], exp.LTE]: - expired_environments, filter_expr = self.state_sync.get_expired_environments() + def _cleanup_environments(self) -> None: + expired_environments = self.state_sync.delete_expired_environments() environment_snapshot_adapters: t.Dict[str, t.Dict[str, EngineAdapter]] = {} for environment in expired_environments: @@ -2344,8 +2342,6 @@ def _cleanup_environments(self) -> t.Tuple[t.List[Environment], exp.LTE]: environment_snapshot_adapters=environment_snapshot_adapters, ) - return expired_environments, filter_expr - def _try_connection(self, connection_name: str, validator: t.Callable[[], None]) -> None: connection_name = connection_name.capitalize() try: diff --git a/tests/core/test_context.py b/tests/core/test_context.py index 16cc969a4c..8f8b2967d8 100644 --- a/tests/core/test_context.py +++ b/tests/core/test_context.py @@ -43,7 +43,6 @@ from sqlmesh.utils.date import ( make_inclusive_end, now, - now_timestamp, to_date, to_timestamp, yesterday_ds, @@ -780,34 +779,28 @@ def test_janitor(sushi_context, mocker: MockerFixture) -> None: adapter_mock = mocker.MagicMock() adapter_mock.dialect = "duckdb" state_sync_mock = mocker.MagicMock() - now_ts = now_timestamp() - filter_expr = exp.LTE( - this=exp.column("expiration_ts"), - expression=exp.Literal.number(now_ts), - ) - state_sync_mock.get_expired_environments.return_value = ( - [ - Environment( - name="test_environment", - suffix_target=EnvironmentSuffixTarget.TABLE, - snapshots=[x.table_info for x in sushi_context.snapshots.values()], - start_at="2022-01-01", - end_at="2022-01-01", - plan_id="test_plan_id", - previous_plan_id="test_plan_id", - ), - Environment( - name="test_environment", - suffix_target=EnvironmentSuffixTarget.SCHEMA, - snapshots=[x.table_info for x in sushi_context.snapshots.values()], - start_at="2022-01-01", - end_at="2022-01-01", - plan_id="test_plan_id", - previous_plan_id="test_plan_id", - ), - ], - filter_expr, - ) + + state_sync_mock.delete_expired_environments.return_value = [ + Environment( + name="test_environment", + suffix_target=EnvironmentSuffixTarget.TABLE, + snapshots=[x.table_info for x in sushi_context.snapshots.values()], + start_at="2022-01-01", + end_at="2022-01-01", + plan_id="test_plan_id", + previous_plan_id="test_plan_id", + ), + Environment( + name="test_environment", + suffix_target=EnvironmentSuffixTarget.SCHEMA, + snapshots=[x.table_info for x in sushi_context.snapshots.values()], + start_at="2022-01-01", + end_at="2022-01-01", + plan_id="test_plan_id", + previous_plan_id="test_plan_id", + ), + ] + sushi_context._engine_adapters = {sushi_context.config.default_gateway: adapter_mock} sushi_context._state_sync = state_sync_mock state_sync_mock.get_expired_snapshots.return_value = (set({}), [], {}) From d16f515faba660a824e0c5b8f8ea8a0174d605d3 Mon Sep 17 00:00:00 2001 From: Themis Valtinos <73662635+themisvaltinos@users.noreply.github.com> Date: Wed, 9 Apr 2025 14:24:37 +0300 Subject: [PATCH 04/14] Revise to contain gateway info in SnapshotTableCleanupTask --- sqlmesh/core/context.py | 13 ++----- sqlmesh/core/snapshot/definition.py | 1 + sqlmesh/core/snapshot/evaluator.py | 6 +-- sqlmesh/core/state_sync/base.py | 44 ---------------------- sqlmesh/core/state_sync/db/facade.py | 27 -------------- sqlmesh/core/state_sync/db/snapshot.py | 32 +++------------- sqlmesh/schedulers/airflow/state_sync.py | 47 ------------------------ tests/core/test_context.py | 1 - tests/core/test_snapshot_evaluator.py | 10 +++-- 9 files changed, 18 insertions(+), 163 deletions(-) diff --git a/sqlmesh/core/context.py b/sqlmesh/core/context.py index 650395b75a..44cac67e3f 100644 --- a/sqlmesh/core/context.py +++ b/sqlmesh/core/context.py @@ -2299,22 +2299,15 @@ def _context_diff( ) def _run_janitor(self, ignore_ttl: bool = False) -> None: - # Get expired environments and removes their views and schemas + # Clean up expired environments by removing their views and schemas self._cleanup_environments() - # Get expired snapshots and corresponding gateways per snapshot when applied - expired_snapshots_ids, cleanup_targets, snapshot_gateways = ( - self.state_sync.get_expired_snapshots(ignore_ttl=ignore_ttl) - ) - - # Clean up snapshots and intervals from the state sync - self.state_sync.delete_snapshots(expired_snapshots_ids) - self.state_sync.cleanup_intervals(cleanup_targets, expired_snapshots_ids) + # Identify and delete expired snapshots + cleanup_targets = self.state_sync.delete_expired_snapshots(ignore_ttl=ignore_ttl) # Remove the expired snapshots tables self.snapshot_evaluator.cleanup( target_snapshots=cleanup_targets, - snapshot_gateways=snapshot_gateways, on_complete=self.console.update_cleanup_progress, ) diff --git a/sqlmesh/core/snapshot/definition.py b/sqlmesh/core/snapshot/definition.py index 3acb1527d8..e2e9c26d15 100644 --- a/sqlmesh/core/snapshot/definition.py +++ b/sqlmesh/core/snapshot/definition.py @@ -1338,6 +1338,7 @@ def __getstate__(self) -> t.Dict[t.Any, t.Any]: class SnapshotTableCleanupTask(PydanticModel): snapshot: SnapshotTableInfo dev_table_only: bool + gateway: t.Optional[str] = None SnapshotIdLike = t.Union[SnapshotId, SnapshotTableInfo, Snapshot] diff --git a/sqlmesh/core/snapshot/evaluator.py b/sqlmesh/core/snapshot/evaluator.py index d770584d74..2e8c6d5732 100644 --- a/sqlmesh/core/snapshot/evaluator.py +++ b/sqlmesh/core/snapshot/evaluator.py @@ -426,7 +426,6 @@ def migrate( def cleanup( self, target_snapshots: t.Iterable[SnapshotTableCleanupTask], - snapshot_gateways: t.Optional[t.Dict[str, str]] = None, on_complete: t.Optional[t.Callable[[str], None]] = None, ) -> None: """Cleans up the given snapshots by removing its table @@ -438,6 +437,7 @@ def cleanup( snapshots_to_dev_table_only = { t.snapshot.snapshot_id: t.dev_table_only for t in target_snapshots } + snapshot_gateways = {t.snapshot.snapshot_id: t.gateway for t in target_snapshots} with self.concurrent_context(): concurrent_apply_to_snapshots( @@ -445,9 +445,7 @@ def cleanup( lambda s: self._cleanup_snapshot( s, snapshots_to_dev_table_only[s.snapshot_id], - self.get_adapter( - snapshot_gateways.get(s.snapshot_id.name) if snapshot_gateways else None - ), + self._get_adapter(snapshot_gateways[s.snapshot_id]), on_complete, ), self.ddl_concurrent_tasks, diff --git a/sqlmesh/core/state_sync/base.py b/sqlmesh/core/state_sync/base.py index 44471a41a9..771dd94172 100644 --- a/sqlmesh/core/state_sync/base.py +++ b/sqlmesh/core/state_sync/base.py @@ -7,7 +7,6 @@ import typing as t from sqlglot import __version__ as SQLGLOT_VERSION -from sqlglot import exp from sqlmesh import migrations from sqlmesh.core.environment import Environment, EnvironmentNamingInfo, EnvironmentStatements @@ -320,23 +319,6 @@ def delete_expired_snapshots( The list of table cleanup tasks. """ - @abc.abstractmethod - def get_expired_snapshots( - self, ignore_ttl: bool = False - ) -> t.Tuple[t.Set[SnapshotId], t.List[SnapshotTableCleanupTask], t.Dict[str, str]]: - """Gets expired snapshots. - - Expired snapshots are snapshots that have exceeded their time-to-live - and are no longer in use within an environment. - - Args: - ignore_ttl: Ignore the TTL on the snapshot when considering it expired. This has the effect of deleting - all snapshots that are not referenced in any environment - - Returns: - A tuple of expired snapshot IDs, cleanup targets and gateway per snapshot dictionary. - """ - @abc.abstractmethod def invalidate_environment(self, name: str) -> None: """Invalidates the target environment by setting its expiration timestamp to now. @@ -345,14 +327,6 @@ def invalidate_environment(self, name: str) -> None: name: The name of the environment to invalidate. """ - @abc.abstractmethod - def cleanup_intervals( - self, - cleanup_targets: t.List[SnapshotTableCleanupTask], - expired_snapshot_ids: t.Set[SnapshotId], - ) -> None: - """Cleans up intervals.""" - @abc.abstractmethod def remove_intervals( self, @@ -421,24 +395,6 @@ def delete_expired_environments(self) -> t.List[Environment]: The list of removed environments. """ - @abc.abstractmethod - def get_expired_environments(self) -> t.Tuple[t.List[Environment], exp.LTE]: - """Returns the expired environments. - - Expired environments are environments that have exceeded their time-to-live value. - - Returns: - The list of environments to remove, the filter to remove environments. - """ - - @abc.abstractmethod - def delete_environments(self, environments: t.List[Environment], filter_expr: exp.LTE) -> None: - """Removes environments. - - Returns: - The list of removed environments. - """ - @abc.abstractmethod def unpause_snapshots( self, snapshots: t.Collection[SnapshotInfoLike], unpaused_dt: TimeLike diff --git a/sqlmesh/core/state_sync/db/facade.py b/sqlmesh/core/state_sync/db/facade.py index 0fe1a74749..0ed17dbdfd 100644 --- a/sqlmesh/core/state_sync/db/facade.py +++ b/sqlmesh/core/state_sync/db/facade.py @@ -288,33 +288,6 @@ def delete_expired_snapshots( def delete_expired_environments(self) -> t.List[Environment]: return self.environment_state.delete_expired_environments() - @transactional() - def delete_environments(self, environments: t.List[Environment], filter_expr: exp.LTE) -> None: - self.environment_state.delete_environments( - environments=environments, filter_expr=filter_expr - ) - - @transactional() - def cleanup_intervals( - self, - cleanup_targets: t.List[SnapshotTableCleanupTask], - expired_snapshot_ids: t.Set[SnapshotId], - ) -> None: - self.interval_state.cleanup_intervals(cleanup_targets, expired_snapshot_ids) - - def get_expired_snapshots( - self, ignore_ttl: bool = False - ) -> t.Tuple[t.Set[SnapshotId], t.List[SnapshotTableCleanupTask], t.Dict[str, str]]: - expired_snapshot_ids, cleanup_targets, snapshot_gateways = ( - self.snapshot_state.get_expired_snapshots( - self.environment_state.get_environments(), ignore_ttl=ignore_ttl - ) - ) - return expired_snapshot_ids, cleanup_targets, snapshot_gateways - - def get_expired_environments(self) -> t.Tuple[t.List[Environment], exp.LTE]: - return self.environment_state.get_expired_environments() - def delete_snapshots(self, snapshot_ids: t.Iterable[SnapshotIdLike]) -> None: self.snapshot_state.delete_snapshots(snapshot_ids) diff --git a/sqlmesh/core/state_sync/db/snapshot.py b/sqlmesh/core/state_sync/db/snapshot.py index 8c5c4ac33a..f25b02a89d 100644 --- a/sqlmesh/core/state_sync/db/snapshot.py +++ b/sqlmesh/core/state_sync/db/snapshot.py @@ -205,26 +205,6 @@ def delete_expired_snapshots( A tuple of expired snapshot IDs and cleanup targets. """ - expired_snapshot_ids, cleanup_targets, _ = self.get_expired_snapshots( - environments=environments, ignore_ttl=ignore_ttl - ) - - if expired_snapshot_ids: - self.delete_snapshots(expired_snapshot_ids) - - return expired_snapshot_ids, cleanup_targets - - def get_expired_snapshots( - self, environments: t.Iterable[Environment], ignore_ttl: bool = False - ) -> t.Tuple[t.Set[SnapshotId], t.List[SnapshotTableCleanupTask], t.Dict[str, str]]: - """Gets the expired snapshots. - - Args: - ignore_ttl: Whether to ignore the TTL of the snapshots. - - Returns: - A tuple of expired snapshot IDs, cleanup targets and gateway per snapshot dictionary. - """ current_ts = now_timestamp(minute_floor=False) expired_query = exp.select("name", "identifier", "version").from_(self.snapshots_table) @@ -241,7 +221,7 @@ def get_expired_snapshots( for name, identifier, version in fetchall(self.engine_adapter, expired_query) } if not expired_candidates: - return set(), [], {} + return set(), [] promoted_snapshot_ids = { snapshot.snapshot_id @@ -261,7 +241,6 @@ def _is_snapshot_used(snapshot: SharedVersionSnapshot) -> bool: ) cleanup_targets = [] expired_snapshot_ids = set() - snapshot_gateways: t.Dict[str, str] = {} for versions_batch in version_batches: snapshots = self._get_snapshots_with_same_version(versions_batch) @@ -275,9 +254,6 @@ def _is_snapshot_used(snapshot: SharedVersionSnapshot) -> bool: expired_snapshot_ids.update([s.snapshot_id for s in expired_snapshots]) for snapshot in expired_snapshots: - if (node := snapshot.raw_snapshot.get("node")) and (gateway := node.get("gateway")): - snapshot_gateways[snapshot.snapshot_id.name] = gateway - shared_version_snapshots = snapshots_by_version[(snapshot.name, snapshot.version)] shared_version_snapshots.discard(snapshot.snapshot_id) @@ -291,10 +267,14 @@ def _is_snapshot_used(snapshot: SharedVersionSnapshot) -> bool: SnapshotTableCleanupTask( snapshot=snapshot.full_snapshot.table_info, dev_table_only=bool(shared_version_snapshots), + gateway=snapshot.raw_snapshot.get("node", {}).get("gateway", None), ) ) - return expired_snapshot_ids, cleanup_targets, snapshot_gateways + if expired_snapshot_ids: + self.delete_snapshots(expired_snapshot_ids) + + return expired_snapshot_ids, cleanup_targets def delete_snapshots(self, snapshot_ids: t.Iterable[SnapshotIdLike]) -> None: """Deletes snapshots. diff --git a/sqlmesh/schedulers/airflow/state_sync.py b/sqlmesh/schedulers/airflow/state_sync.py index 7df3d855b1..fecf15199a 100644 --- a/sqlmesh/schedulers/airflow/state_sync.py +++ b/sqlmesh/schedulers/airflow/state_sync.py @@ -3,7 +3,6 @@ import logging import typing as t -from sqlglot import exp from sqlmesh.core.console import Console from sqlmesh.core.environment import Environment, EnvironmentStatements from sqlmesh.core.snapshot import ( @@ -294,41 +293,6 @@ def delete_expired_environments(self) -> t.List[Environment]: "Deleting expired environments is not supported by the Airflow state sync." ) - def get_expired_environments(self) -> t.Tuple[t.List[Environment], exp.LTE]: - """Returns the expired environments. - - Expired environments are environments that have exceeded their time-to-live value. - - """ - raise NotImplementedError( - "Getting expired environments is not supported by the Airflow state sync." - ) - - def delete_environments(self, environments: t.List[Environment], filter_expr: exp.LTE) -> None: - """Removes environments. - - Returns: - The list of removed environments. - """ - raise NotImplementedError( - "Deleting environments is not supported by the Airflow state sync." - ) - - def get_expired_snapshots( - self, ignore_ttl: bool = False - ) -> t.Tuple[t.Set[SnapshotId], t.List[SnapshotTableCleanupTask], t.Dict[str, str]]: - """Gets expired snapshots. - - Expired snapshots are snapshots that have exceeded their time-to-live - and are no longer in use within an environment. - - Returns: - A tuple of expired snapshot IDs, cleanup targets and gateway per snapshot dictionary. - """ - raise NotImplementedError( - "Getting expired snapshots is not supported by the Airflow state sync." - ) - def unpause_snapshots( self, snapshots: t.Collection[SnapshotInfoLike], unpaused_dt: TimeLike ) -> None: @@ -354,17 +318,6 @@ def compact_intervals(self) -> None: "Compacting intervals is not supported by the Airflow state sync." ) - def cleanup_intervals( - self, - cleanup_targets: t.List[SnapshotTableCleanupTask], - expired_snapshot_ids: t.Set[SnapshotId], - ) -> None: - """Cleans up intervals.""" - - raise NotImplementedError( - "Cleaning up intervals is not supported by the Airflow state sync." - ) - def update_auto_restatements( self, next_auto_restatement_ts: t.Dict[SnapshotNameVersion, t.Optional[int]] ) -> None: diff --git a/tests/core/test_context.py b/tests/core/test_context.py index 8f8b2967d8..1c512b83d9 100644 --- a/tests/core/test_context.py +++ b/tests/core/test_context.py @@ -803,7 +803,6 @@ def test_janitor(sushi_context, mocker: MockerFixture) -> None: sushi_context._engine_adapters = {sushi_context.config.default_gateway: adapter_mock} sushi_context._state_sync = state_sync_mock - state_sync_mock.get_expired_snapshots.return_value = (set({}), [], {}) sushi_context._run_janitor() # Assert that the schemas are dropped just twice for the schema based environment # Make sure that external model schemas/tables are not dropped diff --git a/tests/core/test_snapshot_evaluator.py b/tests/core/test_snapshot_evaluator.py index 8eb8d8087f..31f946d7ea 100644 --- a/tests/core/test_snapshot_evaluator.py +++ b/tests/core/test_snapshot_evaluator.py @@ -3967,13 +3967,15 @@ def test_multiple_engine_cleanup(snapshot: Snapshot, adapters, make_snapshot): f"sqlmesh__test_schema.test_schema__test_model__{snapshot_2.version}", ) - snapshot_gateways = {snapshot.name: "default", snapshot_2.name: "secondary"} evaluator.cleanup( [ - SnapshotTableCleanupTask(snapshot=snapshot.table_info, dev_table_only=True), - SnapshotTableCleanupTask(snapshot=snapshot_2.table_info, dev_table_only=True), + SnapshotTableCleanupTask( + snapshot=snapshot.table_info, dev_table_only=True, gateway="default" + ), + SnapshotTableCleanupTask( + snapshot=snapshot_2.table_info, dev_table_only=True, gateway="secondary" + ), ], - snapshot_gateways, ) # The clean up will happen using the specific gateway the model was created with From 976f99f858846b9232c07374293391836ba4f65d Mon Sep 17 00:00:00 2001 From: Themis Valtinos <73662635+themisvaltinos@users.noreply.github.com> Date: Wed, 9 Apr 2025 15:57:36 +0300 Subject: [PATCH 05/14] Remove unused refactored methods --- sqlmesh/core/state_sync/db/environment.py | 56 +++++++---------------- 1 file changed, 17 insertions(+), 39 deletions(-) diff --git a/sqlmesh/core/state_sync/db/environment.py b/sqlmesh/core/state_sync/db/environment.py index 3a4cbfa8c6..470494902a 100644 --- a/sqlmesh/core/state_sync/db/environment.py +++ b/sqlmesh/core/state_sync/db/environment.py @@ -168,44 +168,6 @@ def delete_expired_environments(self) -> t.List[Environment]: Returns: A list of deleted environments. """ - environments, filter_expr = self.get_expired_environments() - self.delete_environments(environments=environments, filter_expr=filter_expr) - return environments - - def delete_environments(self, environments: t.List[Environment], filter_expr: exp.LTE) -> None: - """Deletes environments and corresponding environment statements. - - Returns: - The list of deleted environments. - """ - if not environments: - return - - # Delete the expired environments - self.engine_adapter.delete_from( - self.environments_table, - where=filter_expr, - ) - - # Delete the expired environments' corresponding environment statements - expired_environments = [ - exp.EQ(this=exp.column("environment_name"), expression=exp.Literal.string(env.name)) - for env in environments - ] - if expired_environments: - self.engine_adapter.delete_from( - self.environment_statements_table, - where=exp.or_(*expired_environments), - ) - - def get_expired_environments(self) -> t.Tuple[t.List[Environment], exp.LTE]: - """Fetches expired environments along with their filter expression and snapshot gateways. - - Returns: - A tuple containing: - - A list of expired environments. - - The filter expression used to identify expired environments. - """ now_ts = now_timestamp() filter_expr = exp.LTE( @@ -223,7 +185,23 @@ def get_expired_environments(self) -> t.Tuple[t.List[Environment], exp.LTE]: environments = [self._environment_from_row(r) for r in rows] - return environments, filter_expr + # Delete the expired environments + self.engine_adapter.delete_from( + self.environments_table, + where=filter_expr, + ) + + # Delete the expired environments' corresponding environment statements + if expired_environments := [ + exp.EQ(this=exp.column("environment_name"), expression=exp.Literal.string(env.name)) + for env in environments + ]: + self.engine_adapter.delete_from( + self.environment_statements_table, + where=exp.or_(*expired_environments), + ) + + return environments def get_environments(self) -> t.List[Environment]: """Fetches all environments. From 44f3a1c17df7effafc697dc9945c5bd3b4a65174 Mon Sep 17 00:00:00 2001 From: Themis Valtinos <73662635+themisvaltinos@users.noreply.github.com> Date: Wed, 9 Apr 2025 23:37:57 +0300 Subject: [PATCH 06/14] Refactors; revise to keep model gateway in SnapshotTableInfo --- sqlmesh/core/context.py | 14 +-------- sqlmesh/core/context_diff.py | 6 +--- sqlmesh/core/environment.py | 21 ++++++-------- sqlmesh/core/plan/builder.py | 2 +- sqlmesh/core/plan/evaluator.py | 8 +---- sqlmesh/core/snapshot/definition.py | 3 +- sqlmesh/core/snapshot/evaluator.py | 12 ++++---- sqlmesh/core/state_sync/common.py | 29 ++++++------------- sqlmesh/core/state_sync/db/environment.py | 7 ++--- sqlmesh/core/state_sync/db/facade.py | 3 +- sqlmesh/core/state_sync/db/snapshot.py | 3 -- .../migrations/v0056_restore_table_indexes.py | 2 +- ... => v0078_add_gateway_managed_property.py} | 4 +-- tests/core/test_integration.py | 2 +- tests/core/test_snapshot_evaluator.py | 18 ++++-------- tests/schedulers/airflow/test_client.py | 2 +- 16 files changed, 42 insertions(+), 94 deletions(-) rename sqlmesh/migrations/{v0078_add_gateway_managed_virtual_layer.py => v0078_add_gateway_managed_property.py} (86%) diff --git a/sqlmesh/core/context.py b/sqlmesh/core/context.py index 44cac67e3f..88ffdd2138 100644 --- a/sqlmesh/core/context.py +++ b/sqlmesh/core/context.py @@ -2316,23 +2316,11 @@ def _run_janitor(self, ignore_ttl: bool = False) -> None: def _cleanup_environments(self) -> None: expired_environments = self.state_sync.delete_expired_environments() - environment_snapshot_adapters: t.Dict[str, t.Dict[str, EngineAdapter]] = {} - for environment in expired_environments: - snapshot_adapters: t.Dict[str, EngineAdapter] = {} - if environment.gateway_managed_virtual_layer: - snapshots = self.state_sync.get_snapshots(environment.snapshots) - for snapshot_id, snapshot in snapshots.items(): - if snapshot.is_model and not snapshot.is_symbolic: - snapshot_adapters[snapshot_id.name] = self._get_engine_adapter( - snapshot.model_gateway - ) - environment_snapshot_adapters[environment.name] = snapshot_adapters - cleanup_expired_views( adapter=self.engine_adapter, environments=expired_environments, console=self.console, - environment_snapshot_adapters=environment_snapshot_adapters, + engine_adapters=self.engine_adapters, ) def _try_connection(self, connection_name: str, validator: t.Callable[[], None]) -> None: diff --git a/sqlmesh/core/context_diff.py b/sqlmesh/core/context_diff.py index 82d2a52cdc..eb05a61155 100644 --- a/sqlmesh/core/context_diff.py +++ b/sqlmesh/core/context_diff.py @@ -121,11 +121,7 @@ def create( env = state_reader.get_environment(environment) create_from_env_exists = False - if ( - env is None - or env.expired - or env.gateway_managed_virtual_layer != gateway_managed_virtual_layer - ): + if env is None or env.expired or env.gateway_managed != gateway_managed_virtual_layer: env = state_reader.get_environment(create_from.lower()) if not env and create_from != c.PROD: diff --git a/sqlmesh/core/environment.py b/sqlmesh/core/environment.py index 95744be343..be3f557999 100644 --- a/sqlmesh/core/environment.py +++ b/sqlmesh/core/environment.py @@ -16,7 +16,7 @@ from sqlmesh.utils.date import TimeLike, now_timestamp from sqlmesh.utils.jinja import JinjaMacroRegistry from sqlmesh.utils.metaprogramming import Executable -from sqlmesh.utils.pydantic import PydanticModel, field_validator +from sqlmesh.utils.pydantic import PydanticModel, field_validator, ValidationInfo T = t.TypeVar("T", bound="EnvironmentNamingInfo") PydanticType = t.TypeVar("PydanticType", bound="PydanticModel") @@ -32,7 +32,7 @@ class EnvironmentNamingInfo(PydanticModel): catalog_name_override: The name of the catalog to use for this environment if an override was provided normalize_name: Indicates whether the environment's name will be normalized. For example, if it's `dev`, then it will become `DEV` when targeting Snowflake. - gateway_managed_virtual_layer: Determines whether the virtual layer's views are created by the model-specific + gateway_managed: Determines whether the virtual layer's views are created by the model-specific gateways, otherwise the default gateway is used. Default: False. """ @@ -40,22 +40,19 @@ class EnvironmentNamingInfo(PydanticModel): suffix_target: EnvironmentSuffixTarget = Field(default=EnvironmentSuffixTarget.SCHEMA) catalog_name_override: t.Optional[str] = None normalize_name: bool = True - gateway_managed_virtual_layer: bool = False + gateway_managed: bool = False @field_validator("name", mode="before") @classmethod def _sanitize_name(cls, v: str) -> str: return word_characters_only(v).lower() - @field_validator("normalize_name", mode="before") + @field_validator("normalize_name", "gateway_managed", mode="before") @classmethod - def _validate_normalize_name(cls, v: t.Any) -> bool: - return True if v is None else bool(v) - - @field_validator("gateway_managed_virtual_layer", mode="before") - @classmethod - def _validate_gateway_managed_virtual_layer(cls, v: t.Any) -> bool: - return False if v is None else bool(v) + def _validate_boolean_field(cls, v: t.Any, info: ValidationInfo) -> bool: + if v is None: + return info.field_name == "normalize_name" + return bool(v) @t.overload @classmethod @@ -202,7 +199,7 @@ def naming_info(self) -> EnvironmentNamingInfo: suffix_target=self.suffix_target, catalog_name_override=self.catalog_name_override, normalize_name=self.normalize_name, - gateway_managed_virtual_layer=self.gateway_managed_virtual_layer, + gateway_managed=self.gateway_managed, ) @property diff --git a/sqlmesh/core/plan/builder.py b/sqlmesh/core/plan/builder.py index bf3e8fd543..5e371d7b28 100644 --- a/sqlmesh/core/plan/builder.py +++ b/sqlmesh/core/plan/builder.py @@ -151,7 +151,7 @@ def __init__( name=self._context_diff.environment, suffix_target=environment_suffix_target, normalize_name=self._context_diff.normalize_environment_name, - gateway_managed_virtual_layer=self._context_diff.gateway_managed_virtual_layer, + gateway_managed=self._context_diff.gateway_managed_virtual_layer, ) self._latest_plan: t.Optional[Plan] = None diff --git a/sqlmesh/core/plan/evaluator.py b/sqlmesh/core/plan/evaluator.py index 685cf3e396..4e4d7e5bfa 100644 --- a/sqlmesh/core/plan/evaluator.py +++ b/sqlmesh/core/plan/evaluator.py @@ -423,14 +423,8 @@ def _demote_snapshots( environment_naming_info: EnvironmentNamingInfo, on_complete: t.Optional[t.Callable[[SnapshotInfoLike], None]] = None, ) -> None: - # In a multi virtual layer setup we need the gateway info from the snapshots for demotion - snapshots_to_demote: t.List[Snapshot] = [] - if environment_naming_info.gateway_managed_virtual_layer: - removed_snapshots = self.state_sync.get_snapshots(target_snapshots) - snapshots_to_demote = [removed_snapshots[s.snapshot_id] for s in target_snapshots] - self.snapshot_evaluator.demote( - snapshots_to_demote or target_snapshots, + target_snapshots, environment_naming_info, on_complete=on_complete, ) diff --git a/sqlmesh/core/snapshot/definition.py b/sqlmesh/core/snapshot/definition.py index e2e9c26d15..386308cb95 100644 --- a/sqlmesh/core/snapshot/definition.py +++ b/sqlmesh/core/snapshot/definition.py @@ -478,6 +478,7 @@ class SnapshotTableInfo(PydanticModel, SnapshotInfoMixin, frozen=True): base_table_name_override: t.Optional[str] = None custom_materialization: t.Optional[str] = None dev_table_suffix: str + model_gateway: t.Optional[str] = None def __lt__(self, other: SnapshotTableInfo) -> bool: return self.name < other.name @@ -1179,6 +1180,7 @@ def table_info(self) -> SnapshotTableInfo: node_type=self.node_type, custom_materialization=custom_materialization, dev_table_suffix=self.dev_table_suffix, + model_gateway=self.model_gateway, ) @property @@ -1338,7 +1340,6 @@ def __getstate__(self) -> t.Dict[t.Any, t.Any]: class SnapshotTableCleanupTask(PydanticModel): snapshot: SnapshotTableInfo dev_table_only: bool - gateway: t.Optional[str] = None SnapshotIdLike = t.Union[SnapshotId, SnapshotTableInfo, Snapshot] diff --git a/sqlmesh/core/snapshot/evaluator.py b/sqlmesh/core/snapshot/evaluator.py index 2e8c6d5732..92a0eaa5a9 100644 --- a/sqlmesh/core/snapshot/evaluator.py +++ b/sqlmesh/core/snapshot/evaluator.py @@ -234,11 +234,11 @@ def promote( table = snapshot.qualified_view_name.table_for_environment( environment_naming_info, dialect=self._get_adapter(snapshot.model_gateway).dialect - if environment_naming_info.gateway_managed_virtual_layer + if environment_naming_info.gateway_managed else self.adapter.dialect, ) tables.append(table) - if environment_naming_info.gateway_managed_virtual_layer: + if environment_naming_info.gateway_managed: table_schema = d.schema_(table.db, catalog=table.catalog) gateway_by_schema[table_schema] = snapshot.model_gateway or "" self._create_schemas(tables=tables, gateways=gateway_by_schema) @@ -437,7 +437,6 @@ def cleanup( snapshots_to_dev_table_only = { t.snapshot.snapshot_id: t.dev_table_only for t in target_snapshots } - snapshot_gateways = {t.snapshot.snapshot_id: t.gateway for t in target_snapshots} with self.concurrent_context(): concurrent_apply_to_snapshots( @@ -445,7 +444,7 @@ def cleanup( lambda s: self._cleanup_snapshot( s, snapshots_to_dev_table_only[s.snapshot_id], - self._get_adapter(snapshot_gateways[s.snapshot_id]), + self._get_adapter(s.model_gateway), on_complete, ), self.ddl_concurrent_tasks, @@ -931,7 +930,7 @@ def _promote_snapshot( if snapshot.is_model: adapter = ( self._get_adapter(snapshot.model_gateway) - if environment_naming_info.gateway_managed_virtual_layer + if environment_naming_info.gateway_managed else self.adapter ) table_name = snapshot.table_name(deployability_index.is_representative(snapshot)) @@ -968,8 +967,7 @@ def _demote_snapshot( ) -> None: adapter = ( self._get_adapter(snapshot.model_gateway) - if environment_naming_info.gateway_managed_virtual_layer - and isinstance(snapshot, Snapshot) + if environment_naming_info.gateway_managed else self.adapter ) view_name = snapshot.qualified_view_name.for_environment( diff --git a/sqlmesh/core/state_sync/common.py b/sqlmesh/core/state_sync/common.py index ae7d36cc09..6fc4a2e405 100644 --- a/sqlmesh/core/state_sync/common.py +++ b/sqlmesh/core/state_sync/common.py @@ -23,8 +23,8 @@ def cleanup_expired_views( adapter: EngineAdapter, environments: t.List[Environment], - environment_snapshot_adapters: t.Optional[t.Dict[str, t.Dict[str, EngineAdapter]]] = None, console: t.Optional[Console] = None, + engine_adapters: t.Optional[t.Dict[str, EngineAdapter]] = None, ) -> None: expired_schema_environments = [ environment for environment in environments if environment.suffix_target.is_schema @@ -33,19 +33,16 @@ def cleanup_expired_views( environment for environment in environments if environment.suffix_target.is_table ] + # We have to use the corresponding adapter if the virtual layer is gateway managed + def get_adapter(gateway_managed: bool, gateway: t.Optional[str] = None) -> EngineAdapter: + if gateway_managed and gateway: + return (engine_adapters or {}).get(gateway, adapter) + return adapter + # Drop the schemas for the expired environments - # Note: We have to use the corresponding adapter if it is a gateway managed virtual layer for engine_adapter, expired_catalog, expired_schema in { ( - ( - engine_adapter := ( - (environment_dict.get(snapshot.name) or adapter) - if environment.gateway_managed_virtual_layer - and environment_snapshot_adapters - and (environment_dict := environment_snapshot_adapters.get(environment.name)) - else adapter - ) - ), + (engine_adapter := get_adapter(environment.gateway_managed, snapshot.model_gateway)), snapshot.qualified_view_name.catalog_for_environment( environment.naming_info, dialect=engine_adapter.dialect ), @@ -74,15 +71,7 @@ def cleanup_expired_views( # Drop the views for the expired environments for engine_adapter, expired_view in { ( - ( - engine_adapter := ( - (environment_dict.get(snapshot.name) or adapter) - if environment.gateway_managed_virtual_layer - and environment_snapshot_adapters - and (environment_dict := environment_snapshot_adapters.get(environment.name)) - else adapter - ) - ), + (engine_adapter := get_adapter(environment.gateway_managed, snapshot.model_gateway)), snapshot.qualified_view_name.for_environment( environment.naming_info, dialect=engine_adapter.dialect ), diff --git a/sqlmesh/core/state_sync/db/environment.py b/sqlmesh/core/state_sync/db/environment.py index 470494902a..a63caf6259 100644 --- a/sqlmesh/core/state_sync/db/environment.py +++ b/sqlmesh/core/state_sync/db/environment.py @@ -48,7 +48,7 @@ def __init__( "catalog_name_override": exp.DataType.build("text"), "previous_finalized_snapshots": exp.DataType.build(blob_type), "normalize_name": exp.DataType.build("boolean"), - "gateway_managed_virtual_layer": exp.DataType.build("boolean"), + "gateway_managed": exp.DataType.build("boolean"), "requirements": exp.DataType.build(blob_type), } @@ -168,7 +168,6 @@ def delete_expired_environments(self) -> t.List[Environment]: Returns: A list of deleted environments. """ - now_ts = now_timestamp() filter_expr = exp.LTE( this=exp.column("expiration_ts"), @@ -182,10 +181,8 @@ def delete_expired_environments(self) -> t.List[Environment]: lock_for_update=True, ), ) - environments = [self._environment_from_row(r) for r in rows] - # Delete the expired environments self.engine_adapter.delete_from( self.environments_table, where=filter_expr, @@ -331,7 +328,7 @@ def _environment_to_df(environment: Environment) -> pd.DataFrame: else None ), "normalize_name": environment.normalize_name, - "gateway_managed_virtual_layer": environment.gateway_managed_virtual_layer, + "gateway_managed": environment.gateway_managed, "requirements": json.dumps(environment.requirements), } ] diff --git a/sqlmesh/core/state_sync/db/facade.py b/sqlmesh/core/state_sync/db/facade.py index 0ed17dbdfd..e5f7493e2a 100644 --- a/sqlmesh/core/state_sync/db/facade.py +++ b/sqlmesh/core/state_sync/db/facade.py @@ -201,8 +201,7 @@ def promote( } if ( not existing_environment.expired - and existing_environment.gateway_managed_virtual_layer - == environment.gateway_managed_virtual_layer + and existing_environment.gateway_managed == environment.gateway_managed ): if environment.previous_plan_id != existing_environment.plan_id: raise ConflictingPlanError( diff --git a/sqlmesh/core/state_sync/db/snapshot.py b/sqlmesh/core/state_sync/db/snapshot.py index f25b02a89d..e46f7d0151 100644 --- a/sqlmesh/core/state_sync/db/snapshot.py +++ b/sqlmesh/core/state_sync/db/snapshot.py @@ -204,7 +204,6 @@ def delete_expired_snapshots( Returns: A tuple of expired snapshot IDs and cleanup targets. """ - current_ts = now_timestamp(minute_floor=False) expired_query = exp.select("name", "identifier", "version").from_(self.snapshots_table) @@ -267,13 +266,11 @@ def _is_snapshot_used(snapshot: SharedVersionSnapshot) -> bool: SnapshotTableCleanupTask( snapshot=snapshot.full_snapshot.table_info, dev_table_only=bool(shared_version_snapshots), - gateway=snapshot.raw_snapshot.get("node", {}).get("gateway", None), ) ) if expired_snapshot_ids: self.delete_snapshots(expired_snapshot_ids) - return expired_snapshot_ids, cleanup_targets def delete_snapshots(self, snapshot_ids: t.Iterable[SnapshotIdLike]) -> None: diff --git a/sqlmesh/migrations/v0056_restore_table_indexes.py b/sqlmesh/migrations/v0056_restore_table_indexes.py index 9d290e6b3c..d6fab1669b 100644 --- a/sqlmesh/migrations/v0056_restore_table_indexes.py +++ b/sqlmesh/migrations/v0056_restore_table_indexes.py @@ -1,4 +1,4 @@ -"""Reads indexes and primary keys in case tables were restored from a backup.""" +"""Readds indexes and primary keys in case tables were restored from a backup.""" from sqlglot import exp from sqlmesh.utils import random_id diff --git a/sqlmesh/migrations/v0078_add_gateway_managed_virtual_layer.py b/sqlmesh/migrations/v0078_add_gateway_managed_property.py similarity index 86% rename from sqlmesh/migrations/v0078_add_gateway_managed_virtual_layer.py rename to sqlmesh/migrations/v0078_add_gateway_managed_property.py index bb43d27a7e..15031372ff 100644 --- a/sqlmesh/migrations/v0078_add_gateway_managed_virtual_layer.py +++ b/sqlmesh/migrations/v0078_add_gateway_managed_property.py @@ -14,7 +14,7 @@ def migrate(state_sync, **kwargs): # type: ignore kind="TABLE", actions=[ exp.ColumnDef( - this=exp.to_column("gateway_managed_virtual_layer"), + this=exp.to_column("gateway_managed"), kind=exp.DataType.build("boolean"), ) ], @@ -23,6 +23,6 @@ def migrate(state_sync, **kwargs): # type: ignore state_sync.engine_adapter.update_table( environments_table, - {"gateway_managed_virtual_layer": False}, + {"gateway_managed": False}, where=exp.true(), ) diff --git a/tests/core/test_integration.py b/tests/core/test_integration.py index 5af1de83c4..e4ab90cbbc 100644 --- a/tests/core/test_integration.py +++ b/tests/core/test_integration.py @@ -4550,7 +4550,7 @@ def test_multi_virtual_layer(mocker): state_environments = context.state_reader.get_environments() state_snapshots = context.state_reader.get_snapshots(context.snapshots.values()) - assert state_environments[0].gateway_managed_virtual_layer + assert state_environments[0].gateway_managed assert len(state_snapshots) == len(state_environments[0].snapshots) assert [snapshot.name for snapshot in plan.directly_modified] == [ diff --git a/tests/core/test_snapshot_evaluator.py b/tests/core/test_snapshot_evaluator.py index 31f946d7ea..c5424d90dc 100644 --- a/tests/core/test_snapshot_evaluator.py +++ b/tests/core/test_snapshot_evaluator.py @@ -3969,12 +3969,8 @@ def test_multiple_engine_cleanup(snapshot: Snapshot, adapters, make_snapshot): evaluator.cleanup( [ - SnapshotTableCleanupTask( - snapshot=snapshot.table_info, dev_table_only=True, gateway="default" - ), - SnapshotTableCleanupTask( - snapshot=snapshot_2.table_info, dev_table_only=True, gateway="secondary" - ), + SnapshotTableCleanupTask(snapshot=snapshot.table_info, dev_table_only=True), + SnapshotTableCleanupTask(snapshot=snapshot_2.table_info, dev_table_only=True), ], ) @@ -4063,10 +4059,8 @@ def model_with_statements(context, **kwargs): cascade=False, ) - environment_naming_info_gw = EnvironmentNamingInfo( - name="test_env", gateway_managed_virtual_layer=True - ) - # Validate that promoting with gateway_managed_virtual_layer leads to this gateway being used for virtual layer + environment_naming_info_gw = EnvironmentNamingInfo(name="test_env", gateway_managed=True) + # Validate that promoting with gateway_managed leads to this gateway being used for virtual layer evaluator.promote([snapshot], environment_naming_info_gw) view_args = engine_adapters["secondary"].create_view.call_args_list assert len(view_args) == 1 @@ -4116,9 +4110,7 @@ def test_multiple_engine_virtual_layer(snapshot: Snapshot, adapters, make_snapsh f"sqlmesh__test_schema.test_schema__test_model__{snapshot_2.version}", ) - environment_naming_info = EnvironmentNamingInfo( - name="test_env", gateway_managed_virtual_layer=True - ) + environment_naming_info = EnvironmentNamingInfo(name="test_env", gateway_managed=True) engine_adapters["third"].create_table.assert_not_called() evaluator.promote([snapshot, snapshot_2], environment_naming_info) diff --git a/tests/schedulers/airflow/test_client.py b/tests/schedulers/airflow/test_client.py index 92a92dc964..94ed8cb5e8 100644 --- a/tests/schedulers/airflow/test_client.py +++ b/tests/schedulers/airflow/test_client.py @@ -170,7 +170,7 @@ def test_apply_plan(mocker: MockerFixture, snapshot: Snapshot): ], "start_at": "2022-01-01", "end_at": "2022-01-01", - "gateway_managed_virtual_layer": False, + "gateway_managed": False, "plan_id": "test_plan_id", "previous_plan_id": "previous_plan_id", "promoted_snapshot_ids": [ From 52d1fb7b605a5f89b3ca3c2396880eced2c6713a Mon Sep 17 00:00:00 2001 From: Themis Valtinos <73662635+themisvaltinos@users.noreply.github.com> Date: Thu, 10 Apr 2025 11:22:46 +0300 Subject: [PATCH 07/14] Rename migration script; add comment; rebase --- sqlmesh/core/model/definition.py | 1 + ...managed_property.py => v0079_add_gateway_managed_property.py} | 0 2 files changed, 1 insertion(+) rename sqlmesh/migrations/{v0078_add_gateway_managed_property.py => v0079_add_gateway_managed_property.py} (100%) diff --git a/sqlmesh/core/model/definition.py b/sqlmesh/core/model/definition.py index ab3eeec808..56a6cdea93 100644 --- a/sqlmesh/core/model/definition.py +++ b/sqlmesh/core/model/definition.py @@ -1907,6 +1907,7 @@ def create_models_from_blueprints( else: gateway_name = None + # We pop to avoid pydantic validation issues since catalogs is not a model property if ( (catalogs := loader_kwargs.pop("catalogs", None)) and gateway_name diff --git a/sqlmesh/migrations/v0078_add_gateway_managed_property.py b/sqlmesh/migrations/v0079_add_gateway_managed_property.py similarity index 100% rename from sqlmesh/migrations/v0078_add_gateway_managed_property.py rename to sqlmesh/migrations/v0079_add_gateway_managed_property.py From 83ab3d7ac928a50a6e6c950b027533b28a3742e4 Mon Sep 17 00:00:00 2001 From: Themis Valtinos <73662635+themisvaltinos@users.noreply.github.com> Date: Fri, 11 Apr 2025 18:44:50 +0300 Subject: [PATCH 08/14] Refactors; account for identical db names between warehouses when creating schemas --- sqlmesh/core/context.py | 12 ++--- sqlmesh/core/loader.py | 4 +- sqlmesh/core/model/decorator.py | 2 + sqlmesh/core/model/definition.py | 8 +-- sqlmesh/core/plan/evaluator.py | 4 +- sqlmesh/core/snapshot/evaluator.py | 69 ++++++++++++++---------- sqlmesh/core/state_sync/common.py | 8 +-- sqlmesh/engines/commands.py | 2 +- tests/core/state_sync/test_state_sync.py | 2 +- 9 files changed, 62 insertions(+), 49 deletions(-) diff --git a/sqlmesh/core/context.py b/sqlmesh/core/context.py index 88ffdd2138..0b911e0691 100644 --- a/sqlmesh/core/context.py +++ b/sqlmesh/core/context.py @@ -364,7 +364,7 @@ def __init__( self._environment_statements: t.List[EnvironmentStatements] = [] self._excluded_requirements: t.Set[str] = set() self._default_catalog: t.Optional[str] = None - self._catalogs: t.Dict[str, str] = {} + self._default_catalog_per_gateway: t.Dict[str, str] = {} self._linters: t.Dict[str, Linter] = {} self._loaded: bool = False @@ -2225,15 +2225,15 @@ def engine_adapters(self) -> t.Dict[str, EngineAdapter]: return self._engine_adapters @cached_property - def catalogs(self) -> t.Dict[str, str]: + def default_catalog_per_gateway(self) -> t.Dict[str, str]: """Returns the catalogs for each engine adapter in a multi virtual layer setup when the catalog isn't shared.""" if self.gateway_managed_virtual_layer: - self._catalogs = { + self._default_catalog_per_gateway = { name: adapter.default_catalog for name, adapter in self.engine_adapters.items() if adapter.default_catalog } - return self._catalogs + return self._default_catalog_per_gateway def _get_engine_adapter(self, gateway: t.Optional[str] = None) -> EngineAdapter: if gateway: @@ -2317,10 +2317,10 @@ def _cleanup_environments(self) -> None: expired_environments = self.state_sync.delete_expired_environments() cleanup_expired_views( - adapter=self.engine_adapter, + default_adapter=self.engine_adapter, + engine_adapters=self.engine_adapters, environments=expired_environments, console=self.console, - engine_adapters=self.engine_adapters, ) def _try_connection(self, connection_name: str, validator: t.Callable[[], None]) -> None: diff --git a/sqlmesh/core/loader.py b/sqlmesh/core/loader.py index e9ee4282f6..a80fd8a3bd 100644 --- a/sqlmesh/core/loader.py +++ b/sqlmesh/core/loader.py @@ -468,7 +468,7 @@ def _load() -> t.List[Model]: default_catalog=self.context.default_catalog, infer_names=self.config.model_naming.infer_names, signal_definitions=signals, - catalogs=self.context.catalogs, + default_catalog_per_gateway=self.context.default_catalog_per_gateway, ) except Exception as ex: raise ConfigError(f"Failed to load model definition at '{path}'.\n{ex}") @@ -526,7 +526,7 @@ def _load_python_models( default_catalog=self.context.default_catalog, infer_names=self.config.model_naming.infer_names, audit_definitions=audits, - catalogs=self.context.catalogs, + default_catalog_per_gateway=self.context.default_catalog_per_gateway, ): if model.enabled: models[model.fqn] = model diff --git a/sqlmesh/core/model/decorator.py b/sqlmesh/core/model/decorator.py index 4193ac099d..acc602783f 100644 --- a/sqlmesh/core/model/decorator.py +++ b/sqlmesh/core/model/decorator.py @@ -93,6 +93,7 @@ def models( path: Path, module_path: Path, dialect: t.Optional[str] = None, + default_catalog_per_gateway: t.Optional[t.Dict[str, str]] = None, **loader_kwargs: t.Any, ) -> t.List[Model]: return create_models_from_blueprints( @@ -103,6 +104,7 @@ def models( path=path, module_path=module_path, dialect=dialect, + default_catalog_per_gateway=default_catalog_per_gateway, **loader_kwargs, ) diff --git a/sqlmesh/core/model/definition.py b/sqlmesh/core/model/definition.py index 56a6cdea93..2c3c59c8a6 100644 --- a/sqlmesh/core/model/definition.py +++ b/sqlmesh/core/model/definition.py @@ -1886,6 +1886,7 @@ def create_models_from_blueprints( path: Path = Path(), module_path: Path = Path(), dialect: DialectType = None, + default_catalog_per_gateway: t.Optional[t.Dict[str, str]] = None, **loader_kwargs: t.Any, ) -> t.List[Model]: model_blueprints: t.List[Model] = [] @@ -1907,11 +1908,10 @@ def create_models_from_blueprints( else: gateway_name = None - # We pop to avoid pydantic validation issues since catalogs is not a model property if ( - (catalogs := loader_kwargs.pop("catalogs", None)) + default_catalog_per_gateway and gateway_name - and (catalog := catalogs.get(gateway_name)) + and (catalog := default_catalog_per_gateway.get(gateway_name)) ): loader_kwargs["default_catalog"] = catalog @@ -1935,6 +1935,7 @@ def load_sql_based_models( path: Path = Path(), module_path: Path = Path(), dialect: DialectType = None, + default_catalog_per_gateway: t.Optional[t.Dict[str, str]] = None, **loader_kwargs: t.Any, ) -> t.List[Model]: gateway: t.Optional[exp.Expression] = None @@ -1972,6 +1973,7 @@ def load_sql_based_models( path=path, module_path=module_path, dialect=dialect, + default_catalog_per_gateway=default_catalog_per_gateway, **loader_kwargs, ) diff --git a/sqlmesh/core/plan/evaluator.py b/sqlmesh/core/plan/evaluator.py index 4e4d7e5bfa..6faed438f5 100644 --- a/sqlmesh/core/plan/evaluator.py +++ b/sqlmesh/core/plan/evaluator.py @@ -424,9 +424,7 @@ def _demote_snapshots( on_complete: t.Optional[t.Callable[[SnapshotInfoLike], None]] = None, ) -> None: self.snapshot_evaluator.demote( - target_snapshots, - environment_naming_info, - on_complete=on_complete, + target_snapshots, environment_naming_info, on_complete=on_complete ) def _restate(self, plan: EvaluatablePlan, snapshots_by_name: t.Dict[str, Snapshot]) -> None: diff --git a/sqlmesh/core/snapshot/evaluator.py b/sqlmesh/core/snapshot/evaluator.py index 92a0eaa5a9..661856da90 100644 --- a/sqlmesh/core/snapshot/evaluator.py +++ b/sqlmesh/core/snapshot/evaluator.py @@ -227,21 +227,20 @@ def promote( on_complete: A callback to call on each successfully promoted snapshot. """ - gateway_by_schema: t.Dict[t.Any, str] = {} - tables: t.List[t.Any] = [] + tables_by_gateway: t.Dict[t.Union[str, None], t.List[exp.Table]] = defaultdict(list) for snapshot in target_snapshots: if snapshot.is_model and not snapshot.is_symbolic: + gateway = ( + snapshot.model_gateway if environment_naming_info.gateway_managed else None + ) + adapter = self._get_adapter(gateway) table = snapshot.qualified_view_name.table_for_environment( - environment_naming_info, - dialect=self._get_adapter(snapshot.model_gateway).dialect - if environment_naming_info.gateway_managed - else self.adapter.dialect, + environment_naming_info, dialect=adapter.dialect ) - tables.append(table) - if environment_naming_info.gateway_managed: - table_schema = d.schema_(table.db, catalog=table.catalog) - gateway_by_schema[table_schema] = snapshot.model_gateway or "" - self._create_schemas(tables=tables, gateways=gateway_by_schema) + tables_by_gateway[gateway].append(table) + + for gateway, tables in tables_by_gateway.items(): + self._create_schemas(tables=tables, gateway=gateway) deployability_index = deployability_index or DeployabilityIndex.all_deployable() with self.concurrent_context(): @@ -301,8 +300,9 @@ def create( allow_destructive_snapshots: Set of snapshots that are allowed to have destructive schema changes. """ snapshots_with_table_names = defaultdict(set) - tables_by_schema = defaultdict(set) - gateway_by_schema: t.Dict[exp.Table, str] = {} + tables_by_gateway_and_schema: t.Dict[t.Union[str, None], t.Dict[exp.Table, set[str]]] = ( + defaultdict(lambda: defaultdict(set)) + ) table_deployability: t.Dict[str, bool] = {} allow_destructive_snapshots = allow_destructive_snapshots or set() @@ -324,24 +324,32 @@ def create( snapshots_with_table_names[snapshot].add(table.name) table_deployability[table.name] = is_deployable table_schema = d.schema_(table.db, catalog=table.catalog) - tables_by_schema[table_schema].add(table.name) - gateway_by_schema[table_schema] = snapshot.model.gateway or "" + tables_by_gateway_and_schema[snapshot.model_gateway][table_schema].add(table.name) - def _get_data_objects(schema: exp.Table, gateway: t.Optional[str] = None) -> t.Set[str]: + def _get_data_objects( + schema: exp.Table, + object_names: t.Optional[t.Set[str]] = None, + gateway: t.Optional[str] = None, + ) -> t.Set[str]: logger.info("Listing data objects in schema %s", schema.sql()) - objs = self.get_adapter(gateway).get_data_objects(schema, tables_by_schema[schema]) + objs = self._get_adapter(gateway).get_data_objects(schema, object_names) return {obj.name for obj in objs} with self.concurrent_context(): - existing_objects = { - obj - for objs in concurrent_apply_to_values( - list(tables_by_schema), - lambda s: _get_data_objects(s, gateway_by_schema[s]), - self.ddl_concurrent_tasks, - ) - for obj in objs - } + existing_objects: t.Set[str] = set() + for gateway, tables_by_schema in tables_by_gateway_and_schema.items(): + objs_for_gateway = { + obj + for objs in concurrent_apply_to_values( + list(tables_by_schema), + lambda s: _get_data_objects( + schema=s, object_names=tables_by_schema.get(s), gateway=gateway + ), + self.ddl_concurrent_tasks, + ) + for obj in objs + } + existing_objects.update(objs_for_gateway) snapshots_to_create = [] target_deployability_flags: t.Dict[str, t.List[bool]] = defaultdict(list) @@ -359,7 +367,10 @@ def _get_data_objects(schema: exp.Table, gateway: t.Optional[str] = None) -> t.S return if on_start: on_start(len(snapshots_to_create)) - self._create_schemas(tables_by_schema, gateway_by_schema) + + for gateway, tables_by_schema in tables_by_gateway_and_schema.items(): + self._create_schemas(tables=tables_by_schema, gateway=gateway) + self._create_snapshots( snapshots_to_create=snapshots_to_create, snapshots=snapshots, @@ -1075,7 +1086,7 @@ def _audit( def _create_schemas( self, tables: t.Iterable[t.Union[exp.Table, str]], - gateways: t.Optional[t.Dict[exp.Table, str]] = None, + gateway: t.Optional[str] = None, ) -> None: table_exprs = [exp.to_table(t) for t in tables] unique_schemas = {(t.args["db"], t.args.get("catalog")) for t in table_exprs if t and t.db} @@ -1084,7 +1095,7 @@ def _create_schemas( for schema_name, catalog in unique_schemas: schema = schema_(schema_name, catalog) logger.info("Creating schema '%s'", schema) - adapter = self.get_adapter(gateways.get(schema)) if gateways else self.adapter + adapter = self._get_adapter(gateway) adapter.create_schema(schema) def get_adapter(self, gateway: t.Optional[str] = None) -> EngineAdapter: diff --git a/sqlmesh/core/state_sync/common.py b/sqlmesh/core/state_sync/common.py index 6fc4a2e405..7c7f444c0c 100644 --- a/sqlmesh/core/state_sync/common.py +++ b/sqlmesh/core/state_sync/common.py @@ -21,10 +21,10 @@ def cleanup_expired_views( - adapter: EngineAdapter, + default_adapter: EngineAdapter, + engine_adapters: t.Dict[str, EngineAdapter], environments: t.List[Environment], console: t.Optional[Console] = None, - engine_adapters: t.Optional[t.Dict[str, EngineAdapter]] = None, ) -> None: expired_schema_environments = [ environment for environment in environments if environment.suffix_target.is_schema @@ -36,8 +36,8 @@ def cleanup_expired_views( # We have to use the corresponding adapter if the virtual layer is gateway managed def get_adapter(gateway_managed: bool, gateway: t.Optional[str] = None) -> EngineAdapter: if gateway_managed and gateway: - return (engine_adapters or {}).get(gateway, adapter) - return adapter + return engine_adapters.get(gateway, default_adapter) + return default_adapter # Drop the schemas for the expired environments for engine_adapter, expired_catalog, expired_schema in { diff --git a/sqlmesh/engines/commands.py b/sqlmesh/engines/commands.py index 6b144ee9e5..c16e759cbf 100644 --- a/sqlmesh/engines/commands.py +++ b/sqlmesh/engines/commands.py @@ -140,7 +140,7 @@ def cleanup( if isinstance(command_payload, str): command_payload = CleanupCommandPayload.parse_raw(command_payload) - cleanup_expired_views(evaluator.adapter, command_payload.environments) + cleanup_expired_views(evaluator.adapter, evaluator.adapters, command_payload.environments) evaluator.cleanup(command_payload.tasks) diff --git a/tests/core/state_sync/test_state_sync.py b/tests/core/state_sync/test_state_sync.py index b19869a0e5..1704719c7b 100644 --- a/tests/core/state_sync/test_state_sync.py +++ b/tests/core/state_sync/test_state_sync.py @@ -2635,7 +2635,7 @@ def test_cleanup_expired_views( previous_plan_id="test_plan_id", catalog_name_override="catalog_override", ) - cleanup_expired_views(adapter, [schema_environment, table_environment]) + cleanup_expired_views(adapter, {}, [schema_environment, table_environment]) assert adapter.drop_schema.called assert adapter.drop_view.called assert adapter.drop_schema.call_args_list == [ From 0197c269827fe1ac926b5bf9bdf9d2d68d2f810b Mon Sep 17 00:00:00 2001 From: Iaroslav Zeigerman Date: Fri, 11 Apr 2025 09:29:39 -0700 Subject: [PATCH 09/14] A few fixes (#4129) --- sqlmesh/core/context.py | 4 ++-- sqlmesh/core/context_diff.py | 8 ++++++-- sqlmesh/core/state_sync/db/facade.py | 6 ++---- 3 files changed, 10 insertions(+), 8 deletions(-) diff --git a/sqlmesh/core/context.py b/sqlmesh/core/context.py index 0b911e0691..6cfe4676e7 100644 --- a/sqlmesh/core/context.py +++ b/sqlmesh/core/context.py @@ -364,7 +364,7 @@ def __init__( self._environment_statements: t.List[EnvironmentStatements] = [] self._excluded_requirements: t.Set[str] = set() self._default_catalog: t.Optional[str] = None - self._default_catalog_per_gateway: t.Dict[str, str] = {} + self._default_catalog_per_gateway: t.Optional[t.Dict[str, str]] = None self._linters: t.Dict[str, Linter] = {} self._loaded: bool = False @@ -2227,7 +2227,7 @@ def engine_adapters(self) -> t.Dict[str, EngineAdapter]: @cached_property def default_catalog_per_gateway(self) -> t.Dict[str, str]: """Returns the catalogs for each engine adapter in a multi virtual layer setup when the catalog isn't shared.""" - if self.gateway_managed_virtual_layer: + if self._default_catalog_per_gateway is None: self._default_catalog_per_gateway = { name: adapter.default_catalog for name, adapter in self.engine_adapters.items() diff --git a/sqlmesh/core/context_diff.py b/sqlmesh/core/context_diff.py index eb05a61155..c7e9dea9a0 100644 --- a/sqlmesh/core/context_diff.py +++ b/sqlmesh/core/context_diff.py @@ -53,7 +53,9 @@ class ContextDiff(PydanticModel): """Whether the currently stored environment record is in unfinalized state.""" normalize_environment_name: bool """Whether the environment name should be normalized.""" - gateway_managed_virtual_layer: bool = False + previous_gateway_managed_virtual_layer: bool + """Whether the previous environment's virtual layer's views were created by the model specified gateways.""" + gateway_managed_virtual_layer: bool """Whether the virtual layer's views will be created by the model specified gateways.""" create_from: str """The name of the environment the target environment will be created from if new.""" @@ -121,7 +123,7 @@ def create( env = state_reader.get_environment(environment) create_from_env_exists = False - if env is None or env.expired or env.gateway_managed != gateway_managed_virtual_layer: + if env is None or env.expired: env = state_reader.get_environment(create_from.lower()) if not env and create_from != c.PROD: @@ -229,6 +231,7 @@ def create( diff_rendered=diff_rendered, previous_environment_statements=previous_environment_statements, environment_statements=environment_statements, + previous_gateway_managed_virtual_layer=env.gateway_managed if env else False, gateway_managed_virtual_layer=gateway_managed_virtual_layer, ) @@ -277,6 +280,7 @@ def has_changes(self) -> bool: or self.is_unfinalized_environment or self.has_requirement_changes or self.has_environment_statements_changes + or self.previous_gateway_managed_virtual_layer != self.gateway_managed_virtual_layer ) @property diff --git a/sqlmesh/core/state_sync/db/facade.py b/sqlmesh/core/state_sync/db/facade.py index e5f7493e2a..9e021995eb 100644 --- a/sqlmesh/core/state_sync/db/facade.py +++ b/sqlmesh/core/state_sync/db/facade.py @@ -199,10 +199,7 @@ def promote( ) != table_infos[name].qualified_view_name.for_environment(environment.naming_info) } - if ( - not existing_environment.expired - and existing_environment.gateway_managed == environment.gateway_managed - ): + if not existing_environment.expired: if environment.previous_plan_id != existing_environment.plan_id: raise ConflictingPlanError( f"Plan '{environment.plan_id}' is no longer valid for the target environment '{environment.name}'. " @@ -229,6 +226,7 @@ def promote( existing_environment and existing_environment.finalized_ts and not existing_environment.expired + and existing_environment.gateway_managed == environment.gateway_managed ): # Only promote new snapshots. added_table_infos -= set(existing_environment.promoted_snapshots) From b0a74dfd5e5a1839f44d66fcbd73f84197e8756d Mon Sep 17 00:00:00 2001 From: Themis Valtinos <73662635+themisvaltinos@users.noreply.github.com> Date: Fri, 11 Apr 2025 21:20:42 +0300 Subject: [PATCH 10/14] refactor; add comments --- sqlmesh/core/context.py | 2 +- sqlmesh/core/context_diff.py | 2 + sqlmesh/core/snapshot/evaluator.py | 2 + tests/core/test_context.py | 10 +- tests/core/test_integration.py | 17 +++- tests/core/test_plan.py | 98 +++++++++++++++++++ .../multi_virtual_layer/audits/.gitkeep | 0 .../fixtures}/multi_virtual_layer/config.yaml | 0 .../multi_virtual_layer/macros/.gitkeep | 0 .../multi_virtual_layer/macros/__init__.py | 0 .../multi_virtual_layer/models/.gitkeep | 0 .../models/local_schema/model_one.sql | 0 .../models/local_schema/model_two.sql | 0 .../models/memory_schema/model_one.sql | 0 .../models/memory_schema/model_two.sql | 0 .../multi_virtual_layer/tests/.gitkeep | 0 16 files changed, 121 insertions(+), 10 deletions(-) rename {examples => tests/fixtures}/multi_virtual_layer/audits/.gitkeep (100%) rename {examples => tests/fixtures}/multi_virtual_layer/config.yaml (100%) rename {examples => tests/fixtures}/multi_virtual_layer/macros/.gitkeep (100%) rename {examples => tests/fixtures}/multi_virtual_layer/macros/__init__.py (100%) rename {examples => tests/fixtures}/multi_virtual_layer/models/.gitkeep (100%) rename {examples => tests/fixtures}/multi_virtual_layer/models/local_schema/model_one.sql (100%) rename {examples => tests/fixtures}/multi_virtual_layer/models/local_schema/model_two.sql (100%) rename {examples => tests/fixtures}/multi_virtual_layer/models/memory_schema/model_one.sql (100%) rename {examples => tests/fixtures}/multi_virtual_layer/models/memory_schema/model_two.sql (100%) rename {examples => tests/fixtures}/multi_virtual_layer/tests/.gitkeep (100%) diff --git a/sqlmesh/core/context.py b/sqlmesh/core/context.py index 6cfe4676e7..f49a739fa3 100644 --- a/sqlmesh/core/context.py +++ b/sqlmesh/core/context.py @@ -2226,7 +2226,7 @@ def engine_adapters(self) -> t.Dict[str, EngineAdapter]: @cached_property def default_catalog_per_gateway(self) -> t.Dict[str, str]: - """Returns the catalogs for each engine adapter in a multi virtual layer setup when the catalog isn't shared.""" + """Returns the default catalogs for each engine adapter.""" if self._default_catalog_per_gateway is None: self._default_catalog_per_gateway = { name: adapter.default_catalog diff --git a/sqlmesh/core/context_diff.py b/sqlmesh/core/context_diff.py index c7e9dea9a0..ec03bf46bf 100644 --- a/sqlmesh/core/context_diff.py +++ b/sqlmesh/core/context_diff.py @@ -270,6 +270,8 @@ def create_no_diff(cls, environment: str, state_reader: StateReader) -> ContextD previous_requirements=env.requirements, requirements=env.requirements, previous_environment_statements=[], + previous_gateway_managed_virtual_layer=env.gateway_managed, + gateway_managed_virtual_layer=env.gateway_managed, ) @property diff --git a/sqlmesh/core/snapshot/evaluator.py b/sqlmesh/core/snapshot/evaluator.py index 661856da90..36ead0341a 100644 --- a/sqlmesh/core/snapshot/evaluator.py +++ b/sqlmesh/core/snapshot/evaluator.py @@ -239,6 +239,7 @@ def promote( ) tables_by_gateway[gateway].append(table) + # A schema can be shared across multiple engines, so we need to group by gateway for gateway, tables in tables_by_gateway.items(): self._create_schemas(tables=tables, gateway=gateway) @@ -337,6 +338,7 @@ def _get_data_objects( with self.concurrent_context(): existing_objects: t.Set[str] = set() + # A schema can be shared across multiple engines, so we need to group tables by both gateway and schema for gateway, tables_by_schema in tables_by_gateway_and_schema.items(): objs_for_gateway = { obj diff --git a/tests/core/test_context.py b/tests/core/test_context.py index 1c512b83d9..4593e2ca33 100644 --- a/tests/core/test_context.py +++ b/tests/core/test_context.py @@ -327,21 +327,15 @@ def test_evaluate_limit(): def test_gateway_specific_adapters(copy_to_temp_path, mocker): path = copy_to_temp_path("examples/sushi") ctx = Context(paths=path, config="isolated_systems_config", gateway="prod") - assert len(ctx._engine_adapters) == 1 + assert len(ctx._engine_adapters) == 3 assert ctx.engine_adapter == ctx._engine_adapters["prod"] - - with pytest.raises(SQLMeshError): - assert ctx._get_engine_adapter("non_existing") - - # This will create the requested engine adapter assert ctx._get_engine_adapter("dev") == ctx._engine_adapters["dev"] ctx = Context(paths=path, config="isolated_systems_config") - assert len(ctx._engine_adapters) == 1 + assert len(ctx._engine_adapters) == 3 assert ctx.engine_adapter == ctx._engine_adapters["dev"] ctx = Context(paths=path, config="isolated_systems_config") - assert len(ctx.engine_adapters) == 3 assert ctx.engine_adapter == ctx._get_engine_adapter() assert ctx._get_engine_adapter("test") == ctx._engine_adapters["test"] diff --git a/tests/core/test_integration.py b/tests/core/test_integration.py index e4ab90cbbc..5498000b82 100644 --- a/tests/core/test_integration.py +++ b/tests/core/test_integration.py @@ -11,6 +11,7 @@ import pytest from pathlib import Path import os +from sqlmesh.utils.concurrency import NodeExecutionFailedError import time_machine from pytest_mock.plugin import MockerFixture from sqlglot import exp @@ -4497,7 +4498,7 @@ def test_multi(mocker): @use_terminal_console def test_multi_virtual_layer(mocker): - context = Context(paths=["examples/multi_virtual_layer"]) + context = Context(paths=["tests/fixtures/multi_virtual_layer"]) local_db = "db.duckdb" if os.path.exists(local_db): @@ -4569,6 +4570,20 @@ def test_multi_virtual_layer(mocker): == " item_id global_one macro_one extra\n0 gateway_2 88 1 c" ) + # Changing the flag should show a diff + context.gateway_managed_virtual_layer = False + plan = context.plan_builder().build() + assert not plan.requires_backfill + assert ( + plan.context_diff.previous_gateway_managed_virtual_layer + != plan.context_diff.gateway_managed_virtual_layer + ) + assert plan.context_diff.has_changes + + # This should error since the default_gateway won't have access to create the view on a non-shared catalog + with pytest.raises(NodeExecutionFailedError, match=r"Execution failed for node SnapshotId*"): + context.apply(plan) + if os.path.exists(local_db): os.remove(local_db) diff --git a/tests/core/test_plan.py b/tests/core/test_plan.py index dcb9876204..7bb598eecd 100644 --- a/tests/core/test_plan.py +++ b/tests/core/test_plan.py @@ -83,6 +83,8 @@ def test_forward_only_plan_sets_version(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) plan_builder = PlanBuilder(context_diff, DuckDBEngineAdapter.SCHEMA_DIFFER, forward_only=True) @@ -134,6 +136,8 @@ def test_forward_only_dev(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) yesterday_ds_mock = mocker.patch("sqlmesh.core.plan.builder.yesterday_ds") @@ -194,6 +198,8 @@ def test_forward_only_metadata_change_dev(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) yesterday_ds_mock = mocker.patch("sqlmesh.core.plan.builder.yesterday_ds") @@ -243,6 +249,8 @@ def test_forward_only_plan_added_models(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) PlanBuilder(context_diff, DuckDBEngineAdapter.SCHEMA_DIFFER, forward_only=True).build() @@ -287,6 +295,8 @@ def test_forward_only_plan_categorizes_change_model_kind_as_breaking( previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) PlanBuilder(context_diff, DuckDBEngineAdapter.SCHEMA_DIFFER, forward_only=True).build() @@ -333,6 +343,8 @@ def test_paused_forward_only_parent(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) PlanBuilder(context_diff, DuckDBEngineAdapter.SCHEMA_DIFFER, forward_only=False).build() @@ -362,6 +374,8 @@ def test_forward_only_plan_allow_destructive_models( previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) with pytest.raises( @@ -436,6 +450,8 @@ def test_forward_only_plan_allow_destructive_models( previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) with pytest.raises( @@ -492,6 +508,8 @@ def test_forward_only_model_on_destructive_change( previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) with pytest.raises( @@ -550,6 +568,8 @@ def test_forward_only_model_on_destructive_change( previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) PlanBuilder(context_diff_2, schema_differ).build() @@ -634,6 +654,8 @@ def test_forward_only_model_on_destructive_change( previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) PlanBuilder(context_diff_3, schema_differ).build() @@ -668,6 +690,8 @@ def test_forward_only_model_on_destructive_change_no_column_types( previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) logger = logging.getLogger("sqlmesh.core.plan.builder") @@ -704,6 +728,8 @@ def test_missing_intervals_lookback(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) plan = Plan( @@ -896,6 +922,8 @@ def test_restate_symbolic_model(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) plan = PlanBuilder( @@ -930,6 +958,8 @@ def test_restate_seed_model(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) plan = PlanBuilder( @@ -954,6 +984,8 @@ def test_restate_missing_model(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) with pytest.raises( @@ -983,6 +1015,8 @@ def test_new_snapshots_with_restatements(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) with pytest.raises( @@ -1016,6 +1050,8 @@ def test_end_validation(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) schema_differ = DuckDBEngineAdapter.SCHEMA_DIFFER @@ -1081,6 +1117,8 @@ def test_forward_only_revert_not_allowed(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) schema_differ = DuckDBEngineAdapter.SCHEMA_DIFFER @@ -1139,6 +1177,8 @@ def test_forward_only_plan_seed_models(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) PlanBuilder(context_diff, DuckDBEngineAdapter.SCHEMA_DIFFER, forward_only=True).build() @@ -1173,6 +1213,8 @@ def test_start_inference(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) snapshot_b.add_interval("2022-01-01", now()) @@ -1211,6 +1253,8 @@ def test_auto_categorization(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) PlanBuilder(context_diff, DuckDBEngineAdapter.SCHEMA_DIFFER).build() @@ -1257,6 +1301,8 @@ def test_auto_categorization_missing_schema_downstream(make_snapshot, mocker: Mo previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) PlanBuilder(context_diff, DuckDBEngineAdapter.SCHEMA_DIFFER).build() @@ -1287,6 +1333,8 @@ def test_broken_references(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) # Make sure the downstream snapshot doesn't have any parents, @@ -1322,6 +1370,8 @@ def test_broken_references_external_model(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) # Make sure the downstream snapshot doesn't have any parents, @@ -1363,6 +1413,8 @@ def test_effective_from(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) schema_differ = DuckDBEngineAdapter.SCHEMA_DIFFER @@ -1446,6 +1498,8 @@ def test_effective_from_non_evaluatble_model(make_snapshot, mocker: MockerFixtur previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) schema_differ = DuckDBEngineAdapter.SCHEMA_DIFFER @@ -1482,6 +1536,8 @@ def test_new_environment_no_changes(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) schema_differ = DuckDBEngineAdapter.SCHEMA_DIFFER @@ -1527,6 +1583,8 @@ def test_new_environment_with_changes(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) # Modified the existing model. @@ -1610,6 +1668,8 @@ def test_forward_only_models(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) schema_differ = DuckDBEngineAdapter.SCHEMA_DIFFER @@ -1654,6 +1714,8 @@ def test_forward_only_models_model_kind_changed(make_snapshot, mocker: MockerFix previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) PlanBuilder(context_diff, DuckDBEngineAdapter.SCHEMA_DIFFER, is_dev=True).build() @@ -1731,6 +1793,8 @@ def test_indirectly_modified_forward_only_model(make_snapshot, mocker: MockerFix previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) plan = PlanBuilder(context_diff, DuckDBEngineAdapter.SCHEMA_DIFFER, is_dev=True).build() @@ -1785,6 +1849,8 @@ def test_added_model_with_forward_only_parent(make_snapshot, mocker: MockerFixtu previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) PlanBuilder(context_diff, DuckDBEngineAdapter.SCHEMA_DIFFER, is_dev=True).build() @@ -1823,6 +1889,8 @@ def test_added_forward_only_model(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) PlanBuilder(context_diff, DuckDBEngineAdapter.SCHEMA_DIFFER).build() @@ -1855,6 +1923,8 @@ def test_disable_restatement(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) schema_differ = DuckDBEngineAdapter.SCHEMA_DIFFER @@ -1922,6 +1992,8 @@ def test_revert_to_previous_value(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) plan_builder = PlanBuilder(context_diff, DuckDBEngineAdapter.SCHEMA_DIFFER) @@ -2134,6 +2206,8 @@ def test_add_restatements( previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) plan = PlanBuilder( @@ -2211,6 +2285,8 @@ def test_dev_plan_depends_past(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) schema_differ = DuckDBEngineAdapter.SCHEMA_DIFFER @@ -2314,6 +2390,8 @@ def test_dev_plan_depends_past_non_deployable(make_snapshot, mocker: MockerFixtu previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) schema_differ = DuckDBEngineAdapter.SCHEMA_DIFFER @@ -2380,6 +2458,8 @@ def test_models_selected_for_backfill(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) schema_differ = DuckDBEngineAdapter.SCHEMA_DIFFER @@ -2432,6 +2512,8 @@ def test_categorized_uncategorized(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) plan_builder = PlanBuilder( @@ -2487,6 +2569,8 @@ def test_environment_previous_finalized_snapshots(make_snapshot, mocker: MockerF previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=[snapshot_c.table_info, snapshot_d.table_info], + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) schema_differ = DuckDBEngineAdapter.SCHEMA_DIFFER @@ -2541,6 +2625,8 @@ def test_metadata_change(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) plan = PlanBuilder(context_diff, DuckDBEngineAdapter.SCHEMA_DIFFER, is_dev=True).build() @@ -2581,6 +2667,8 @@ def test_plan_start_when_preview_enabled(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) default_start_for_preview = "2024-06-09" @@ -2630,6 +2718,8 @@ def test_interval_end_per_model(make_snapshot): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) plan_builder = PlanBuilder( @@ -2705,6 +2795,8 @@ def test_unaligned_start_model_with_forward_only_preview(make_snapshot): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) plan_builder = PlanBuilder( @@ -2755,6 +2847,8 @@ def test_restate_production_model_in_dev(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) mock_console = mocker.Mock() @@ -2856,6 +2950,8 @@ def test_restate_daily_to_monthly(make_snapshot, mocker: MockerFixture): previous_plan_id=None, previously_promoted_snapshot_ids=set(), previous_finalized_snapshots=None, + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) schema_differ = DuckDBEngineAdapter.SCHEMA_DIFFER @@ -2909,6 +3005,8 @@ def test_plan_environment_statements_diff(make_snapshot): python_env={}, ) ], + previous_gateway_managed_virtual_layer=False, + gateway_managed_virtual_layer=False, ) assert context_diff.has_changes diff --git a/examples/multi_virtual_layer/audits/.gitkeep b/tests/fixtures/multi_virtual_layer/audits/.gitkeep similarity index 100% rename from examples/multi_virtual_layer/audits/.gitkeep rename to tests/fixtures/multi_virtual_layer/audits/.gitkeep diff --git a/examples/multi_virtual_layer/config.yaml b/tests/fixtures/multi_virtual_layer/config.yaml similarity index 100% rename from examples/multi_virtual_layer/config.yaml rename to tests/fixtures/multi_virtual_layer/config.yaml diff --git a/examples/multi_virtual_layer/macros/.gitkeep b/tests/fixtures/multi_virtual_layer/macros/.gitkeep similarity index 100% rename from examples/multi_virtual_layer/macros/.gitkeep rename to tests/fixtures/multi_virtual_layer/macros/.gitkeep diff --git a/examples/multi_virtual_layer/macros/__init__.py b/tests/fixtures/multi_virtual_layer/macros/__init__.py similarity index 100% rename from examples/multi_virtual_layer/macros/__init__.py rename to tests/fixtures/multi_virtual_layer/macros/__init__.py diff --git a/examples/multi_virtual_layer/models/.gitkeep b/tests/fixtures/multi_virtual_layer/models/.gitkeep similarity index 100% rename from examples/multi_virtual_layer/models/.gitkeep rename to tests/fixtures/multi_virtual_layer/models/.gitkeep diff --git a/examples/multi_virtual_layer/models/local_schema/model_one.sql b/tests/fixtures/multi_virtual_layer/models/local_schema/model_one.sql similarity index 100% rename from examples/multi_virtual_layer/models/local_schema/model_one.sql rename to tests/fixtures/multi_virtual_layer/models/local_schema/model_one.sql diff --git a/examples/multi_virtual_layer/models/local_schema/model_two.sql b/tests/fixtures/multi_virtual_layer/models/local_schema/model_two.sql similarity index 100% rename from examples/multi_virtual_layer/models/local_schema/model_two.sql rename to tests/fixtures/multi_virtual_layer/models/local_schema/model_two.sql diff --git a/examples/multi_virtual_layer/models/memory_schema/model_one.sql b/tests/fixtures/multi_virtual_layer/models/memory_schema/model_one.sql similarity index 100% rename from examples/multi_virtual_layer/models/memory_schema/model_one.sql rename to tests/fixtures/multi_virtual_layer/models/memory_schema/model_one.sql diff --git a/examples/multi_virtual_layer/models/memory_schema/model_two.sql b/tests/fixtures/multi_virtual_layer/models/memory_schema/model_two.sql similarity index 100% rename from examples/multi_virtual_layer/models/memory_schema/model_two.sql rename to tests/fixtures/multi_virtual_layer/models/memory_schema/model_two.sql diff --git a/examples/multi_virtual_layer/tests/.gitkeep b/tests/fixtures/multi_virtual_layer/tests/.gitkeep similarity index 100% rename from examples/multi_virtual_layer/tests/.gitkeep rename to tests/fixtures/multi_virtual_layer/tests/.gitkeep From eee257e96eafdc385b83538fba60855e8f045010 Mon Sep 17 00:00:00 2001 From: Themis Valtinos <73662635+themisvaltinos@users.noreply.github.com> Date: Fri, 11 Apr 2025 21:45:13 +0300 Subject: [PATCH 11/14] fix rebase --- sqlmesh/core/snapshot/evaluator.py | 12 ++++++------ 1 file changed, 6 insertions(+), 6 deletions(-) diff --git a/sqlmesh/core/snapshot/evaluator.py b/sqlmesh/core/snapshot/evaluator.py index 36ead0341a..86450b8a93 100644 --- a/sqlmesh/core/snapshot/evaluator.py +++ b/sqlmesh/core/snapshot/evaluator.py @@ -233,7 +233,7 @@ def promote( gateway = ( snapshot.model_gateway if environment_naming_info.gateway_managed else None ) - adapter = self._get_adapter(gateway) + adapter = self.get_adapter(gateway) table = snapshot.qualified_view_name.table_for_environment( environment_naming_info, dialect=adapter.dialect ) @@ -333,7 +333,7 @@ def _get_data_objects( gateway: t.Optional[str] = None, ) -> t.Set[str]: logger.info("Listing data objects in schema %s", schema.sql()) - objs = self._get_adapter(gateway).get_data_objects(schema, object_names) + objs = self.get_adapter(gateway).get_data_objects(schema, object_names) return {obj.name for obj in objs} with self.concurrent_context(): @@ -457,7 +457,7 @@ def cleanup( lambda s: self._cleanup_snapshot( s, snapshots_to_dev_table_only[s.snapshot_id], - self._get_adapter(s.model_gateway), + self.get_adapter(s.model_gateway), on_complete, ), self.ddl_concurrent_tasks, @@ -942,7 +942,7 @@ def _promote_snapshot( ) -> None: if snapshot.is_model: adapter = ( - self._get_adapter(snapshot.model_gateway) + self.get_adapter(snapshot.model_gateway) if environment_naming_info.gateway_managed else self.adapter ) @@ -979,7 +979,7 @@ def _demote_snapshot( on_complete: t.Optional[t.Callable[[SnapshotInfoLike], None]], ) -> None: adapter = ( - self._get_adapter(snapshot.model_gateway) + self.get_adapter(snapshot.model_gateway) if environment_naming_info.gateway_managed else self.adapter ) @@ -1097,7 +1097,7 @@ def _create_schemas( for schema_name, catalog in unique_schemas: schema = schema_(schema_name, catalog) logger.info("Creating schema '%s'", schema) - adapter = self._get_adapter(gateway) + adapter = self.get_adapter(gateway) adapter.create_schema(schema) def get_adapter(self, gateway: t.Optional[str] = None) -> EngineAdapter: From 1a6d84b4fd5e289d617f0cd7b5f592aceec156a6 Mon Sep 17 00:00:00 2001 From: Themis Valtinos <73662635+themisvaltinos@users.noreply.github.com> Date: Sat, 12 Apr 2025 00:48:43 +0300 Subject: [PATCH 12/14] extend integration test --- tests/core/test_integration.py | 30 ++++++++++++++++++++++++++++++ 1 file changed, 30 insertions(+) diff --git a/tests/core/test_integration.py b/tests/core/test_integration.py index 5498000b82..aecec06398 100644 --- a/tests/core/test_integration.py +++ b/tests/core/test_integration.py @@ -4570,6 +4570,36 @@ def test_multi_virtual_layer(mocker): == " item_id global_one macro_one extra\n0 gateway_2 88 1 c" ) + # Create dev environment + model = context.get_model("db.local_schema.model_one") + context.upsert_model(model.copy(update={"query": model.query.select("'d' AS extra")})) + plan = context.plan_builder("dev").build() + context.apply(plan) + + dev_environment = context.state_sync.get_environment("dev") + assert dev_environment is not None + metadata = DuckDBMetadata.from_context(context) + start_schemas = set(metadata.schemas) + assert sorted(start_schemas) == sorted( + {"local_schema", "local_schema__dev", "sqlmesh", "sqlmesh__local_schema"} + ) + + # Invalidate dev environment + context.invalidate_environment("dev") + invalidate_environment = context.state_sync.get_environment("dev") + assert invalidate_environment is not None + schemas_prior_to_janitor = set(metadata.schemas) + assert invalidate_environment.expiration_ts < dev_environment.expiration_ts # type: ignore + assert sorted(start_schemas) == sorted(schemas_prior_to_janitor) + + # Run janitor + context._run_janitor() + removed_schemas = start_schemas - set(metadata.schemas) + assert context.state_sync.get_environment("dev") is None + assert removed_schemas == {"local_schema__dev"} + prod_environment = context.state_sync.get_environment("prod") + assert len(prod_environment.snapshots_) == 4 + # Changing the flag should show a diff context.gateway_managed_virtual_layer = False plan = context.plan_builder().build() From 3821729227fcc852c101460b41c8aee1b5766d3e Mon Sep 17 00:00:00 2001 From: Themis Valtinos <73662635+themisvaltinos@users.noreply.github.com> Date: Sat, 12 Apr 2025 20:00:04 +0300 Subject: [PATCH 13/14] use tmp_path in integration tests, clean up and extend them --- tests/core/test_integration.py | 106 ++++++++++++------ .../fixtures/multi_virtual_layer/config.yaml | 28 ----- .../model_one.sql | 0 .../model_two.sql | 2 +- .../model_one.sql | 2 +- .../model_two.sql | 4 +- 6 files changed, 74 insertions(+), 68 deletions(-) delete mode 100644 tests/fixtures/multi_virtual_layer/config.yaml rename tests/fixtures/multi_virtual_layer/models/{local_schema => first_schema}/model_one.sql (100%) rename tests/fixtures/multi_virtual_layer/models/{local_schema => first_schema}/model_two.sql (70%) rename tests/fixtures/multi_virtual_layer/models/{memory_schema => second_schema}/model_one.sql (86%) rename tests/fixtures/multi_virtual_layer/models/{memory_schema => second_schema}/model_two.sql (58%) diff --git a/tests/core/test_integration.py b/tests/core/test_integration.py index aecec06398..c59bf87b96 100644 --- a/tests/core/test_integration.py +++ b/tests/core/test_integration.py @@ -6,11 +6,12 @@ from unittest import mock from unittest.mock import patch +import os import numpy as np import pandas as pd import pytest from pathlib import Path -import os +from sqlmesh.core.config.naming import NameInferenceConfig from sqlmesh.utils.concurrency import NodeExecutionFailedError import time_machine from pytest_mock.plugin import MockerFixture @@ -4497,26 +4498,44 @@ def test_multi(mocker): @use_terminal_console -def test_multi_virtual_layer(mocker): - context = Context(paths=["tests/fixtures/multi_virtual_layer"]) +def test_multi_virtual_layer(copy_to_temp_path): + paths = copy_to_temp_path("tests/fixtures/multi_virtual_layer") + path = Path(paths[0]) + first_db_path = str(path / "db_1.db") + second_db_path = str(path / "db_2.db") - local_db = "db.duckdb" - if os.path.exists(local_db): - os.remove(local_db) + config = Config( + gateways={ + "first": GatewayConfig( + connection=DuckDBConnectionConfig(database=first_db_path), + variables={"overriden_var": "gateway_1"}, + ), + "second": GatewayConfig( + connection=DuckDBConnectionConfig(database=second_db_path), + variables={"overriden_var": "gateway_2"}, + ), + }, + model_defaults=ModelDefaultsConfig(dialect="duckdb"), + model_naming=NameInferenceConfig(infer_names=True), + default_gateway="first", + gateway_managed_virtual_layer=True, + variables={"overriden_var": "global", "global_one": 88}, + ) + + context = Context(paths=paths, config=config) # For the model without gateway the default should be used and the gateway variable should overide the global assert ( - context.render("local_schema.model_one").sql() + context.render("first_schema.model_one").sql() == 'SELECT \'gateway_1\' AS "item_id", 88 AS "global_one", 1 AS "macro_one"' ) # For model with gateway specified the appropriate variable should be used to overide assert ( - context.render("memory.memory_schema.model_one").sql() + context.render("db_2.second_schema.model_one").sql() == 'SELECT \'gateway_2\' AS "item_id", 88 AS "global_one", 1 AS "macro_one"' ) - # context._new_state_sync().reset(default_catalog=context.default_catalog) plan = context.plan_builder().build() assert len(plan.new_snapshots) == 4 context.apply(plan) @@ -4524,25 +4543,25 @@ def test_multi_virtual_layer(mocker): # Validate the tables that source from the first tables are correct as well with evaluate assert ( context.evaluate( - "local_schema.model_two", start=now(), end=now(), execution_time=now() + "first_schema.model_two", start=now(), end=now(), execution_time=now() ).to_string() == " item_id global_one\n0 gateway_1 88" ) assert ( context.evaluate( - "memory.memory_schema.model_two", start=now(), end=now(), execution_time=now() + "db_2.second_schema.model_two", start=now(), end=now(), execution_time=now() ).to_string() == " item_id global_one\n0 gateway_2 88" ) assert sorted(set(snapshot.name for snapshot in plan.directly_modified)) == [ - '"db"."local_schema"."model_one"', - '"db"."local_schema"."model_two"', - '"memory"."memory_schema"."model_one"', - '"memory"."memory_schema"."model_two"', + '"db_1"."first_schema"."model_one"', + '"db_1"."first_schema"."model_two"', + '"db_2"."second_schema"."model_one"', + '"db_2"."second_schema"."model_two"', ] - model = context.get_model("memory.memory_schema.model_one") + model = context.get_model("db_1.first_schema.model_one") context.upsert_model(model.copy(update={"query": model.query.select("'c' AS extra")})) plan = context.plan_builder().build() @@ -4553,52 +4572,70 @@ def test_multi_virtual_layer(mocker): assert state_environments[0].gateway_managed assert len(state_snapshots) == len(state_environments[0].snapshots) - assert [snapshot.name for snapshot in plan.directly_modified] == [ - '"memory"."memory_schema"."model_one"' + '"db_1"."first_schema"."model_one"' ] assert [x.name for x in list(plan.indirectly_modified.values())[0]] == [ - '"memory"."memory_schema"."model_two"' + '"db_1"."first_schema"."model_two"' ] assert len(plan.missing_intervals) == 1 - assert ( context.evaluate( - "memory.memory_schema.model_one", start=now(), end=now(), execution_time=now() + "db_1.first_schema.model_one", start=now(), end=now(), execution_time=now() ).to_string() - == " item_id global_one macro_one extra\n0 gateway_2 88 1 c" + == " item_id global_one macro_one extra\n0 gateway_1 88 1 c" ) - # Create dev environment - model = context.get_model("db.local_schema.model_one") + # Create dev environment with changed models + model = context.get_model("db_2.second_schema.model_one") context.upsert_model(model.copy(update={"query": model.query.select("'d' AS extra")})) + model = context.get_model("first_schema.model_two") + context.upsert_model(model.copy(update={"query": model.query.select("'d2' AS col")})) plan = context.plan_builder("dev").build() context.apply(plan) dev_environment = context.state_sync.get_environment("dev") assert dev_environment is not None - metadata = DuckDBMetadata.from_context(context) - start_schemas = set(metadata.schemas) - assert sorted(start_schemas) == sorted( - {"local_schema", "local_schema__dev", "sqlmesh", "sqlmesh__local_schema"} + + metadata_engine_1 = DuckDBMetadata.from_context(context) + start_schemas_1 = set(metadata_engine_1.schemas) + assert sorted(start_schemas_1) == sorted( + {"first_schema__dev", "sqlmesh", "first_schema", "sqlmesh__first_schema"} + ) + + metadata_engine_2 = DuckDBMetadata(context._get_engine_adapter("second")) + start_schemas_2 = set(metadata_engine_2.schemas) + assert sorted(start_schemas_2) == sorted( + {"sqlmesh__second_schema", "second_schema", "second_schema__dev"} ) # Invalidate dev environment context.invalidate_environment("dev") invalidate_environment = context.state_sync.get_environment("dev") assert invalidate_environment is not None - schemas_prior_to_janitor = set(metadata.schemas) assert invalidate_environment.expiration_ts < dev_environment.expiration_ts # type: ignore - assert sorted(start_schemas) == sorted(schemas_prior_to_janitor) + assert sorted(start_schemas_1) == sorted(set(metadata_engine_1.schemas)) + assert sorted(start_schemas_2) == sorted(set(metadata_engine_2.schemas)) # Run janitor context._run_janitor() - removed_schemas = start_schemas - set(metadata.schemas) assert context.state_sync.get_environment("dev") is None - assert removed_schemas == {"local_schema__dev"} + removed_schemas = start_schemas_1 - set(metadata_engine_1.schemas) + assert removed_schemas == {"first_schema__dev"} + removed_schemas = start_schemas_2 - set(metadata_engine_2.schemas) + assert removed_schemas == {"second_schema__dev"} prod_environment = context.state_sync.get_environment("prod") - assert len(prod_environment.snapshots_) == 4 + + # Remove the second gateway's second model and apply plan + second_model = path / "models/second_schema/model_two.sql" + os.remove(second_model) + assert not second_model.exists() + context = Context(paths=paths, config=config) + plan = context.plan_builder().build() + context.apply(plan) + prod_environment = context.state_sync.get_environment("prod") + assert len(prod_environment.snapshots_) == 3 # Changing the flag should show a diff context.gateway_managed_virtual_layer = False @@ -4614,9 +4651,6 @@ def test_multi_virtual_layer(mocker): with pytest.raises(NodeExecutionFailedError, match=r"Execution failed for node SnapshotId*"): context.apply(plan) - if os.path.exists(local_db): - os.remove(local_db) - def test_multi_dbt(mocker): context = Context(paths=["examples/multi_dbt/bronze", "examples/multi_dbt/silver"]) diff --git a/tests/fixtures/multi_virtual_layer/config.yaml b/tests/fixtures/multi_virtual_layer/config.yaml deleted file mode 100644 index 483472f16b..0000000000 --- a/tests/fixtures/multi_virtual_layer/config.yaml +++ /dev/null @@ -1,28 +0,0 @@ -gateways: - local: - connection: - type: duckdb - database: db.duckdb - variables: - overriden_var: 'gateway_1' - memory: - connection: - type: duckdb - variables: - overriden_var: 'gateway_2' - -default_gateway: local - -model_defaults: - dialect: 'duckdb' - -model_naming: - infer_names: True - -gateway_managed_virtual_layer: True - -variables: - overriden_var: 'global' - global_one: 88 - - diff --git a/tests/fixtures/multi_virtual_layer/models/local_schema/model_one.sql b/tests/fixtures/multi_virtual_layer/models/first_schema/model_one.sql similarity index 100% rename from tests/fixtures/multi_virtual_layer/models/local_schema/model_one.sql rename to tests/fixtures/multi_virtual_layer/models/first_schema/model_one.sql diff --git a/tests/fixtures/multi_virtual_layer/models/local_schema/model_two.sql b/tests/fixtures/multi_virtual_layer/models/first_schema/model_two.sql similarity index 70% rename from tests/fixtures/multi_virtual_layer/models/local_schema/model_two.sql rename to tests/fixtures/multi_virtual_layer/models/first_schema/model_two.sql index 93b927eee8..c09794f02f 100644 --- a/tests/fixtures/multi_virtual_layer/models/local_schema/model_two.sql +++ b/tests/fixtures/multi_virtual_layer/models/first_schema/model_two.sql @@ -6,4 +6,4 @@ SELECT item_id, global_one FROM - local_schema.model_one; \ No newline at end of file + first_schema.model_one; \ No newline at end of file diff --git a/tests/fixtures/multi_virtual_layer/models/memory_schema/model_one.sql b/tests/fixtures/multi_virtual_layer/models/second_schema/model_one.sql similarity index 86% rename from tests/fixtures/multi_virtual_layer/models/memory_schema/model_one.sql rename to tests/fixtures/multi_virtual_layer/models/second_schema/model_one.sql index 9c0fac206a..b4b75d80bf 100644 --- a/tests/fixtures/multi_virtual_layer/models/memory_schema/model_one.sql +++ b/tests/fixtures/multi_virtual_layer/models/second_schema/model_one.sql @@ -1,6 +1,6 @@ MODEL ( kind FULL, - gateway memory + gateway second ); SELECT diff --git a/tests/fixtures/multi_virtual_layer/models/memory_schema/model_two.sql b/tests/fixtures/multi_virtual_layer/models/second_schema/model_two.sql similarity index 58% rename from tests/fixtures/multi_virtual_layer/models/memory_schema/model_two.sql rename to tests/fixtures/multi_virtual_layer/models/second_schema/model_two.sql index c249aa37f6..f7688d70de 100644 --- a/tests/fixtures/multi_virtual_layer/models/memory_schema/model_two.sql +++ b/tests/fixtures/multi_virtual_layer/models/second_schema/model_two.sql @@ -1,10 +1,10 @@ MODEL ( kind FULL, - gateway memory + gateway second ); SELECT item_id, global_one FROM - memory_schema.model_one; \ No newline at end of file + second_schema.model_one; \ No newline at end of file From 64db448c22cf90a7b329fd79f6fa184e179d4db4 Mon Sep 17 00:00:00 2001 From: Themis Valtinos <73662635+themisvaltinos@users.noreply.github.com> Date: Mon, 14 Apr 2025 15:41:02 +0300 Subject: [PATCH 14/14] revise condition to check for none --- sqlmesh/core/model/definition.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/sqlmesh/core/model/definition.py b/sqlmesh/core/model/definition.py index 2c3c59c8a6..207f3853ae 100644 --- a/sqlmesh/core/model/definition.py +++ b/sqlmesh/core/model/definition.py @@ -1911,7 +1911,7 @@ def create_models_from_blueprints( if ( default_catalog_per_gateway and gateway_name - and (catalog := default_catalog_per_gateway.get(gateway_name)) + and (catalog := default_catalog_per_gateway.get(gateway_name)) is not None ): loader_kwargs["default_catalog"] = catalog